mirror of
https://github.com/netbirdio/netbird.git
synced 2026-07-31 21:01:29 +02:00
Compare commits
3 Commits
main
...
add-agents
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8f5fcf99d4 | ||
|
|
8947d0132f | ||
|
|
164d62f3a5 |
514
AGENTS.md
Normal file
514
AGENTS.md
Normal file
@@ -0,0 +1,514 @@
|
||||
# NetBird Agent Guidelines
|
||||
|
||||
**NetBird** is an open-source connectivity platform: a WireGuard®-based overlay
|
||||
network with a control plane. The **agent** (`client/`) runs on user machines as
|
||||
a privileged daemon and manages the WireGuard interface, routing, firewall, and
|
||||
DNS. **Management** (`management/`) is the control plane and REST/gRPC API,
|
||||
**Signal** (`signal/`) brokers peer handshakes, **Relay** (`relay/`) carries
|
||||
traffic when a direct tunnel is impossible, and **Proxy** (`proxy/`) is the
|
||||
identity-aware proxy behind Agent Network.
|
||||
|
||||
This file applies to the whole repository, and is the single source of truth for
|
||||
agent guidance here. `CLAUDE.md` is a one-line pointer to it — keep the guidance
|
||||
in this file, not duplicated there.
|
||||
|
||||
## Contents
|
||||
|
||||
- [NetBird Agent Guidelines](#netbird-agent-guidelines)
|
||||
- [Contents](#contents)
|
||||
- [STOP and ask the user before](#stop-and-ask-the-user-before)
|
||||
- [Quick reference](#quick-reference)
|
||||
- [Structure](#structure)
|
||||
- [Where to look](#where-to-look)
|
||||
- [Repo-wide principles](#repo-wide-principles)
|
||||
- [Error handling](#error-handling)
|
||||
- [Comments](#comments)
|
||||
- [Testing](#testing)
|
||||
- [Pitfalls](#pitfalls)
|
||||
- [Commits, PRs, releases](#commits-prs-releases)
|
||||
- [After you push: CI and review bots](#after-you-push-ci-and-review-bots)
|
||||
- [Discussion and support](#discussion-and-support)
|
||||
|
||||
## STOP and ask the user before
|
||||
|
||||
- **Opening a pull request for anything beyond a trivial fix, without an agreed
|
||||
ticket.** Ask the user directly: *"Is there a discussion or issue for this
|
||||
change?"* NetBird is discussion-first — community reports start in
|
||||
[Discussions](https://github.com/netbirdio/netbird/discussions), DevRel
|
||||
validates them, and only validated discussions become issues. A PR that
|
||||
changes behavior with no linked issue may be closed on arrival. If there is no
|
||||
ticket, offer to draft the discussion post **instead of** the PR, and wait for
|
||||
the user's call. Only typos, broken links, documentation corrections, and
|
||||
one-line fixes that already have an issue can skip this.
|
||||
- **Designing in any high-risk area** (see
|
||||
[CONTRIBUTING.md](CONTRIBUTING.md#high-risk-areas)): public API and OpenAPI
|
||||
schema, gRPC protos, behavior existing deployments would notice after an
|
||||
upgrade, peer connectivity (ICE, NAT traversal, relay selection, WireGuard® or
|
||||
Rosenpass key handling), client system integration (routing, firewall, DNS,
|
||||
interface), authentication and authorization, CLI or service flags, config
|
||||
file format, daemon IPC, store schema and migrations, or a new feature. The
|
||||
design gets agreed in the ticket before code is written.
|
||||
- **Writing a store migration or changing a persisted model.** Migrations are
|
||||
one-way in the field and both the GORM and pgx paths may need the change.
|
||||
- **Hand-editing generated code.** `*.pb.go`, `*.gen.go`, and mocks are outputs.
|
||||
Edit the source (`.proto`, `openapi.yml`) and rerun the matching
|
||||
`generate.sh`.
|
||||
- **Adding, removing, or bumping a dependency**, and never vendor a fork.
|
||||
- **Weakening a security control** — authentication, authorization, certificate
|
||||
verification, privilege dropping, or peer identity checks — even when it is
|
||||
the fastest way to make a test pass.
|
||||
- **Force-pushing to `main`**, force-pushing any branch that is already under
|
||||
review, amending pushed commits, or bypassing hooks with `--no-verify`.
|
||||
|
||||
## Quick reference
|
||||
|
||||
```bash
|
||||
# Build
|
||||
go build ./...
|
||||
cd client && CGO_ENABLED=0 go build . # agent
|
||||
cd management && go build . # management service
|
||||
cd signal && go build . # signal service
|
||||
|
||||
# Verify (run before every push)
|
||||
go fmt ./...
|
||||
make lint # golangci-lint on files changed vs origin/main (also the pre-push hook)
|
||||
make lint-all # full-repository lint, matches CI
|
||||
make test-unit # host-safe unit tests, -tags devcert, no sudo
|
||||
make test-privileged # privileged-tagged suite in a Docker container with NET_ADMIN
|
||||
make setup-hooks # wire make lint into .githooks/pre-push
|
||||
|
||||
# Narrow runs
|
||||
go test ./client/internal/dns/...
|
||||
go test -race -run TestPeerConn ./client/internal/peer/...
|
||||
PRIV_RUN=TestNftablesManager PRIV_PKGS=./client/firewall/nftables/... make test-privileged
|
||||
|
||||
# Code generation (never hand-edit the output)
|
||||
./shared/management/http/api/generate.sh # REST types from openapi.yml
|
||||
./shared/management/proto/generate.sh
|
||||
./shared/signal/proto/generate.sh
|
||||
./client/proto/generate.sh
|
||||
./flow/proto/generate.sh
|
||||
|
||||
# Run locally (lab only, never on a machine you rely on)
|
||||
sudo ./client/netbird up --log-level debug --log-file console
|
||||
sudo ./client/netbird down # teardown: restores routing, firewall, DNS
|
||||
./signal/signal run --log-level debug --log-file console
|
||||
./management/management management --log-level debug --log-file console --config ./management.json
|
||||
```
|
||||
|
||||
`netbird up` needs root and rewrites the host's routing table, firewall rules,
|
||||
DNS configuration, and WireGuard® interface. Run it only in a disposable test
|
||||
environment (a VM, container, or throwaway host) that you can rebuild, never on
|
||||
a workstation or server whose connectivity matters. Run `sudo netbird down`
|
||||
before you stop working, before rebuilding the binary, and on every failure
|
||||
path, so the host's networking state is restored instead of left half-applied.
|
||||
See [Pitfalls](#pitfalls) for why cleanup on every exit path matters.
|
||||
|
||||
## Structure
|
||||
|
||||
```text
|
||||
netbird/
|
||||
├── client/ NetBird agent
|
||||
│ ├── cmd/ agent CLI
|
||||
│ ├── internal/ agent business logic (engine, peer, dns, routemanager, ...)
|
||||
│ ├── server/ daemon for background execution
|
||||
│ ├── proto/ daemon gRPC protos
|
||||
│ ├── iface/ WireGuard® interface management
|
||||
│ ├── firewall/ nftables, iptables, pf, WFP, userspace backends
|
||||
│ ├── ssh/ built-in SSH server and client
|
||||
│ ├── ui/ desktop UI (Wails v3 + React)
|
||||
│ ├── android/, ios/ mobile bindings
|
||||
│ ├── wasm/ WebAssembly build
|
||||
│ └── mdm/, system/ MDM policy, host information
|
||||
├── management/ control plane
|
||||
│ └── server/ account, peer, groups, networks, posture, permissions,
|
||||
│ settings, store, http (REST), idp, integrations, migration
|
||||
├── signal/ handshake broker (peer/, server/)
|
||||
├── relay/ relay service (protocol/, server/, healthcheck/)
|
||||
├── proxy/ identity-aware proxy (llm/, acme/, accesslog/, middleware/, tcp/, udp/)
|
||||
├── agent-network/ Agent Network overview
|
||||
├── shared/ imported by both agent and services
|
||||
│ ├── management/ proto/, client/, http/api (OpenAPI + generated types)
|
||||
│ ├── signal/ proto/, client/
|
||||
│ └── relay/, auth/, sshauth/, metrics/
|
||||
├── e2e/ end-to-end suites and harness
|
||||
├── encryption/, dns/, route/, stun/, sharedsock/, util/, flow/
|
||||
├── infrastructure_files/ docker compose and getting-started templates
|
||||
└── release_files/ files packaged into releases
|
||||
```
|
||||
|
||||
## Where to look
|
||||
|
||||
| Task | Location |
|
||||
| --------------------------- | ------------------------------------------------------------ |
|
||||
| REST API / OpenAPI | `shared/management/http/api/` + `management/server/http/` |
|
||||
| Management gRPC protocol | `shared/management/proto/` |
|
||||
| Signal protocol | `shared/signal/proto/` |
|
||||
| Daemon IPC protocol | `client/proto/` |
|
||||
| Peer connection and NAT | `client/internal/peer/` |
|
||||
| Network map handling | `client/internal/engine.go`, `shared/management/networkmap/` |
|
||||
| Routing | `client/internal/routemanager/`, `route/` |
|
||||
| Firewall backends | `client/firewall/` |
|
||||
| DNS | `client/internal/dns/`, `dns/` |
|
||||
| WireGuard® interface | `client/iface/` |
|
||||
| Persistence and migrations | `management/server/store/`, `management/server/migration/` |
|
||||
| IdP integrations | `management/server/idp/` |
|
||||
| Permissions model | `management/server/permissions/` |
|
||||
| LLM routing / Agent Network | `proxy/internal/llm/`, `agent-network/` |
|
||||
| End-to-end tests | `e2e/` |
|
||||
|
||||
## Repo-wide principles
|
||||
|
||||
1. **Run `go fmt` on every modified Go file.** Formatting is not optional.
|
||||
2. **Zero unaddressed diagnostics.** Fix IDE and linter warnings on code you
|
||||
touch, and delete imports, helpers, and parameters your refactor orphaned.
|
||||
Exception: unused parameters in shared code may be consumed by builds outside
|
||||
this repository — do not remove them, ask instead.
|
||||
3. **Function comments are mandatory for exported functions**, written as full
|
||||
sentences with a period, starting with the identifier name.
|
||||
4. **Prefer private functions and constants.** Export only what a caller outside
|
||||
the package genuinely needs.
|
||||
5. **Early returns and guard clauses.** Handle errors and edge cases first
|
||||
instead of nesting `if`/`else` chains.
|
||||
6. **Split complex functions.** If a function trips a complexity warning, break
|
||||
it into named helpers rather than silencing the warning.
|
||||
7. **Avoid LLM-slop tells:** em dashes, hedging narration, restating the diff in
|
||||
prose, trailing summaries. Defaults, not absolute bans. Applies to code,
|
||||
comments, commit messages, and PR descriptions alike.
|
||||
8. **Concurrency: do a two-pass race analysis after every change** that adds
|
||||
shared state. Guard maps and slices with a mutex, keep critical sections
|
||||
short, and run `go test -race` on the touched packages.
|
||||
9. **Cross-platform builds must keep working.** The agent targets Linux, macOS,
|
||||
Windows, FreeBSD, Android, and iOS. When you add a platform-specific file,
|
||||
add the counterpart or a build-tagged fallback for the others.
|
||||
10. **Never hand-edit generated files.** Change the source and regenerate.
|
||||
11. **Never log secrets** — private keys, setup keys, tokens, PAT values — and
|
||||
keep peer IPs and hostnames out of logs above debug level.
|
||||
|
||||
## Error handling
|
||||
|
||||
Use single-assignment form when the error is only needed inside the `if`:
|
||||
|
||||
```go
|
||||
// Good
|
||||
if err := someCall(); err != nil {
|
||||
return fmt.Errorf("context: %w", err)
|
||||
}
|
||||
|
||||
// Bad - unnecessary split
|
||||
err := someCall()
|
||||
if err != nil {
|
||||
return fmt.Errorf("context: %w", err)
|
||||
}
|
||||
```
|
||||
|
||||
Use multiple assignment when the value is needed after the block:
|
||||
|
||||
```go
|
||||
result, err := someCall()
|
||||
if err != nil {
|
||||
return fmt.Errorf("context: %w", err)
|
||||
}
|
||||
```
|
||||
|
||||
Add short, meaningful context, and **do not** start `fmt.Errorf` messages with
|
||||
obvious words like "failed to" or "error":
|
||||
|
||||
```go
|
||||
// Good
|
||||
return fmt.Errorf("parse remote address: %w", err)
|
||||
return fmt.Errorf("listen on %s: %w", addr, err)
|
||||
|
||||
// Bad
|
||||
return fmt.Errorf("failed to parse remote address: %w", err)
|
||||
return fmt.Errorf("error listening on %s: %w", addr, err)
|
||||
|
||||
// "failed" is fine in log messages
|
||||
log.Debugf("failed to parse remote address: %v", err)
|
||||
```
|
||||
|
||||
Skip the wrapping when a function only extracts or delegates and the wrap would
|
||||
add nothing:
|
||||
|
||||
```go
|
||||
func parseAddr(addr string) (string, int, error) {
|
||||
host, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
// ...
|
||||
}
|
||||
```
|
||||
|
||||
Log the errors you choose not to act on:
|
||||
|
||||
- `log.Debugf()` for errors that do not affect program flow but help debugging.
|
||||
- `log.Tracef()` for very verbose errors that would otherwise spam logs.
|
||||
- **Never ignore** errors from writes, network sends, or critical cleanup.
|
||||
- Close errors may be ignored for read-only operations; log them at debug for
|
||||
writes.
|
||||
|
||||
## Comments
|
||||
|
||||
Comment the **why**, never the **what**. Default to no comment, and add one only
|
||||
when a hidden constraint or workaround would surprise a future reader. Never
|
||||
reference the current task, PR, or your own changes in a comment.
|
||||
|
||||
```go
|
||||
// Bad - trailing comments explaining the obvious
|
||||
defer localConn.Close() // Close the connection
|
||||
if err != nil { // Check if error occurred
|
||||
|
||||
// Good
|
||||
defer localConn.Close()
|
||||
|
||||
// Good - explains a non-obvious constraint
|
||||
// Use incremental checksum update per RFC 1624 for performance.
|
||||
checksum = updateChecksum(checksum, oldPort, newPort)
|
||||
```
|
||||
|
||||
### Length budget
|
||||
|
||||
- **90 characters per line.** Wrap the comment, do not run past it.
|
||||
- **250 characters per comment**, roughly three wrapped lines. Doc comments on
|
||||
exported identifiers may exceed it when the API genuinely needs the
|
||||
explanation; inline comments inside a function body may not.
|
||||
|
||||
The budget is a smell detector, not a rule to game. Do not compress a needed
|
||||
explanation into cryptic shorthand to fit — if a block of code needs more than
|
||||
250 characters of prose, the code is doing too much. Fix the code:
|
||||
|
||||
- **Extract a named function.** A well-named function replaces the comment: the
|
||||
name says *what*, the body shows *how*, and the comment you no longer write
|
||||
was the *what* anyway. Clean Code calls this "explain yourself in code".
|
||||
- **Extract a named constant or predicate.** `if isExpiredSetupKey(key)` needs
|
||||
no comment; `if key.ExpiresAt.Before(now) && !key.Revoked && key.UsageLimit > 0`
|
||||
does.
|
||||
- **Keep the surviving comment for the why** — the RFC, the kernel quirk, the
|
||||
ordering constraint. That part is usually one or two lines.
|
||||
|
||||
### Long switch and if/else chains
|
||||
|
||||
A `switch` whose cases carry multi-line explanations is the usual place this
|
||||
budget is breached, and the comment is a symptom. In order of preference:
|
||||
|
||||
1. **Extract each case body into a named function.** The case becomes one line,
|
||||
the name carries the meaning, and the switch reads as a table of contents.
|
||||
2. **Replace the switch with a lookup table** — `map[Kind]handlerFunc` — when the
|
||||
branches are uniform. Adding a case stops meaning editing a growing function.
|
||||
3. **Replace conditional with polymorphism** when branches vary by type and the
|
||||
same switch shape starts appearing in more than one place. Clean Code's rule
|
||||
of thumb: tolerate a switch statement if it appears **once**, is buried in a
|
||||
factory that returns an interface, and no other switch dispatches on the same
|
||||
type. A second switch over the same enum is the signal to introduce the
|
||||
interface.
|
||||
|
||||
Do not restructure a switch purely to satisfy the budget when the cases are one
|
||||
line each and self-evident — a flat, boring `switch` over an enum is fine and
|
||||
needs no comments at all.
|
||||
|
||||
Explanatory comments in tests are welcome — they document the scenario being set
|
||||
up, and the 250-character budget does not apply to them.
|
||||
|
||||
## Testing
|
||||
|
||||
- **Unit tests** live beside the code as `_test.go`. `make test-unit` runs the
|
||||
host-safe set with `-tags devcert` and no sudo.
|
||||
- **Privileged tests** carry the `privileged` build tag and mutate host
|
||||
networking. They run through `make test-privileged`, inside a Docker container
|
||||
with `NET_ADMIN`. Never bypass that harness by running them directly on the
|
||||
host.
|
||||
- **End-to-end suites** live in `e2e/` with a shared harness.
|
||||
- **Test real behavior, not API existence.** Assert on the observable end state
|
||||
a consumer would see — bytes that arrived, the packet after translation, the
|
||||
row after the write — not merely that a method exists or returns an error.
|
||||
- **Avoid mocks for code we own.** Exercise the real store, manager, or
|
||||
controller and assert what the caller actually receives.
|
||||
- **`require` for setup and preconditions, `assert` for the conditions under
|
||||
test.** Use `require` whenever a later line would panic or be meaningless
|
||||
otherwise.
|
||||
- **Message guidance:** optional for `NoError`/`Error`; always give context for
|
||||
comparison, boolean, and collection assertions.
|
||||
|
||||
```go
|
||||
server, err := StartTestServer()
|
||||
require.NoError(t, err, "Test server setup must succeed")
|
||||
defer server.Close()
|
||||
|
||||
result, err := client.DoOperation()
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, expectedResult, result, "Result should match expected")
|
||||
```
|
||||
|
||||
## Pitfalls
|
||||
|
||||
- **The agent runs as root.** Anything touching routing, firewall, DNS, or the
|
||||
interface can take a user's machine off the network. Prefer a reversible
|
||||
change and make sure cleanup runs on every exit path.
|
||||
- **Management has two account loaders** (GORM and pgx). Adding a relation to an
|
||||
account often means updating both, or it silently comes back empty in
|
||||
production.
|
||||
- **`go test ./...` without `-tags devcert` skips tests** that need the
|
||||
development certificate. Use `make test-unit`.
|
||||
- **`make lint` only checks the diff against `origin/main`.** CI runs
|
||||
`make lint-all`; run it too before pushing a large change.
|
||||
- **Protos are consumed by released clients.** An old agent must keep working
|
||||
against a new Management, so fields are added, never renumbered or removed.
|
||||
- **Windows requires the wintun driver**, and the daemon serves a named pipe
|
||||
(`npipe://netbird`) rather than loopback TCP. Loopback TCP carries no caller
|
||||
identity, so privileged operations are refused over it.
|
||||
|
||||
## Commits, PRs, releases
|
||||
|
||||
- **PR titles must start with a bracketed tag.** Before you propose a title,
|
||||
**read [`.github/workflows/pr-title-check.yml`](.github/workflows/pr-title-check.yml)
|
||||
and take the allowed tags from the `allowedTags` array in that file.** It is
|
||||
the only source of truth, it changes as components are added, and the check
|
||||
runs on every title edit — a tag that is not in that array is a red build. Do
|
||||
not rely on a list memorized from anywhere else, including this file.
|
||||
|
||||
```text
|
||||
[client] Authorize daemon IPC callers by their local identity
|
||||
[management,client] Add MDM policy support
|
||||
```
|
||||
|
||||
Multiple tags are comma-separated inside one pair of brackets. Match the tag
|
||||
to the component you actually changed, not to the one you read the most.
|
||||
|
||||
- **Use the repository's PR template.** Fill in
|
||||
[`.github/pull_request_template.md`](.github/pull_request_template.md) rather
|
||||
than replacing it with your own summary: describe the change, link the issue,
|
||||
tick the checklist honestly (including "ran locally" and "single purpose"),
|
||||
and complete the documentation section. Do not tick a box you have not
|
||||
verified, and do not delete rows that do not apply.
|
||||
|
||||
- **Keep the PR description short.** Under 1000 words on top of the template's
|
||||
own text, and usually far less — a few paragraphs. Reviewers read the diff;
|
||||
the description exists to explain what the diff cannot say for itself. This is
|
||||
well below what an agent will produce by default, so cut before you post.
|
||||
|
||||
- **Body: why before what.** Lead with the problem and the reason for this
|
||||
approach, then the shape of the change. No bullet list of files changed, no
|
||||
per-function walkthrough, no restating the diff in prose, no trailing summary
|
||||
section, no self-congratulatory closing line.
|
||||
|
||||
- **No `Co-Authored-By` or tool-attribution trailers in the PR description**,
|
||||
and none in commits either. Contributors own their contributions. Whatever
|
||||
tooling produced the diff, the person opening the PR is its author: they have
|
||||
read every line, they can explain why it works, they can answer review
|
||||
questions without going back to a model, and they are accountable for the
|
||||
consequences of merging it. Do not add a trailer, footer, or description line
|
||||
that spreads that ownership onto a tool.
|
||||
|
||||
- **Commit subjects follow the same `[scope] Subject` convention.** Keep the
|
||||
subject short, and use the body for why before what. No bullet lists of files
|
||||
changed.
|
||||
|
||||
- **Push review fixes as separate commits.** The PR is squashed on merge, so
|
||||
there is no reason to rewrite history mid-review; many small commits make the
|
||||
re-review readable.
|
||||
|
||||
- **Do not force-push a branch that is under review.** A force-push detaches
|
||||
existing review comments from the lines they were written against, destroys
|
||||
the "changes since your last review" diff a reviewer relies on, and discards
|
||||
the CI history that showed which commit broke what. Add commits instead —
|
||||
including for fixups and reverts. Force-push only when there is no
|
||||
alternative: a rebase to clear a genuine conflict, or removing a secret or a
|
||||
large binary that was committed by mistake. When you must, ask the user first,
|
||||
then say so in a PR comment so reviewers know their anchors moved. Never
|
||||
force-push `main`, and never force-push a branch you do not own.
|
||||
|
||||
- **One PR, one purpose.** Split refactors out of fixes and fixes out of
|
||||
features.
|
||||
|
||||
- **Keep the PR small.** Size is the single strongest predictor of how long a PR
|
||||
waits. Aim for **under ~400 changed lines across under ~20 files**; past
|
||||
roughly **1000 lines or 50 files** a community PR is likely to be sent back to
|
||||
be split, or left unreviewed until it is. Large PRs from outside the core team
|
||||
may be blocked outright when the size was never agreed in the ticket —
|
||||
reviewing a sprawling change against a privileged networking daemon is a
|
||||
security risk in itself, not just a time cost.
|
||||
|
||||
Judge the size by hand-written code: exclude generated output, `go.sum`,
|
||||
vendored files, and test fixtures from the estimate, but do not use their
|
||||
presence to argue a 3000-line PR is small.
|
||||
|
||||
When a change genuinely cannot be small — a protocol migration, a
|
||||
cross-component rename — agree the split in the ticket **before** writing
|
||||
code, and land it as a sequence of PRs that each build, test, and make sense
|
||||
on their own. Propose that split to the user rather than opening one large PR
|
||||
and hoping.
|
||||
|
||||
- **User-facing changes need a docs PR** in
|
||||
[netbirdio/docs](https://github.com/netbirdio/docs), linked from the PR
|
||||
description.
|
||||
|
||||
## After you push: CI and review bots
|
||||
|
||||
Opening the PR is not the end of the task. Watch the run, read what the bots
|
||||
say, and drive the PR to green before you report the work as done.
|
||||
|
||||
```bash
|
||||
gh pr checks <pr> --watch # all checks, live
|
||||
gh run view <run-id> --log-failed # only the failing steps
|
||||
gh pr view <pr> --comments # bot and human review comments
|
||||
```
|
||||
|
||||
**Never report a change as finished while checks are pending or red**, and never
|
||||
describe a red PR as passing. If you ran out of turn before CI finished, say
|
||||
which checks were still running.
|
||||
|
||||
### The checks
|
||||
|
||||
- **Go tests** — `golang-test-{linux,darwin,windows,freebsd}.yml`, sharded per
|
||||
component. A failure in a component you did not touch is usually a real
|
||||
interaction, not noise; read the log before assuming flake.
|
||||
- **golangci-lint** — `golangci-lint.yml` runs the full repository, while
|
||||
`make lint` only checks your diff. A clean local lint does not guarantee green
|
||||
CI on a large change.
|
||||
- **PR Title Check** — `pr-title-check.yml`, see above.
|
||||
- **Codecov** — uploaded from the Linux test workflow with per-component flags
|
||||
(`unit,client`, `unit,management`, `unit,relay`, `unit,proxy`, `unit,signal`,
|
||||
`integration,management`). Coverage on new code should not go backwards. Add
|
||||
tests for the paths you introduced; do not adjust thresholds or exclude files
|
||||
to clear the report.
|
||||
- **CodeRabbit** — configured in [`.coderabbit.yaml`](.coderabbit.yaml): `chill`
|
||||
profile, auto-review on every non-draft PR, TypeScript/JavaScript/SVG paths
|
||||
filtered out. Chat auto-reply is on, so `@coderabbitai` in a comment reaches
|
||||
it.
|
||||
- **SonarCloud** — project `netbirdio_netbird`, quality gate on new code (bugs,
|
||||
vulnerabilities, code smells, duplication, coverage).
|
||||
- **Snyk** — dependency and code scanning.
|
||||
|
||||
Sonar and Snyk report as GitHub App checks rather than workflows in this
|
||||
repository, so their detail lives on the PR check, not in the Actions logs.
|
||||
|
||||
### Handling bot findings
|
||||
|
||||
- **Read every comment and act on it.** Either fix it, or reply with the reason
|
||||
it does not apply. Do not bulk-resolve threads to clear the count, and do not
|
||||
silently ignore a finding because the check is advisory.
|
||||
- **Bots are frequently wrong here.** NetBird has privileged, platform-specific,
|
||||
and concurrency-heavy code that static analysis reads poorly. A confident
|
||||
CodeRabbit or Sonar comment can still be nonsense. Verify the claim against
|
||||
the code before you change anything — never edit correct code just to silence
|
||||
a bot.
|
||||
- **Security findings get the opposite default.** For a Snyk or Sonar
|
||||
vulnerability, or a CodeRabbit comment about authentication, authorization,
|
||||
certificate verification, or key handling, assume it is real until you have
|
||||
disproved it. Surface it to the user rather than dismissing it yourself.
|
||||
- **A new vulnerable dependency is a stop.** Bumping or replacing dependencies
|
||||
needs the user's decision, as above.
|
||||
- **Never change a workflow, threshold, lint exclusion, or bot config to make a
|
||||
check pass.** If a check is genuinely wrong, say so and let the user decide.
|
||||
- **Do not paper over flakes with blind re-runs.** Identify the failure first. If
|
||||
it is a known flake, name it; if you cannot tell, report it as unresolved
|
||||
rather than re-running until it goes green.
|
||||
|
||||
## Discussion and support
|
||||
|
||||
- Discussions: <https://github.com/netbirdio/netbird/discussions>
|
||||
- Slack: <https://docs.netbird.io/slack-url>
|
||||
- Docs: <https://docs.netbird.io>
|
||||
- Security: <https://github.com/netbirdio/netbird/security/policy> — never in public
|
||||
- Contribution process: [CONTRIBUTING.md](CONTRIBUTING.md)
|
||||
1
CLAUDE.md
Normal file
1
CLAUDE.md
Normal file
@@ -0,0 +1 @@
|
||||
See [AGENTS.md](AGENTS.md) for the agent guidelines in this repository.
|
||||
@@ -66,11 +66,41 @@ Typical bug fixes, internal refactors, documentation updates, and tests do not
|
||||
need a design discussion, but should still be tied to an issue so the work is
|
||||
visible and nobody duplicates it.
|
||||
|
||||
### Using AI coding agents
|
||||
|
||||
We have no policy for or against using an AI agent to write NetBird code. That
|
||||
choice is yours, and we are not going to interrogate anyone about their tools.
|
||||
|
||||
What we do have is a lot of incoming contributions that were plainly drafted with
|
||||
one, and enough experience reviewing them to see the same avoidable problems
|
||||
again and again: no ticket behind the change, a diff far too large to review, a
|
||||
description longer than the code it describes, an approach that was never going
|
||||
to be accepted, and an author who cannot answer questions about their own PR.
|
||||
None of that is caused by the tooling — it is what happens when a tool is pointed
|
||||
at a repository whose expectations it has never been told.
|
||||
|
||||
So rather than a rule, there is a guide. [AGENTS.md](AGENTS.md) restates the
|
||||
expectations from this document in the form agents read automatically
|
||||
(`CLAUDE.md` points to it), so pointing your tool at the repository is usually
|
||||
enough. Among other things it tells the agent to ask you for the
|
||||
discussion or issue before drafting a PR, to keep the change small and
|
||||
single-purpose, to run the tests locally, to use this repository's PR template
|
||||
and title tags, and to write a description a reviewer can get through.
|
||||
|
||||
The guardrails are the point, and they are the same ones we apply to everyone: an
|
||||
agreed ticket, a change you have actually run, a diff small enough to review with
|
||||
care, and an author who can explain it. Whatever wrote the diff, you are its
|
||||
author — you own every line you submit and the consequences of opening a PR with it.
|
||||
|
||||
We may assess whether a contribution is maintainable and whether its merged code
|
||||
aligns with our security standards and design expectations.
|
||||
|
||||
## Contents
|
||||
|
||||
- [Contributing to NetBird](#contributing-to-netbird)
|
||||
- [Ticket first, PR second](#ticket-first-pr-second)
|
||||
- [High-risk areas](#high-risk-areas)
|
||||
- [Using AI coding agents](#using-ai-coding-agents)
|
||||
- [Contents](#contents)
|
||||
- [Code of conduct](#code-of-conduct)
|
||||
- [Directory structure](#directory-structure)
|
||||
|
||||
@@ -301,7 +301,7 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (strin
|
||||
uploadCtx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path, false)
|
||||
key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("upload debug bundle: %w", err)
|
||||
}
|
||||
|
||||
@@ -29,9 +29,8 @@ const errCloseConnection = "Failed to close connection: %v"
|
||||
var (
|
||||
logFileCount uint32
|
||||
systemInfoFlag bool
|
||||
uploadBundleFlag bool
|
||||
uploadBundleURLFlag string
|
||||
uploadBundleInsecureFlag bool
|
||||
uploadBundleFlag bool
|
||||
uploadBundleURLFlag string
|
||||
)
|
||||
|
||||
var debugCmd = &cobra.Command{
|
||||
@@ -175,11 +174,10 @@ func debugBundle(cmd *cobra.Command, _ []string) error {
|
||||
}
|
||||
if uploadBundleFlag {
|
||||
request.UploadURL = uploadBundleURLFlag
|
||||
request.UploadInsecure = uploadBundleInsecureFlag
|
||||
}
|
||||
resp, err := client.DebugBundle(cmd.Context(), request)
|
||||
if err != nil {
|
||||
return daemonCallError("bundle debug", err)
|
||||
return fmt.Errorf("failed to bundle debug: %v", status.Convert(err).Message())
|
||||
}
|
||||
cmd.Printf("Local file:\n%s\n", resp.GetPath())
|
||||
|
||||
@@ -375,11 +373,10 @@ func runForDuration(cmd *cobra.Command, args []string) error {
|
||||
}
|
||||
if uploadBundleFlag {
|
||||
request.UploadURL = uploadBundleURLFlag
|
||||
request.UploadInsecure = uploadBundleInsecureFlag
|
||||
}
|
||||
resp, err := client.DebugBundle(cmd.Context(), request)
|
||||
if err != nil {
|
||||
return daemonCallError("bundle debug", err)
|
||||
return fmt.Errorf("failed to bundle debug: %v", status.Convert(err).Message())
|
||||
}
|
||||
|
||||
if needsRestoreUp {
|
||||
@@ -527,12 +524,10 @@ func init() {
|
||||
debugBundleCmd.Flags().BoolVarP(&systemInfoFlag, "system-info", "S", true, "Adds system information to the debug bundle")
|
||||
debugBundleCmd.Flags().BoolVarP(&uploadBundleFlag, "upload-bundle", "U", false, "Uploads the debug bundle to a server")
|
||||
debugBundleCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle")
|
||||
debugBundleCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root")
|
||||
|
||||
forCmd.Flags().Uint32VarP(&logFileCount, "log-file-count", "C", 1, "Number of rotated log files to include in debug bundle")
|
||||
forCmd.Flags().BoolVarP(&systemInfoFlag, "system-info", "S", true, "Adds system information to the debug bundle")
|
||||
forCmd.Flags().BoolVarP(&uploadBundleFlag, "upload-bundle", "U", false, "Uploads the debug bundle to a server")
|
||||
forCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle")
|
||||
forCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root")
|
||||
forCmd.Flags().Bool("capture", false, "Capture packets during the debug duration and include in bundle")
|
||||
}
|
||||
|
||||
@@ -6,11 +6,6 @@ import (
|
||||
"runtime"
|
||||
)
|
||||
|
||||
// UILogFile is the file name the desktop UI writes its log to. It is defined
|
||||
// here so the UI (writer), the daemon's RegisterUILog validation, and the debug
|
||||
// bundle collector all share one definition.
|
||||
const UILogFile = "gui-client.log"
|
||||
|
||||
var StateDir string
|
||||
|
||||
func init() {
|
||||
|
||||
@@ -229,6 +229,7 @@ scutil_dns.txt (macOS only):
|
||||
|
||||
const (
|
||||
clientLogFile = "client.log"
|
||||
uiLogFile = "gui-client.log"
|
||||
errorLogFile = "netbird.err"
|
||||
stdoutLogFile = "netbird.out"
|
||||
|
||||
@@ -247,20 +248,6 @@ type MetricsExporter interface {
|
||||
Export(w io.Writer) error
|
||||
}
|
||||
|
||||
// LogOpener opens a log file for inclusion in the bundle. It exists so that log
|
||||
// files whose path was supplied by an IPC caller can be opened under a check
|
||||
// the daemon defines, instead of being opened with the daemon's privileges
|
||||
// unconditionally.
|
||||
type LogOpener func(path string) (*os.File, error)
|
||||
|
||||
func openLogFile(path string) (*os.File, error) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open %s: %w", path, err)
|
||||
}
|
||||
return f, nil
|
||||
}
|
||||
|
||||
type BundleGenerator struct {
|
||||
anonymizer *anonymize.Anonymizer
|
||||
|
||||
@@ -270,7 +257,6 @@ type BundleGenerator struct {
|
||||
syncResponse *mgmProto.SyncResponse
|
||||
logPath string
|
||||
uiLogPath string
|
||||
uiLogOpener LogOpener
|
||||
tempDir string
|
||||
statePath string
|
||||
cpuProfile []byte
|
||||
@@ -299,20 +285,14 @@ type GeneratorDependencies struct {
|
||||
SyncResponse *mgmProto.SyncResponse
|
||||
LogPath string
|
||||
UILogPath string // Absolute path to the desktop UI's gui-client.log, reported via RegisterUILog. Empty if no UI registered one.
|
||||
// UILogOpener opens the UI log and its rotated siblings. The path comes from
|
||||
// a local IPC caller, so the daemon must not open it with plain os.Open: the
|
||||
// opener is where the caller's right to that file is enforced. Defaults to
|
||||
// os.Open, which is only correct where the path is not caller-supplied
|
||||
// (mobile).
|
||||
UILogOpener LogOpener
|
||||
TempDir string // Directory for temporary bundle zip files. If empty, os.TempDir() is used.
|
||||
StatePath string // Path to the state file. If empty, the ServiceManager default path is used.
|
||||
CPUProfile []byte
|
||||
CapturePath string
|
||||
RefreshStatus func()
|
||||
ClientMetrics MetricsExporter
|
||||
DaemonVersion string
|
||||
CliVersion string
|
||||
TempDir string // Directory for temporary bundle zip files. If empty, os.TempDir() is used.
|
||||
StatePath string // Path to the state file. If empty, the ServiceManager default path is used.
|
||||
CPUProfile []byte
|
||||
CapturePath string
|
||||
RefreshStatus func()
|
||||
ClientMetrics MetricsExporter
|
||||
DaemonVersion string
|
||||
CliVersion string
|
||||
}
|
||||
|
||||
func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGenerator {
|
||||
@@ -322,11 +302,6 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen
|
||||
logFileCount = 1
|
||||
}
|
||||
|
||||
uiLogOpener := deps.UILogOpener
|
||||
if uiLogOpener == nil {
|
||||
uiLogOpener = openLogFile
|
||||
}
|
||||
|
||||
return &BundleGenerator{
|
||||
anonymizer: anonymize.NewAnonymizer(anonymize.DefaultAddresses()),
|
||||
|
||||
@@ -335,7 +310,6 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen
|
||||
syncResponse: deps.SyncResponse,
|
||||
logPath: deps.LogPath,
|
||||
uiLogPath: deps.UILogPath,
|
||||
uiLogOpener: uiLogOpener,
|
||||
tempDir: deps.TempDir,
|
||||
statePath: deps.StatePath,
|
||||
cpuProfile: deps.CPUProfile,
|
||||
@@ -1022,11 +996,11 @@ func (g *BundleGenerator) addLogfile() error {
|
||||
|
||||
logDir := filepath.Dir(g.logPath)
|
||||
|
||||
if err := g.addSingleLogfile(openLogFile, g.logPath, clientLogFile); err != nil {
|
||||
if err := g.addSingleLogfile(g.logPath, clientLogFile); err != nil {
|
||||
return fmt.Errorf("add client log file to zip: %w", err)
|
||||
}
|
||||
|
||||
g.addRotatedLogFiles(openLogFile, logDir, clientLogPrefix)
|
||||
g.addRotatedLogFiles(logDir, clientLogPrefix)
|
||||
|
||||
stdErrLogPath := filepath.Join(logDir, errorLogFile)
|
||||
stdoutLogPath := filepath.Join(logDir, stdoutLogFile)
|
||||
@@ -1035,11 +1009,11 @@ func (g *BundleGenerator) addLogfile() error {
|
||||
stdoutLogPath = darwinStdoutLogPath
|
||||
}
|
||||
|
||||
if err := g.addSingleLogfile(openLogFile, stdErrLogPath, errorLogFile); err != nil {
|
||||
if err := g.addSingleLogfile(stdErrLogPath, errorLogFile); err != nil {
|
||||
log.Warnf("Failed to add %s to zip: %v", errorLogFile, err)
|
||||
}
|
||||
|
||||
if err := g.addSingleLogfile(openLogFile, stdoutLogPath, stdoutLogFile); err != nil {
|
||||
if err := g.addSingleLogfile(stdoutLogPath, stdoutLogFile); err != nil {
|
||||
log.Warnf("Failed to add %s to zip: %v", stdoutLogFile, err)
|
||||
}
|
||||
|
||||
@@ -1056,18 +1030,18 @@ func (g *BundleGenerator) addUILog() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := g.addSingleLogfile(g.uiLogOpener, g.uiLogPath, configs.UILogFile); err != nil {
|
||||
if err := g.addSingleLogfile(g.uiLogPath, uiLogFile); err != nil {
|
||||
return fmt.Errorf("add UI log file to zip: %w", err)
|
||||
}
|
||||
|
||||
g.addRotatedLogFiles(g.uiLogOpener, filepath.Dir(g.uiLogPath), uiLogPrefix)
|
||||
g.addRotatedLogFiles(filepath.Dir(g.uiLogPath), uiLogPrefix)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// addSingleLogfile adds a single log file to the archive
|
||||
func (g *BundleGenerator) addSingleLogfile(open LogOpener, logPath, targetName string) error {
|
||||
logFile, err := open(logPath)
|
||||
func (g *BundleGenerator) addSingleLogfile(logPath, targetName string) error {
|
||||
logFile, err := os.Open(logPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("open log file %s: %w", targetName, err)
|
||||
}
|
||||
@@ -1092,8 +1066,8 @@ func (g *BundleGenerator) addSingleLogfile(open LogOpener, logPath, targetName s
|
||||
}
|
||||
|
||||
// addSingleLogFileGz adds a single gzipped log file to the archive
|
||||
func (g *BundleGenerator) addSingleLogFileGz(open LogOpener, logPath, targetName string) error {
|
||||
f, err := open(logPath)
|
||||
func (g *BundleGenerator) addSingleLogFileGz(logPath, targetName string) error {
|
||||
f, err := os.Open(logPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("open gz log file %s: %w", targetName, err)
|
||||
}
|
||||
@@ -1140,7 +1114,7 @@ func (g *BundleGenerator) addSingleLogFileGz(open LogOpener, logPath, targetName
|
||||
// addRotatedLogFiles adds rotated log files to the bundle based on logFileCount.
|
||||
// prefix is the base log name without extension (e.g. "client", "gui-client");
|
||||
// the glob matches both files rotated by us and by logrotate on linux.
|
||||
func (g *BundleGenerator) addRotatedLogFiles(open LogOpener, logDir, prefix string) {
|
||||
func (g *BundleGenerator) addRotatedLogFiles(logDir, prefix string) {
|
||||
if g.logFileCount == 0 {
|
||||
return
|
||||
}
|
||||
@@ -1180,9 +1154,9 @@ func (g *BundleGenerator) addRotatedLogFiles(open LogOpener, logDir, prefix stri
|
||||
for i := 0; i < maxFiles; i++ {
|
||||
name := filepath.Base(files[i])
|
||||
if strings.HasSuffix(name, ".gz") {
|
||||
err = g.addSingleLogFileGz(open, files[i], name)
|
||||
err = g.addSingleLogFileGz(files[i], name)
|
||||
} else {
|
||||
err = g.addSingleLogfile(open, files[i], name)
|
||||
err = g.addSingleLogfile(files[i], name)
|
||||
}
|
||||
if err != nil {
|
||||
log.Warnf("failed to add rotated log %s: %v", name, err)
|
||||
|
||||
@@ -27,7 +27,7 @@ func (g *BundleGenerator) addPlatformLog() error {
|
||||
}
|
||||
|
||||
swiftLogPath := filepath.Join(filepath.Dir(g.logPath), swiftLogFile)
|
||||
if err := g.addSingleLogfile(openLogFile, swiftLogPath, swiftLogFile); err != nil {
|
||||
if err := g.addSingleLogfile(swiftLogPath, swiftLogFile); err != nil {
|
||||
// The Swift log is best-effort: the app may not have written it yet.
|
||||
log.Warnf("failed to add %s to debug bundle: %v", swiftLogFile, err)
|
||||
}
|
||||
|
||||
@@ -97,7 +97,7 @@ func runAddRotatedLogFilesPrefix(t *testing.T, dir, prefix string, logFileCount
|
||||
archive: zip.NewWriter(&buf),
|
||||
logFileCount: logFileCount,
|
||||
}
|
||||
g.addRotatedLogFiles(openLogFile, dir, prefix)
|
||||
g.addRotatedLogFiles(dir, prefix)
|
||||
require.NoError(t, g.archive.Close())
|
||||
|
||||
zr, err := zip.NewReader(bytes.NewReader(buf.Bytes()), int64(buf.Len()))
|
||||
|
||||
@@ -1,64 +0,0 @@
|
||||
package debug
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"bytes"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/client/configs"
|
||||
)
|
||||
|
||||
// bundleEntries generates a bundle with the given generator and returns the
|
||||
// set of entry names in the resulting archive.
|
||||
func bundleEntries(t *testing.T, g *BundleGenerator) map[string]struct{} {
|
||||
t.Helper()
|
||||
|
||||
path, err := g.Generate()
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = os.Remove(path) })
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
require.NoError(t, err)
|
||||
|
||||
zr, err := zip.NewReader(bytes.NewReader(data), int64(len(data)))
|
||||
require.NoError(t, err)
|
||||
|
||||
names := make(map[string]struct{}, len(zr.File))
|
||||
for _, f := range zr.File {
|
||||
names[f.Name] = struct{}{}
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
func TestBundleIncludesUILogWhenOpenerAllows(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), configs.UILogFile)
|
||||
require.NoError(t, os.WriteFile(path, []byte("gui log"), 0600))
|
||||
|
||||
g := NewBundleGenerator(GeneratorDependencies{
|
||||
UILogPath: path,
|
||||
UILogOpener: openLogFile,
|
||||
}, BundleConfig{})
|
||||
|
||||
require.Contains(t, bundleEntries(t, g), configs.UILogFile)
|
||||
}
|
||||
|
||||
// A UILogOpener that refuses (as the ownership check does for a foreign file)
|
||||
// keeps the UI log out of the bundle without failing bundle generation.
|
||||
func TestBundleExcludesUILogWhenOpenerRefuses(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), configs.UILogFile)
|
||||
require.NoError(t, os.WriteFile(path, []byte("secret"), 0600))
|
||||
|
||||
g := NewBundleGenerator(GeneratorDependencies{
|
||||
UILogPath: path,
|
||||
UILogOpener: func(string) (*os.File, error) {
|
||||
return nil, fmt.Errorf("not owned by the caller")
|
||||
},
|
||||
}, BundleConfig{})
|
||||
|
||||
require.NotContains(t, bundleEntries(t, g), configs.UILogFile)
|
||||
}
|
||||
@@ -3,12 +3,10 @@ package debug
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"crypto/tls"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
neturl "net/url"
|
||||
"os"
|
||||
|
||||
"github.com/netbirdio/netbird/upload-server/types"
|
||||
@@ -16,80 +14,20 @@ import (
|
||||
|
||||
const maxBundleUploadSize = 50 * 1024 * 1024
|
||||
|
||||
// requireHTTPS refuses any URL the daemon would fetch or upload to that is not
|
||||
// https. The daemon runs as root and the bundle carries its logs and state, so a
|
||||
// plaintext hop is a place to intercept the bundle or the presigned redirect.
|
||||
// The server-side gate already enforces this for the desktop path; this also
|
||||
// covers the mobile and job-runner callers that reach this package directly.
|
||||
// Skipped when the caller opted into an insecure upload (self-hosted server).
|
||||
func requireHTTPS(what, rawURL string) error {
|
||||
parsed, err := neturl.Parse(rawURL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse %s: %w", what, err)
|
||||
}
|
||||
if parsed.Scheme != "https" {
|
||||
return fmt.Errorf("%s must use https, got scheme %q", what, parsed.Scheme)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// uploadClient returns the HTTP client for the upload requests. The default
|
||||
// client verifies TLS and refuses a redirect that would downgrade to a non-https
|
||||
// hop, so a bundle can never leave over http after an https start. The insecure
|
||||
// variant accepts http and untrusted certificates, and is only reachable for a
|
||||
// privileged caller that passed --upload-bundle-insecure (see
|
||||
// requirePrivilegeForUploadURL).
|
||||
func uploadClient(insecure bool) *http.Client {
|
||||
if !insecure {
|
||||
return &http.Client{CheckRedirect: rejectInsecureRedirect}
|
||||
}
|
||||
return &http.Client{
|
||||
Transport: &http.Transport{
|
||||
//nolint:gosec // opt-in, privileged, self-hosted upload servers
|
||||
TLSClientConfig: &tls.Config{InsecureSkipVerify: true, MinVersion: tls.VersionTLS12},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// rejectInsecureRedirect refuses a redirect to a non-https target and keeps the
|
||||
// standard library's 10-hop limit that a custom CheckRedirect would otherwise
|
||||
// disable.
|
||||
func rejectInsecureRedirect(req *http.Request, via []*http.Request) error {
|
||||
if req.URL.Scheme != "https" {
|
||||
return fmt.Errorf("refusing redirect to non-https URL %s", req.URL.Redacted())
|
||||
}
|
||||
if len(via) >= 10 {
|
||||
return fmt.Errorf("stopped after 10 redirects")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func UploadDebugBundle(ctx context.Context, url, managementURL, filePath string, insecure bool) (key string, err error) {
|
||||
if !insecure {
|
||||
if err := requireHTTPS("upload service URL", url); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
response, err := getUploadURL(ctx, url, managementURL, insecure)
|
||||
func UploadDebugBundle(ctx context.Context, url, managementURL, filePath string) (key string, err error) {
|
||||
response, err := getUploadURL(ctx, url, managementURL)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
if !insecure {
|
||||
if err := requireHTTPS("upload URL from service", response.URL); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
err = upload(ctx, filePath, response, insecure)
|
||||
err = upload(ctx, filePath, response)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return response.Key, nil
|
||||
}
|
||||
|
||||
func upload(ctx context.Context, filePath string, response *types.GetURLResponse, insecure bool) error {
|
||||
func upload(ctx context.Context, filePath string, response *types.GetURLResponse) error {
|
||||
fileData, err := os.Open(filePath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("open file: %w", err)
|
||||
@@ -114,7 +52,7 @@ func upload(ctx context.Context, filePath string, response *types.GetURLResponse
|
||||
req.ContentLength = stat.Size()
|
||||
req.Header.Set("Content-Type", "application/octet-stream")
|
||||
|
||||
putResp, err := uploadClient(insecure).Do(req)
|
||||
putResp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("upload failed: %v", err)
|
||||
}
|
||||
@@ -127,23 +65,16 @@ func upload(ctx context.Context, filePath string, response *types.GetURLResponse
|
||||
return nil
|
||||
}
|
||||
|
||||
func getUploadURL(ctx context.Context, serviceURL string, managementURL string, insecure bool) (*types.GetURLResponse, error) {
|
||||
parsed, err := neturl.Parse(serviceURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse upload service URL: %w", err)
|
||||
}
|
||||
q := parsed.Query()
|
||||
q.Set("id", getURLHash(managementURL))
|
||||
parsed.RawQuery = q.Encode()
|
||||
|
||||
getReq, err := http.NewRequestWithContext(ctx, "GET", parsed.String(), nil)
|
||||
func getUploadURL(ctx context.Context, url string, managementURL string) (*types.GetURLResponse, error) {
|
||||
id := getURLHash(managementURL)
|
||||
getReq, err := http.NewRequestWithContext(ctx, "GET", url+"?id="+id, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create GET request: %w", err)
|
||||
}
|
||||
|
||||
getReq.Header.Set(types.ClientHeader, types.ClientHeaderValue)
|
||||
|
||||
resp, err := uploadClient(insecure).Do(getReq)
|
||||
resp, err := http.DefaultClient.Do(getReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get presigned URL: %w", err)
|
||||
}
|
||||
|
||||
@@ -5,7 +5,6 @@ import (
|
||||
"errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
@@ -44,7 +43,7 @@ func TestUpload(t *testing.T) {
|
||||
fileContent := []byte("test file content")
|
||||
err := os.WriteFile(file, fileContent, 0640)
|
||||
require.NoError(t, err)
|
||||
key, err := UploadDebugBundle(context.Background(), testURL+types.GetURLPath, testURL, file, true)
|
||||
key, err := UploadDebugBundle(context.Background(), testURL+types.GetURLPath, testURL, file)
|
||||
require.NoError(t, err)
|
||||
id := getURLHash(testURL)
|
||||
require.Contains(t, key, id+"/")
|
||||
@@ -80,47 +79,3 @@ func waitForServer(t *testing.T, addr string) {
|
||||
}
|
||||
t.Fatalf("server did not start listening on %s in time", addr)
|
||||
}
|
||||
|
||||
func TestRequireHTTPS(t *testing.T) {
|
||||
require.NoError(t, requireHTTPS("upload URL", "https://upload.example/path"))
|
||||
require.Error(t, requireHTTPS("upload URL", "http://upload.example/path"))
|
||||
require.Error(t, requireHTTPS("upload URL", "ftp://upload.example/path"))
|
||||
require.Error(t, requireHTTPS("upload URL", "://malformed"))
|
||||
}
|
||||
|
||||
func TestRejectInsecureRedirect(t *testing.T) {
|
||||
httpsReq, err := http.NewRequest(http.MethodGet, "https://a.example/", nil)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, rejectInsecureRedirect(httpsReq, nil), "https redirect target must be allowed")
|
||||
|
||||
httpReq, err := http.NewRequest(http.MethodGet, "http://a.example/", nil)
|
||||
require.NoError(t, err)
|
||||
require.Error(t, rejectInsecureRedirect(httpReq, nil), "http redirect target must be refused")
|
||||
|
||||
require.Error(t, rejectInsecureRedirect(httpsReq, make([]*http.Request, 10)), "the 10-redirect limit must be enforced")
|
||||
}
|
||||
|
||||
// The secure client refuses to follow an https response that redirects to http,
|
||||
// so a bundle can't be downgraded onto plaintext mid-flight.
|
||||
func TestUploadClientRefusesHTTPSToHTTPRedirect(t *testing.T) {
|
||||
plain := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
t.Cleanup(plain.Close)
|
||||
|
||||
secure := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, plain.URL, http.StatusFound)
|
||||
}))
|
||||
t.Cleanup(secure.Close)
|
||||
|
||||
client := uploadClient(false)
|
||||
// Trust the test server's cert without disabling verification globally.
|
||||
client.Transport = secure.Client().Transport
|
||||
|
||||
resp, err := client.Get(secure.URL)
|
||||
if resp != nil {
|
||||
_ = resp.Body.Close()
|
||||
}
|
||||
require.Error(t, err, "redirect from https to http must be refused")
|
||||
require.Contains(t, err.Error(), "non-https")
|
||||
}
|
||||
|
||||
@@ -1,63 +0,0 @@
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
)
|
||||
|
||||
// OpenOwnedFile opens path for reading on behalf of the IPC caller identified by
|
||||
// id, and fails unless the opened file is a regular file that id owns.
|
||||
//
|
||||
// It exists for the paths a local caller hands to the daemon over the IPC. The
|
||||
// daemon runs as root, so opening such a path unchecked lets any local user read
|
||||
// any file through it. Ownership is the invariant that keeps the daemon from
|
||||
// reading, with its own privileges, a file the caller could not read itself: a
|
||||
// symlink or hard link planted at the path resolves to a file someone else owns
|
||||
// and is refused.
|
||||
//
|
||||
// The check is made against the open descriptor rather than the path, so
|
||||
// swapping the path between the check and the read cannot change the answer.
|
||||
//
|
||||
// A privileged caller is exempt: it can read the file directly, so refusing it
|
||||
// here would protect nothing. The regular-file requirement still applies to
|
||||
// everyone, since a fifo or device planted at the path is never a log file.
|
||||
func OpenOwnedFile(id Identity, path string) (*os.File, error) {
|
||||
f, err := openForRead(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := checkOwnership(id, f); err != nil {
|
||||
if cerr := f.Close(); cerr != nil {
|
||||
return nil, fmt.Errorf("%w (close: %v)", err, cerr)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return f, nil
|
||||
}
|
||||
|
||||
func checkOwnership(id Identity, f *os.File) error {
|
||||
info, err := f.Stat()
|
||||
if err != nil {
|
||||
return fmt.Errorf("stat %s: %w", f.Name(), err)
|
||||
}
|
||||
|
||||
if !info.Mode().IsRegular() {
|
||||
return fmt.Errorf("%s is not a regular file", f.Name())
|
||||
}
|
||||
|
||||
if IsPrivilegedCaller(id) {
|
||||
return nil
|
||||
}
|
||||
|
||||
owned, err := fileOwnedBy(id, f)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read owner of %s: %w", f.Name(), err)
|
||||
}
|
||||
if !owned {
|
||||
return fmt.Errorf("%s is not owned by the caller (%s)", f.Name(), id)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -1,64 +0,0 @@
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// otherIdentity is an unprivileged caller that owns nothing the test creates.
|
||||
func otherIdentity(t *testing.T) Identity {
|
||||
t.Helper()
|
||||
if runtime.GOOS == "windows" {
|
||||
return Identity{SID: "S-1-5-21-1-2-3-1001"}
|
||||
}
|
||||
return Identity{UID: uint32(os.Geteuid() + 1), GID: uint32(os.Getegid() + 1)}
|
||||
}
|
||||
|
||||
func TestOpenOwnedFileReadsFileOwnedByCaller(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "gui-client.log")
|
||||
require.NoError(t, os.WriteFile(path, []byte("hello"), 0600))
|
||||
|
||||
id, err := CurrentProcessIdentity()
|
||||
require.NoError(t, err)
|
||||
|
||||
f, err := OpenOwnedFile(id, path)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = f.Close() })
|
||||
|
||||
content, err := io.ReadAll(f)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "hello", string(content))
|
||||
}
|
||||
|
||||
func TestOpenOwnedFileRefusesFileOwnedByAnother(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "gui-client.log")
|
||||
require.NoError(t, os.WriteFile(path, []byte("secret"), 0600))
|
||||
|
||||
_, err := OpenOwnedFile(otherIdentity(t), path)
|
||||
require.ErrorContains(t, err, "not owned by the caller")
|
||||
}
|
||||
|
||||
func TestOpenOwnedFileRefusesNonRegularFile(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
|
||||
// The caller owns the directory, so this is the regular-file requirement
|
||||
// talking, not the ownership check.
|
||||
id, err := CurrentProcessIdentity()
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = OpenOwnedFile(id, dir)
|
||||
require.ErrorContains(t, err, "not a regular file")
|
||||
}
|
||||
|
||||
func TestOpenOwnedFileRefusesMissingFile(t *testing.T) {
|
||||
id, err := CurrentProcessIdentity()
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = OpenOwnedFile(id, filepath.Join(t.TempDir(), "absent.log"))
|
||||
require.Error(t, err)
|
||||
}
|
||||
@@ -1,35 +0,0 @@
|
||||
//go:build !windows
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
// openForRead opens a caller-supplied path without following a symlink at its
|
||||
// final component and without blocking: a fifo planted at the path would
|
||||
// otherwise stall the open until a writer appears, and the daemon holds a lock
|
||||
// while it collects the file.
|
||||
func openForRead(path string) (*os.File, error) {
|
||||
f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW|syscall.O_NONBLOCK, 0)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open %s: %w", path, err)
|
||||
}
|
||||
return f, nil
|
||||
}
|
||||
|
||||
func fileOwnedBy(id Identity, f *os.File) (bool, error) {
|
||||
info, err := f.Stat()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
stat, ok := info.Sys().(*syscall.Stat_t)
|
||||
if !ok {
|
||||
return false, fmt.Errorf("no owner information in %T", info.Sys())
|
||||
}
|
||||
|
||||
return stat.Uid == id.UID, nil
|
||||
}
|
||||
@@ -1,57 +0,0 @@
|
||||
//go:build !windows
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// A symlink is the shape the arbitrary-read attempt takes: the caller owns the
|
||||
// link, the file it points at belongs to someone else.
|
||||
func TestOpenOwnedFileRefusesSymlink(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
target := filepath.Join(dir, "target.log")
|
||||
require.NoError(t, os.WriteFile(target, []byte("secret"), 0600))
|
||||
|
||||
link := filepath.Join(dir, "gui-client.log")
|
||||
require.NoError(t, os.Symlink(target, link))
|
||||
|
||||
id, err := CurrentProcessIdentity()
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = OpenOwnedFile(id, link)
|
||||
// O_NOFOLLOW on a symlink reports ELOOP on Linux/Darwin and EMLINK on FreeBSD.
|
||||
if !errors.Is(err, syscall.ELOOP) && !errors.Is(err, syscall.EMLINK) {
|
||||
t.Fatalf("symlink open: got %v, want ELOOP or EMLINK", err)
|
||||
}
|
||||
}
|
||||
|
||||
// A fifo would block the open until a writer showed up, stalling the daemon
|
||||
// while it holds its lock.
|
||||
func TestOpenOwnedFileRefusesFifoWithoutBlocking(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "gui-client.log")
|
||||
require.NoError(t, syscall.Mkfifo(path, 0600))
|
||||
|
||||
id, err := CurrentProcessIdentity()
|
||||
require.NoError(t, err)
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := OpenOwnedFile(id, path)
|
||||
done <- err
|
||||
}()
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
require.ErrorContains(t, err, "not a regular file")
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("opening a fifo blocked")
|
||||
}
|
||||
}
|
||||
@@ -1,59 +0,0 @@
|
||||
//go:build windows
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
// openForRead opens a caller-supplied path without following a reparse point at
|
||||
// it. FILE_FLAG_OPEN_REPARSE_POINT is the Windows analogue of O_NOFOLLOW: it
|
||||
// opens a symlink/junction itself rather than its target, so the regular-file
|
||||
// check in checkOwnership refuses a link the caller planted to redirect the
|
||||
// read. FILE_FLAG_BACKUP_SEMANTICS lets a directory open too (as os.Open does),
|
||||
// so a directory planted at the path is refused as non-regular rather than
|
||||
// erroring here. The share mode matches os.Open so a log being written stays
|
||||
// openable.
|
||||
func openForRead(path string) (*os.File, error) {
|
||||
p, err := windows.UTF16PtrFromString(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("convert path %s: %w", path, err)
|
||||
}
|
||||
|
||||
handle, err := windows.CreateFile(
|
||||
p,
|
||||
windows.GENERIC_READ,
|
||||
windows.FILE_SHARE_READ|windows.FILE_SHARE_WRITE|windows.FILE_SHARE_DELETE,
|
||||
nil,
|
||||
windows.OPEN_EXISTING,
|
||||
windows.FILE_FLAG_OPEN_REPARSE_POINT|windows.FILE_FLAG_BACKUP_SEMANTICS,
|
||||
0,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open %s: %w", path, err)
|
||||
}
|
||||
|
||||
return os.NewFile(uintptr(handle), path), nil
|
||||
}
|
||||
|
||||
// fileOwnedBy compares the file's owner SID with the caller's. Files an elevated
|
||||
// process creates are owned by BUILTIN\Administrators rather than by the user,
|
||||
// but such a caller is privileged and never reaches this check.
|
||||
func fileOwnedBy(id Identity, f *os.File) (bool, error) {
|
||||
// x/sys/windows GetSecurityInfo frees the OS buffer itself and returns a
|
||||
// Go-heap copy, so there is nothing to LocalFree here.
|
||||
sd, err := windows.GetSecurityInfo(windows.Handle(f.Fd()), windows.SE_FILE_OBJECT, windows.OWNER_SECURITY_INFORMATION)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("read security info: %w", err)
|
||||
}
|
||||
|
||||
owner, _, err := sd.Owner()
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("read owner: %w", err)
|
||||
}
|
||||
|
||||
return id.SID != "" && owner.String() == id.SID, nil
|
||||
}
|
||||
@@ -1,78 +0,0 @@
|
||||
//go:build windows
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
// fileOwnerSID reads the owner SID of path the same way OpenOwnedFile does, so
|
||||
// the test can construct an Identity that matches (or deliberately does not).
|
||||
func fileOwnerSID(t *testing.T, path string) string {
|
||||
t.Helper()
|
||||
f, err := os.Open(path)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = f.Close() })
|
||||
|
||||
sd, err := windows.GetSecurityInfo(windows.Handle(f.Fd()), windows.SE_FILE_OBJECT, windows.OWNER_SECURITY_INFORMATION)
|
||||
require.NoError(t, err)
|
||||
owner, _, err := sd.Owner()
|
||||
require.NoError(t, err)
|
||||
return owner.String()
|
||||
}
|
||||
|
||||
// The allow branch of fileOwnedBy is the SID-equality path the legitimate GUI
|
||||
// flow depends on. Running elevated, a created file is owned by
|
||||
// BUILTIN\Administrators; an Identity carrying that SID with Elevated=false and
|
||||
// no groups is unprivileged by IsPrivileged (which reads the token, not the
|
||||
// SID's RID), so this exercises the real GetSecurityInfo equality rather than
|
||||
// the privileged-caller shortcut.
|
||||
func TestOpenOwnedFileWindowsOwnerMatchAllows(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "gui-client.log")
|
||||
require.NoError(t, os.WriteFile(path, []byte("hello"), 0600))
|
||||
|
||||
ownerSID := fileOwnerSID(t, path)
|
||||
id := Identity{SID: ownerSID}
|
||||
require.False(t, id.IsPrivileged(), "identity built from the owner SID must be unprivileged for this to test the match path")
|
||||
|
||||
f, err := OpenOwnedFile(id, path)
|
||||
require.NoError(t, err)
|
||||
_ = f.Close()
|
||||
}
|
||||
|
||||
func TestOpenOwnedFileWindowsOwnerMismatchRefuses(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "gui-client.log")
|
||||
require.NoError(t, os.WriteFile(path, []byte("secret"), 0600))
|
||||
|
||||
other := Identity{SID: "S-1-5-21-9-9-9-9999"}
|
||||
require.False(t, other.IsPrivileged())
|
||||
|
||||
_, err := OpenOwnedFile(other, path)
|
||||
require.ErrorContains(t, err, "not owned by the caller")
|
||||
}
|
||||
|
||||
// FILE_FLAG_OPEN_REPARSE_POINT must make OpenOwnedFile refuse a symlink the same
|
||||
// way O_NOFOLLOW does on Unix, so a planted link can't redirect the read to
|
||||
// another file. Creating a symlink needs a privilege the runner may lack, so the
|
||||
// test skips rather than fails when it can't.
|
||||
func TestOpenOwnedFileWindowsRefusesSymlink(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
target := filepath.Join(dir, "target.log")
|
||||
require.NoError(t, os.WriteFile(target, []byte("secret"), 0600))
|
||||
|
||||
link := filepath.Join(dir, "gui-client.log")
|
||||
if err := os.Symlink(target, link); err != nil {
|
||||
t.Skipf("cannot create symlink (privilege not held?): %v", err)
|
||||
}
|
||||
|
||||
id, err := CurrentProcessIdentity()
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = OpenOwnedFile(id, link)
|
||||
require.Error(t, err, "a symlink must be refused")
|
||||
}
|
||||
@@ -262,7 +262,7 @@ func (c *Client) DebugBundle(anonymize bool) (string, error) {
|
||||
uploadCtx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path, false)
|
||||
key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("upload debug bundle: %w", err)
|
||||
}
|
||||
|
||||
@@ -54,7 +54,7 @@ func (e *Executor) BundleJob(ctx context.Context, debugBundleDependencies debug.
|
||||
}
|
||||
}()
|
||||
|
||||
key, err := debug.UploadDebugBundle(ctx, types.DefaultBundleURL, mgmURL, path, false)
|
||||
key, err := debug.UploadDebugBundle(ctx, types.DefaultBundleURL, mgmURL, path)
|
||||
if err != nil {
|
||||
log.Errorf("failed to upload debug bundle: %v", err)
|
||||
return "", fmt.Errorf("upload debug bundle: %w", err)
|
||||
|
||||
@@ -2771,18 +2771,14 @@ func (x *ForwardingRulesResponse) GetRules() []*ForwardingRule {
|
||||
|
||||
// DebugBundler
|
||||
type DebugBundleRequest struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Anonymize bool `protobuf:"varint,1,opt,name=anonymize,proto3" json:"anonymize,omitempty"`
|
||||
SystemInfo bool `protobuf:"varint,3,opt,name=systemInfo,proto3" json:"systemInfo,omitempty"`
|
||||
UploadURL string `protobuf:"bytes,4,opt,name=uploadURL,proto3" json:"uploadURL,omitempty"`
|
||||
LogFileCount uint32 `protobuf:"varint,5,opt,name=logFileCount,proto3" json:"logFileCount,omitempty"`
|
||||
CliVersion string `protobuf:"bytes,6,opt,name=cliVersion,proto3" json:"cliVersion,omitempty"`
|
||||
// uploadInsecure allows uploading to an http endpoint or one with an
|
||||
// untrusted TLS certificate. Restricted to privileged callers; for
|
||||
// self-hosted upload servers.
|
||||
UploadInsecure bool `protobuf:"varint,7,opt,name=uploadInsecure,proto3" json:"uploadInsecure,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Anonymize bool `protobuf:"varint,1,opt,name=anonymize,proto3" json:"anonymize,omitempty"`
|
||||
SystemInfo bool `protobuf:"varint,3,opt,name=systemInfo,proto3" json:"systemInfo,omitempty"`
|
||||
UploadURL string `protobuf:"bytes,4,opt,name=uploadURL,proto3" json:"uploadURL,omitempty"`
|
||||
LogFileCount uint32 `protobuf:"varint,5,opt,name=logFileCount,proto3" json:"logFileCount,omitempty"`
|
||||
CliVersion string `protobuf:"bytes,6,opt,name=cliVersion,proto3" json:"cliVersion,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *DebugBundleRequest) Reset() {
|
||||
@@ -2850,13 +2846,6 @@ func (x *DebugBundleRequest) GetCliVersion() string {
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *DebugBundleRequest) GetUploadInsecure() bool {
|
||||
if x != nil {
|
||||
return x.UploadInsecure
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
type DebugBundleResponse struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Path string `protobuf:"bytes,1,opt,name=path,proto3" json:"path,omitempty"`
|
||||
@@ -7253,7 +7242,7 @@ const file_daemon_proto_rawDesc = "" +
|
||||
"\x12translatedHostname\x18\x04 \x01(\tR\x12translatedHostname\x128\n" +
|
||||
"\x0etranslatedPort\x18\x05 \x01(\v2\x10.daemon.PortInfoR\x0etranslatedPort\"G\n" +
|
||||
"\x17ForwardingRulesResponse\x12,\n" +
|
||||
"\x05rules\x18\x01 \x03(\v2\x16.daemon.ForwardingRuleR\x05rules\"\xdc\x01\n" +
|
||||
"\x05rules\x18\x01 \x03(\v2\x16.daemon.ForwardingRuleR\x05rules\"\xb4\x01\n" +
|
||||
"\x12DebugBundleRequest\x12\x1c\n" +
|
||||
"\tanonymize\x18\x01 \x01(\bR\tanonymize\x12\x1e\n" +
|
||||
"\n" +
|
||||
@@ -7263,8 +7252,7 @@ const file_daemon_proto_rawDesc = "" +
|
||||
"\flogFileCount\x18\x05 \x01(\rR\flogFileCount\x12\x1e\n" +
|
||||
"\n" +
|
||||
"cliVersion\x18\x06 \x01(\tR\n" +
|
||||
"cliVersion\x12&\n" +
|
||||
"\x0euploadInsecure\x18\a \x01(\bR\x0euploadInsecure\"}\n" +
|
||||
"cliVersion\"}\n" +
|
||||
"\x13DebugBundleResponse\x12\x12\n" +
|
||||
"\x04path\x18\x01 \x01(\tR\x04path\x12 \n" +
|
||||
"\vuploadedKey\x18\x02 \x01(\tR\vuploadedKey\x120\n" +
|
||||
|
||||
@@ -536,10 +536,6 @@ message DebugBundleRequest {
|
||||
string uploadURL = 4;
|
||||
uint32 logFileCount = 5;
|
||||
string cliVersion = 6;
|
||||
// uploadInsecure allows uploading to an http endpoint or one with an
|
||||
// untrusted TLS certificate. Restricted to privileged callers; for
|
||||
// self-hosted upload servers.
|
||||
bool uploadInsecure = 7;
|
||||
}
|
||||
|
||||
message DebugBundleResponse {
|
||||
|
||||
@@ -7,62 +7,18 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"runtime/pprof"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"google.golang.org/grpc/codes"
|
||||
gstatus "google.golang.org/grpc/status"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/debug"
|
||||
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/version"
|
||||
)
|
||||
|
||||
// DebugBundle creates a debug bundle and returns the location.
|
||||
func (s *Server) DebugBundle(callerCtx context.Context, req *proto.DebugBundleRequest) (resp *proto.DebugBundleResponse, err error) {
|
||||
if err := requirePrivilegeForUploadURL(callerCtx, req.GetUploadURL(), req.GetUploadInsecure()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// The UI log is opened as whoever asked for this bundle, so a caller only
|
||||
// collects a log it owns (privileged callers excepted). ok is false on a
|
||||
// socket that carries no identity, which skips the UI log.
|
||||
callerID, callerIdentified := ipcauth.CallerIdentity(callerCtx)
|
||||
|
||||
path, managementURL, err := s.generateDebugBundle(req, uiLogOpener(callerID, callerIdentified))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if req.GetUploadURL() == "" {
|
||||
return &proto.DebugBundleResponse{Path: path}, nil
|
||||
}
|
||||
|
||||
// The upload runs without s.mutex held: it does network I/O to a possibly
|
||||
// slow destination and must not block the other RPCs that take the lock. The
|
||||
// bounded context is a backstop against a hung connection.
|
||||
uploadCtx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
defer cancel()
|
||||
key, err := debug.UploadDebugBundle(uploadCtx, req.GetUploadURL(), managementURL, path, req.GetUploadInsecure())
|
||||
if err != nil {
|
||||
log.Errorf("failed to upload debug bundle to %s: %v", req.GetUploadURL(), err)
|
||||
return &proto.DebugBundleResponse{Path: path, UploadFailureReason: err.Error()}, nil
|
||||
}
|
||||
|
||||
log.Infof("debug bundle uploaded to %s with key %s", req.GetUploadURL(), key)
|
||||
|
||||
return &proto.DebugBundleResponse{Path: path, UploadedKey: key}, nil
|
||||
}
|
||||
|
||||
// generateDebugBundle builds the bundle under s.mutex and returns its path plus
|
||||
// the management URL captured under the lock, so the caller can run the upload
|
||||
// without holding the lock.
|
||||
func (s *Server) generateDebugBundle(req *proto.DebugBundleRequest, uiOpener debug.LogOpener) (path string, managementURL string, err error) {
|
||||
func (s *Server) DebugBundle(_ context.Context, req *proto.DebugBundleRequest) (resp *proto.DebugBundleResponse, err error) {
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
|
||||
@@ -112,7 +68,6 @@ func (s *Server) generateDebugBundle(req *proto.DebugBundleRequest, uiOpener deb
|
||||
SyncResponse: syncResponse,
|
||||
LogPath: s.logFile,
|
||||
UILogPath: s.uiLogPath,
|
||||
UILogOpener: uiOpener,
|
||||
CPUProfile: cpuProfileData,
|
||||
CapturePath: capturePath,
|
||||
RefreshStatus: refreshStatus,
|
||||
@@ -127,16 +82,23 @@ func (s *Server) generateDebugBundle(req *proto.DebugBundleRequest, uiOpener deb
|
||||
},
|
||||
)
|
||||
|
||||
path, err = bundleGenerator.Generate()
|
||||
path, err := bundleGenerator.Generate()
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("generate debug bundle: %w", err)
|
||||
return nil, fmt.Errorf("generate debug bundle: %w", err)
|
||||
}
|
||||
|
||||
if s.config != nil && s.config.ManagementURL != nil {
|
||||
managementURL = s.config.ManagementURL.String()
|
||||
if req.GetUploadURL() == "" {
|
||||
return &proto.DebugBundleResponse{Path: path}, nil
|
||||
}
|
||||
key, err := debug.UploadDebugBundle(context.Background(), req.GetUploadURL(), s.config.ManagementURL.String(), path)
|
||||
if err != nil {
|
||||
log.Errorf("failed to upload debug bundle to %s: %v", req.GetUploadURL(), err)
|
||||
return &proto.DebugBundleResponse{Path: path, UploadFailureReason: err.Error()}, nil
|
||||
}
|
||||
|
||||
return path, managementURL, nil
|
||||
log.Infof("debug bundle uploaded to %s with key %s", req.GetUploadURL(), key)
|
||||
|
||||
return &proto.DebugBundleResponse{Path: path, UploadedKey: key}, nil
|
||||
}
|
||||
|
||||
// GetLogLevel gets the current logging level for the server.
|
||||
@@ -176,34 +138,12 @@ func (s *Server) SetLogLevel(_ context.Context, req *proto.SetLogLevelRequest) (
|
||||
// RegisterUILog records the desktop UI's absolute log path so DebugBundle can
|
||||
// collect the GUI log. The daemon runs as root and can't resolve the user's
|
||||
// config dir, so the UI reports it. Last-writer-wins (one UI per socket).
|
||||
//
|
||||
// The path arrives over an IPC any local user can reach and is later opened by
|
||||
// a root daemon, so it is constrained to the file name the UI writes and to a
|
||||
// local absolute path. Authorization happens when DebugBundle opens it: the
|
||||
// bundle refuses a file its requester does not own. A caller the daemon cannot
|
||||
// identify cannot register a path at all.
|
||||
func (s *Server) RegisterUILog(callerCtx context.Context, req *proto.RegisterUILogRequest) (*proto.RegisterUILogResponse, error) {
|
||||
if _, ok := ipcauth.CallerIdentity(callerCtx); !ok {
|
||||
return nil, gstatus.Error(codes.PermissionDenied,
|
||||
"registering a UI log path requires a control channel that carries the caller's identity")
|
||||
}
|
||||
|
||||
path := filepath.Clean(req.GetPath())
|
||||
if !filepath.IsAbs(path) || filepath.Base(path) != uiLogFileName {
|
||||
return nil, gstatus.Errorf(codes.InvalidArgument, "UI log path must be an absolute path ending in %s", uiLogFileName)
|
||||
}
|
||||
// filepath.IsAbs accepts a Windows UNC path (\\host\share\...) and a device
|
||||
// path (\\.\, \\?\); opening one would make the root daemon reach a remote
|
||||
// or device namespace. Require a plain local path.
|
||||
if strings.HasPrefix(path, `\\`) {
|
||||
return nil, gstatus.Error(codes.InvalidArgument, "UI log path must be a local path, not a UNC or device path")
|
||||
}
|
||||
|
||||
func (s *Server) RegisterUILog(_ context.Context, req *proto.RegisterUILogRequest) (*proto.RegisterUILogResponse, error) {
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
|
||||
s.uiLogPath = path
|
||||
log.Infof("registered UI log path %s", s.uiLogPath)
|
||||
s.uiLogPath = req.GetPath()
|
||||
log.Infof("registered UI log path: %s", s.uiLogPath)
|
||||
|
||||
return &proto.RegisterUILogResponse{}, nil
|
||||
}
|
||||
|
||||
@@ -1,99 +0,0 @@
|
||||
//go:build !android && !ios
|
||||
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"google.golang.org/grpc/codes"
|
||||
gstatus "google.golang.org/grpc/status"
|
||||
|
||||
"github.com/netbirdio/netbird/client/configs"
|
||||
"github.com/netbirdio/netbird/client/internal/debug"
|
||||
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||
"github.com/netbirdio/netbird/upload-server/types"
|
||||
)
|
||||
|
||||
// uiLogFileName is the only file name the daemon accepts as a UI log path. The
|
||||
// UI (writer), this validation, and the bundle collector all read it from
|
||||
// configs so they cannot drift.
|
||||
const uiLogFileName = configs.UILogFile
|
||||
|
||||
// uiLogOpener opens the registered UI log, and its rotated siblings, on behalf
|
||||
// of the caller requesting the bundle: OpenOwnedFile then collects the log only
|
||||
// when that caller owns it (or is privileged). identified is false on a socket
|
||||
// that carries no caller identity, in which case nothing is opened.
|
||||
func uiLogOpener(id ipcauth.Identity, identified bool) debug.LogOpener {
|
||||
return func(path string) (*os.File, error) {
|
||||
if !identified {
|
||||
return nil, fmt.Errorf("bundle requester has no verified identity")
|
||||
}
|
||||
return ipcauth.OpenOwnedFile(id, path)
|
||||
}
|
||||
}
|
||||
|
||||
// requirePrivilegeForUploadURL restricts where the daemon may send a debug
|
||||
// bundle. The bundle holds the daemon's own logs and state, and the daemon
|
||||
// fetches the upload URL itself, so an unrestricted endpoint turns the daemon
|
||||
// into both an exfiltration channel and a request forwarder that reaches
|
||||
// services only it can talk to.
|
||||
//
|
||||
// The upload service NetBird publishes is open to any caller, since that is what
|
||||
// the CLI and the desktop UI use. Any other endpoint, self-hosted upload servers
|
||||
// included, requires a privileged caller. Plaintext is refused for everyone: the
|
||||
// daemon fetches the URL and then PUTs the bundle to whatever that fetch returns,
|
||||
// so an http hop is a place to intercept the bundle or the redirect.
|
||||
//
|
||||
// insecure relaxes transport security (http, or an untrusted TLS certificate)
|
||||
// for a self-hosted server. It weakens a root-privileged upload, so it is
|
||||
// refused for an unprivileged caller regardless of the host.
|
||||
func requirePrivilegeForUploadURL(ctx context.Context, rawURL string, insecure bool) error {
|
||||
if rawURL == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
parsed, err := url.Parse(rawURL)
|
||||
if err != nil {
|
||||
return gstatus.Errorf(codes.InvalidArgument, "parse upload URL: %v", err)
|
||||
}
|
||||
|
||||
// --insecure relaxes https to http or an untrusted certificate; it does not
|
||||
// widen the URL to arbitrary schemes, so a host and http/https are required
|
||||
// before the insecure branch takes over.
|
||||
if parsed.Host == "" || (parsed.Scheme != "https" && parsed.Scheme != "http") {
|
||||
return gstatus.Errorf(codes.InvalidArgument, "upload URL must be http or https with a host")
|
||||
}
|
||||
|
||||
if insecure {
|
||||
return denyPrivileged(ctx,
|
||||
"uploading a debug bundle without transport security (--upload-bundle-insecure)",
|
||||
ipcauth.ElevatedCommand("netbird debug bundle -U --upload-bundle-insecure --upload-bundle-url <url>"))
|
||||
}
|
||||
|
||||
if parsed.Scheme != "https" {
|
||||
return gstatus.Errorf(codes.InvalidArgument, "upload URL must use https, got scheme %q", parsed.Scheme)
|
||||
}
|
||||
|
||||
if isDefaultUploadService(parsed) {
|
||||
return nil
|
||||
}
|
||||
|
||||
return denyPrivileged(ctx,
|
||||
"uploading a debug bundle to an upload service other than the default one",
|
||||
ipcauth.ElevatedCommand("netbird debug bundle -U --upload-bundle-url <url>"))
|
||||
}
|
||||
|
||||
// isDefaultUploadService reports whether the URL points at the upload service
|
||||
// NetBird runs. Only the host is compared: the service's path may differ between
|
||||
// releases, and the host is what decides who receives the bundle.
|
||||
func isDefaultUploadService(parsed *url.URL) bool {
|
||||
defaultURL, err := url.Parse(types.DefaultBundleURL)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return parsed.Scheme == defaultURL.Scheme && strings.EqualFold(parsed.Host, defaultURL.Host)
|
||||
}
|
||||
@@ -1,157 +0,0 @@
|
||||
//go:build !android && !ios
|
||||
|
||||
package server
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"testing"
|
||||
|
||||
"google.golang.org/grpc/codes"
|
||||
gstatus "google.golang.org/grpc/status"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
"github.com/netbirdio/netbird/upload-server/types"
|
||||
)
|
||||
|
||||
func TestRegisterUILogRefusesUnidentifiedCaller(t *testing.T) {
|
||||
s := &Server{}
|
||||
|
||||
_, err := s.RegisterUILog(noIdentityCtx(), &proto.RegisterUILogRequest{
|
||||
Path: filepath.Join(t.TempDir(), uiLogFileName),
|
||||
})
|
||||
|
||||
if gstatus.Code(err) != codes.PermissionDenied {
|
||||
t.Fatalf("code = %v, want PermissionDenied", gstatus.Code(err))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterUILogRefusesForeignPath(t *testing.T) {
|
||||
secret := "/etc/shadow"
|
||||
if runtime.GOOS == "windows" {
|
||||
secret = `C:\Windows\System32\config\SAM`
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
path string
|
||||
}{
|
||||
{"empty", ""},
|
||||
{"relative", filepath.Join("netbird", uiLogFileName)},
|
||||
{"another file", secret},
|
||||
{"directory of the log", t.TempDir()},
|
||||
{"unc path", `\\attacker\share\` + uiLogFileName},
|
||||
{"device path", `\\.\C:\` + uiLogFileName},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
s := &Server{}
|
||||
|
||||
_, err := s.RegisterUILog(userCtx(), &proto.RegisterUILogRequest{Path: tc.path})
|
||||
|
||||
if gstatus.Code(err) != codes.InvalidArgument {
|
||||
t.Fatalf("code = %v, want InvalidArgument", gstatus.Code(err))
|
||||
}
|
||||
if s.uiLogPath != "" {
|
||||
t.Fatalf("path %q was recorded despite the refusal", s.uiLogPath)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterUILogRecordsPath(t *testing.T) {
|
||||
s := &Server{}
|
||||
path := filepath.Join(t.TempDir(), uiLogFileName)
|
||||
|
||||
if _, err := s.RegisterUILog(userCtx(), &proto.RegisterUILogRequest{Path: path}); err != nil {
|
||||
t.Fatalf("register: %v", err)
|
||||
}
|
||||
|
||||
if s.uiLogPath != path {
|
||||
t.Fatalf("path = %q, want %q", s.uiLogPath, path)
|
||||
}
|
||||
}
|
||||
|
||||
// The UI log is opened as the bundle requester, so a second local user cannot
|
||||
// collect a log they do not own, and an unidentified requester collects nothing.
|
||||
func TestUILogOpenerBindsToRequester(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), uiLogFileName)
|
||||
if err := os.WriteFile(path, []byte("log line"), 0600); err != nil {
|
||||
t.Fatalf("write log: %v", err)
|
||||
}
|
||||
|
||||
// A different unprivileged user than the file's owner: refused.
|
||||
if _, err := uiLogOpener(unprivilegedIdentity(), true)(path); err == nil {
|
||||
t.Fatal("expected a file the requester does not own to be refused")
|
||||
}
|
||||
|
||||
// No verified identity: refused.
|
||||
if _, err := uiLogOpener(ipcauth.Identity{}, false)(path); err == nil {
|
||||
t.Fatal("expected an unidentified requester to be refused")
|
||||
}
|
||||
|
||||
// The requester that owns the file: allowed. The test process created it, so
|
||||
// its own identity is the owner (and a privileged runner is exempt anyway).
|
||||
owner, err := ipcauth.CurrentProcessIdentity()
|
||||
if err != nil {
|
||||
t.Fatalf("current identity: %v", err)
|
||||
}
|
||||
f, err := uiLogOpener(owner, true)(path)
|
||||
if err != nil {
|
||||
t.Fatalf("expected the owning requester to be allowed, got %v", err)
|
||||
}
|
||||
_ = f.Close()
|
||||
}
|
||||
|
||||
func TestRequirePrivilegeForUploadURL(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
url string
|
||||
insecure bool
|
||||
unprivOK bool
|
||||
invalid bool
|
||||
rootAlso bool
|
||||
}{
|
||||
{name: "no upload", url: "", unprivOK: true},
|
||||
{name: "default service", url: types.DefaultBundleURL, unprivOK: true},
|
||||
{name: "default service, other path", url: "https://upload.debug.netbird.io/other", unprivOK: true},
|
||||
{name: "loopback exfiltration endpoint", url: "https://127.0.0.1:8080/upload-url", rootAlso: true},
|
||||
{name: "custom upload service", url: "https://attacker.example/upload-url", rootAlso: true},
|
||||
{name: "plaintext default host", url: "http://upload.debug.netbird.io/upload-url", invalid: true},
|
||||
{name: "plaintext custom host", url: "http://attacker.example/upload-url", invalid: true},
|
||||
{name: "unsupported scheme", url: "file:///etc/shadow", invalid: true},
|
||||
// insecure relaxes transport security; privileged only, whatever the host.
|
||||
{name: "insecure http custom", url: "http://selfhosted.local/upload-url", insecure: true, rootAlso: true},
|
||||
{name: "insecure https custom", url: "https://selfhosted.local/upload-url", insecure: true, rootAlso: true},
|
||||
{name: "insecure default host", url: types.DefaultBundleURL, insecure: true, rootAlso: true},
|
||||
// --insecure must not widen the URL to non-http(s) schemes or a hostless URL.
|
||||
{name: "insecure file scheme", url: "file:///etc/shadow", insecure: true, invalid: true},
|
||||
{name: "insecure hostless", url: "https:///upload-url", insecure: true, invalid: true},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := requirePrivilegeForUploadURL(userCtx(), tc.url, tc.insecure)
|
||||
|
||||
switch {
|
||||
case tc.invalid:
|
||||
if gstatus.Code(err) != codes.InvalidArgument {
|
||||
t.Fatalf("code = %v, want InvalidArgument", gstatus.Code(err))
|
||||
}
|
||||
return
|
||||
case tc.unprivOK:
|
||||
assertAllowed(t, err)
|
||||
return
|
||||
default:
|
||||
assertDenied(t, err)
|
||||
}
|
||||
|
||||
if tc.rootAlso {
|
||||
assertAllowed(t, requirePrivilegeForUploadURL(rootCtx(), tc.url, tc.insecure))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -72,9 +72,6 @@ type Server struct {
|
||||
// RegisterUILog. Guarded by mutex. Consumed by DebugBundle so the bundle
|
||||
// can collect the GUI log even though the daemon runs as root and can't
|
||||
// resolve the user's config dir. Last-writer-wins (one UI per socket).
|
||||
// DebugBundle opens it on behalf of the bundle requester and refuses a file
|
||||
// that caller does not own, so a local user cannot read another user's log
|
||||
// or a root-only file through it.
|
||||
uiLogPath string
|
||||
|
||||
oauthAuthFlow oauthAuthFlow
|
||||
|
||||
@@ -8,19 +8,21 @@ import (
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/client/configs"
|
||||
"github.com/netbirdio/netbird/client/ui/guilog"
|
||||
)
|
||||
|
||||
// uiLogFileName must stay in sync with the daemon's "gui-client*.log.*" glob
|
||||
// for rotated siblings (addUILog in client/internal/debug).
|
||||
const uiLogFileName = "gui-client.log"
|
||||
|
||||
// uiLogPath returns the GUI log path with native separators, since the daemon
|
||||
// opens it directly for debug-bundle collection. The file name comes from
|
||||
// configs.UILogFile so the daemon validates and collects the same name.
|
||||
// opens it directly for debug-bundle collection.
|
||||
func uiLogPath() (string, error) {
|
||||
dir, err := os.UserConfigDir()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.Join(dir, "netbird", configs.UILogFile), nil
|
||||
return filepath.Join(dir, "netbird", uiLogFileName), nil
|
||||
}
|
||||
|
||||
// newDebugLog builds the GUI debug log, disabled when userSetLogFile is set
|
||||
|
||||
@@ -76,13 +76,6 @@ type ServerConfig struct {
|
||||
|
||||
SupportedSyncMessageVersions *int `yaml:"supportedSyncMessageVersions,omitempty"`
|
||||
PerAccountSupportedSyncMessageVersions map[string]int `yaml:"perAccountSupportedSyncMessageVersions,omitempty"`
|
||||
|
||||
AgentNetwork AgentNetworkConfig `yaml:"agentNetwork"`
|
||||
}
|
||||
|
||||
// AgentNetworkConfig contains agent-network (LLM gateway) configuration.
|
||||
type AgentNetworkConfig struct {
|
||||
PricingDefaultsFile string `yaml:"pricingDefaultsFile"`
|
||||
}
|
||||
|
||||
// TLSConfig contains TLS/HTTPS settings
|
||||
@@ -730,9 +723,6 @@ func (c *CombinedConfig) ToManagementConfig() (*nbconfig.Config, error) {
|
||||
EmbeddedIdP: embeddedIdP,
|
||||
HighestSupportedSyncMessageVersion: c.Server.SupportedSyncMessageVersions,
|
||||
PerAccountHighestSupportedSyncMessageVersion: c.Server.PerAccountSupportedSyncMessageVersions,
|
||||
AgentNetwork: nbconfig.AgentNetwork{
|
||||
PricingDefaultsFile: c.Server.AgentNetwork.PricingDefaultsFile,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -10,7 +10,6 @@ import (
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -25,7 +24,6 @@ import (
|
||||
"google.golang.org/grpc"
|
||||
|
||||
"github.com/netbirdio/netbird/encryption"
|
||||
agentnetworkpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
mgmtServer "github.com/netbirdio/netbird/management/internals/server"
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
"github.com/netbirdio/netbird/management/server/telemetry"
|
||||
@@ -290,11 +288,6 @@ func (s *serverInstances) createManagementServer(ctx context.Context, cfg *Combi
|
||||
return fmt.Errorf("failed to ensure encryption key: %w", err)
|
||||
}
|
||||
|
||||
if err := loadAgentNetworkPricing(ctx, mgmtConfig); err != nil {
|
||||
cleanupSTUNListeners(s.stunListeners)
|
||||
return fmt.Errorf("failed to load agent-network pricing defaults: %w", err)
|
||||
}
|
||||
|
||||
LogConfigInfo(mgmtConfig)
|
||||
|
||||
s.mgmtSrv, err = createManagementServer(cfg, mgmtConfig)
|
||||
@@ -629,32 +622,6 @@ func handleRelayWebSocket(w http.ResponseWriter, r *http.Request, acceptFn func(
|
||||
acceptFn(conn)
|
||||
}
|
||||
|
||||
// loadAgentNetworkPricing loads the management-side LLM pricing defaults
|
||||
// file for the combined server and starts its periodic reloader. An
|
||||
// explicitly configured PricingDefaultsFile is required to load (a typo
|
||||
// must fail startup rather than silently bill with built-ins the operator
|
||||
// believes they replaced); a relative path is resolved against the data
|
||||
// directory so a bare filename like "pricing.yaml" lands in the datadir
|
||||
// alongside the store. With no path configured, <datadir>/<DefaultFileName>
|
||||
// is probed and may be absent (compiled-in defaults serve).
|
||||
func loadAgentNetworkPricing(ctx context.Context, mgmtConfig *nbconfig.Config) error {
|
||||
pricingPath := mgmtConfig.AgentNetwork.PricingDefaultsFile
|
||||
required := pricingPath != ""
|
||||
if !required {
|
||||
pricingPath = agentnetworkpricing.DefaultFileName
|
||||
}
|
||||
if !filepath.IsAbs(pricingPath) {
|
||||
pricingPath = filepath.Join(mgmtConfig.Datadir, pricingPath)
|
||||
}
|
||||
|
||||
log.Infof("loading agent-network pricing defaults from %s (required: %v)", pricingPath, required)
|
||||
if err := agentnetworkpricing.LoadFile(pricingPath, required); err != nil {
|
||||
return err
|
||||
}
|
||||
agentnetworkpricing.StartReloader(ctx, agentnetworkpricing.ReloadInterval)
|
||||
return nil
|
||||
}
|
||||
|
||||
// logConfig prints all configuration parameters for debugging
|
||||
func logConfig(cfg *CombinedConfig) {
|
||||
log.Info("=== Configuration ===")
|
||||
@@ -731,25 +698,6 @@ func logManagementConfig(cfg *CombinedConfig) {
|
||||
log.Infof(" Relay addresses: %v", cfg.Management.Relays.Addresses)
|
||||
log.Infof(" Relay credentials TTL: %s", cfg.Management.Relays.CredentialsTTL)
|
||||
}
|
||||
|
||||
logAgentNetworkConfig(cfg)
|
||||
}
|
||||
|
||||
func logAgentNetworkConfig(cfg *CombinedConfig) {
|
||||
log.Info(" Agent Network:")
|
||||
pricingPath := cfg.Server.AgentNetwork.PricingDefaultsFile
|
||||
configured := pricingPath != ""
|
||||
if !configured {
|
||||
pricingPath = agentnetworkpricing.DefaultFileName
|
||||
}
|
||||
if !filepath.IsAbs(pricingPath) {
|
||||
pricingPath = filepath.Join(cfg.Management.DataDir, pricingPath)
|
||||
}
|
||||
if configured {
|
||||
log.Infof(" Pricing defaults file: %s", pricingPath)
|
||||
} else {
|
||||
log.Infof(" Pricing defaults file: %s (default, optional)", pricingPath)
|
||||
}
|
||||
}
|
||||
|
||||
// logEnvVars logs all NB_ environment variables that are currently set
|
||||
|
||||
@@ -134,16 +134,3 @@ server:
|
||||
# trustedPeers: [] # CIDRs of trusted peer networks (e.g. ["100.64.0.0/10"])
|
||||
# accessLogRetentionDays: 7 # Days to retain HTTP access logs. 0 (or unset) defaults to 7. Negative values disable cleanup (logs kept indefinitely).
|
||||
# accessLogCleanupIntervalHours: 24 # How often (in hours) to run the access-log cleanup job. 0 (or unset) is treated as "not set" and defaults to 24 hours; cleanup remains enabled. To disable cleanup, set accessLogRetentionDays to a negative value.
|
||||
|
||||
# Agent network (LLM gateway) settings (optional)
|
||||
# agentNetwork:
|
||||
# # Path to the YAML file holding the default LLM pricing table. A relative
|
||||
# # path is resolved against dataDir, so a bare filename like "pricing.yaml"
|
||||
# # lands in the data directory. When empty, {dataDir}/defaults_llm_pricing.yaml
|
||||
# # is probed; if no file is present the compiled-in defaults are used.
|
||||
# # Schema: surface ("openai"/"anthropic"/"bedrock") -> model -> rates in USD
|
||||
# # per 1k tokens (input_per_1k, output_per_1k, and the optional
|
||||
# # cached_input_per_1k / cache_read_per_1k / cache_creation_per_1k). The file
|
||||
# # is re-read periodically (mtime poll). An explicitly configured path that
|
||||
# # fails to load fails startup; runtime reload errors keep the previous table.
|
||||
# pricingDefaultsFile: "pricing.yaml"
|
||||
|
||||
@@ -23,13 +23,8 @@ import (
|
||||
type per1k struct{ in, out, read, write float64 }
|
||||
|
||||
// publishedPer1k hardcodes the vendors' PUBLISHED rates for the models the live matrix can drive,
|
||||
// keyed by the normalized model id the proxy stamps. Deliberately independent of NetBird's own
|
||||
// default pricing table so a wrong default rate or a broken normalization fails the run.
|
||||
//
|
||||
// These rates are also what providerRequest registers as the operator's per-model prices. Since
|
||||
// management now ships operator prices to the cost meter as a per-provider-record table that is
|
||||
// consulted BEFORE the surface defaults, registering the published rate is what keeps this matrix
|
||||
// asserting vendor rates — and exercises the per-record path at the same time.
|
||||
// keyed by the normalized model id the proxy stamps. Deliberately independent of the proxy's
|
||||
// pricing table so a wrong embedded rate or a broken normalization fails the run.
|
||||
var publishedPer1k = map[string]per1k{
|
||||
"gpt-4o-mini": {0.00015, 0.0006, 0.000075, 0},
|
||||
"gpt-4o": {0.0025, 0.01, 0.00125, 0},
|
||||
@@ -40,22 +35,12 @@ var publishedPer1k = map[string]per1k{
|
||||
"anthropic.claude-haiku-4-5": {0.001, 0.005, 0.0001, 0.00125},
|
||||
"anthropic.claude-sonnet-4-5": {0.003, 0.015, 0.0003, 0.00375},
|
||||
"anthropic.claude-sonnet-4-6": {0.003, 0.015, 0.0003, 0.00375},
|
||||
// Gateway-prefixed ids (Vercel AI Gateway, OpenRouter). A gateway model is not in
|
||||
// NetBird's default table, so before operator pricing it could only be recorded at
|
||||
// cost 0. The operator names it and prices it — at the underlying vendor's published
|
||||
// rate, which is what the gateway charges through — so these rows are now priced.
|
||||
"openai/gpt-4o-mini": {0.00015, 0.0006, 0.000075, 0},
|
||||
"openai/gpt-4o": {0.0025, 0.01, 0.00125, 0},
|
||||
}
|
||||
|
||||
// rawCostVerificationSQL is the operator-facing double-check, run straight against the management
|
||||
// sqlite store: recompute each usage row's expected total and cache cost from its own persisted
|
||||
// token buckets and hardcoded published rates. OpenAI counts cached tokens as a subset of input;
|
||||
// Anthropic-shape providers count cache buckets additively.
|
||||
//
|
||||
// The rate rows must stay in sync with publishedPer1k — they are the same vendor rates the matrix
|
||||
// registers as operator prices. The join is on model, so rows written by other tests in this
|
||||
// package (which price their own made-up model ids) are simply not covered here.
|
||||
const rawCostVerificationSQL = `
|
||||
WITH rates(model, in_rate, out_rate, read_rate, write_rate) AS (
|
||||
VALUES
|
||||
@@ -67,9 +52,7 @@ WITH rates(model, in_rate, out_rate, read_rate, write_rate) AS (
|
||||
('kimi-k3', 0.003, 0.015, 0.0003, 0.003),
|
||||
('anthropic.claude-haiku-4-5', 0.001, 0.005, 0.0001, 0.00125),
|
||||
('anthropic.claude-sonnet-4-5', 0.003, 0.015, 0.0003, 0.00375),
|
||||
('anthropic.claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375),
|
||||
('openai/gpt-4o-mini', 0.00015, 0.0006, 0.000075, 0.0),
|
||||
('openai/gpt-4o', 0.0025, 0.01, 0.00125, 0.0)
|
||||
('anthropic.claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375)
|
||||
)
|
||||
SELECT
|
||||
u.provider,
|
||||
@@ -163,11 +146,6 @@ func verifyUsageRowsSQL(t *testing.T, srv *harness.Combined) {
|
||||
require.Positive(t, verified, "raw SQL check must cover at least one usage row")
|
||||
t.Logf("[sql] verified %d usage rows in store.db against published rates", verified)
|
||||
|
||||
// Gateway-prefixed model ids are absent from NetBird's default pricing table, so they are
|
||||
// priced only because the operator registered and priced them on the provider record. Assert
|
||||
// they are priced (not silently 0) — the join above already checked the exact figures for the
|
||||
// ones this matrix drives. A gateway row at cost 0 means the per-record table never reached
|
||||
// the cost meter, which is the regression this guards.
|
||||
gwRows, err := db.Raw(`SELECT model,
|
||||
(input_cost_usd + cached_input_cost_usd + cache_creation_cost_usd + output_cost_usd) AS cost_usd
|
||||
FROM agent_network_request_usage WHERE model LIKE '%/%'`).Rows()
|
||||
@@ -177,8 +155,8 @@ func verifyUsageRowsSQL(t *testing.T, srv *harness.Combined) {
|
||||
var model string
|
||||
var cost float64
|
||||
require.NoError(t, gwRows.Scan(&model, &cost), "scan gateway usage row")
|
||||
t.Logf("[sql] gateway %s: stored=$%.6f (priced from the operator's per-record rate)", model, cost)
|
||||
assert.Positivef(t, cost, "gateway-prefixed model %q is priced on the provider record, so its cost must be > 0", model)
|
||||
t.Logf("[sql] gateway %s: stored=$%.6f (must be 0 — deliberately unpriced)", model, cost)
|
||||
assert.Zerof(t, cost, "gateway-prefixed model %q must store cost 0, never a guessed rate", model)
|
||||
}
|
||||
require.NoError(t, gwRows.Err(), "iterate gateway usage rows")
|
||||
}
|
||||
@@ -199,6 +177,10 @@ func validateAccessLogCost(t *testing.T, pc providerCase, row api.AgentNetworkAc
|
||||
|
||||
rates, known := publishedPer1k[model]
|
||||
if !known {
|
||||
if strings.Contains(model, "/") {
|
||||
assert.Zerof(t, row.CostUsd, "gateway-prefixed model %q is not priced so the cost meter must skip (cost 0)", model)
|
||||
return
|
||||
}
|
||||
t.Logf("[cost] %s: no published rate on file for model %q (env-overridden?); skipping cost validation", pc.name, model)
|
||||
return
|
||||
}
|
||||
@@ -355,17 +337,8 @@ func availableProviders() []providerCase {
|
||||
}
|
||||
|
||||
// providerRequest builds a create request for a matrix provider: enabled, with
|
||||
// its model registered at the vendor's published rates for body-routed
|
||||
// providers, and no models for the path-routed Vertex (whose model lives in the
|
||||
// request path, so it prices from the defaults table management ships).
|
||||
//
|
||||
// The registered rates matter: management synthesizes them into the cost
|
||||
// meter's per-provider-record table, which is consulted before the surface
|
||||
// defaults, so these are the rates the proxy actually bills with. Registering
|
||||
// the published rate keeps the cost assertions vendor-anchored while covering
|
||||
// the operator-pricing path. A model with no published rate on file (an
|
||||
// env-overridden Bedrock profile) falls back to a nominal rate, and
|
||||
// validateAccessLogCost skips its cost check.
|
||||
// a uniquely-priced model for body-routed providers and none for the
|
||||
// path-routed Vertex (whose model lives in the request path).
|
||||
func providerRequest(pc providerCase) api.AgentNetworkProviderRequest {
|
||||
req := api.AgentNetworkProviderRequest{
|
||||
Name: pc.name,
|
||||
@@ -383,23 +356,9 @@ func providerRequest(pc providerCase) api.AgentNetworkProviderRequest {
|
||||
if pc.kind == harness.WireBedrock {
|
||||
modelID = catalogModel(pc)
|
||||
}
|
||||
model := api.AgentNetworkProviderModel{Id: modelID, InputPer1k: 0.001, OutputPer1k: 0.002}
|
||||
if rates, known := publishedPer1k[catalogModel(pc)]; known {
|
||||
model.InputPer1k = rates.in
|
||||
model.OutputPer1k = rates.out
|
||||
// Pin the cache rates too, rather than letting them inherit from the
|
||||
// defaults table: a gateway-prefixed id has no default entry to
|
||||
// inherit from, and an unset rate bills that bucket at the input
|
||||
// rate, which would not match the published-rate recompute.
|
||||
if rates.read > 0 {
|
||||
model.CachedInputPer1k = ptr(rates.read) // OpenAI shape
|
||||
model.CacheReadPer1k = ptr(rates.read) // Anthropic / Bedrock shape
|
||||
}
|
||||
if rates.write > 0 {
|
||||
model.CacheCreationPer1k = ptr(rates.write)
|
||||
}
|
||||
req.Models = &[]api.AgentNetworkProviderModel{
|
||||
{Id: modelID, InputPer1k: 0.001, OutputPer1k: 0.002},
|
||||
}
|
||||
req.Models = &[]api.AgentNetworkProviderModel{model}
|
||||
}
|
||||
return req
|
||||
}
|
||||
|
||||
@@ -1,633 +0,0 @@
|
||||
//go:build e2e
|
||||
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/netbirdio/netbird/e2e/harness"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
// The mock vLLM upstream (harness/vllm.go) always answers with this fixed usage
|
||||
// block, so every request drives deterministic token counts regardless of the
|
||||
// model the client asks for. The proxy prices off the REQUEST model, not the
|
||||
// upstream response model, so a made-up model id billed at operator rates lets
|
||||
// these tests assert exact costs without a real vendor key.
|
||||
const (
|
||||
vllmPromptTokens = 11
|
||||
vllmCompletionTokens = 2
|
||||
)
|
||||
|
||||
// pricedEnv is a connected single-provider agent-network deployment pointed at
|
||||
// the mock vLLM upstream, with the proxy and client up and the endpoint resolved
|
||||
// — ready to drive chat. All containers are torn down via t.Cleanup.
|
||||
type pricedEnv struct {
|
||||
providerID string
|
||||
groupID string // source group of the policy; the client peer's auto-group
|
||||
policyID string // policy that authorises (and meters) the requests
|
||||
upstream string // provider upstream URL, needed to re-send on a PUT update
|
||||
endpoint string
|
||||
proxyIP string
|
||||
client *harness.Client
|
||||
proxy *harness.Proxy
|
||||
}
|
||||
|
||||
// provisionPricedProvider brings up the full path for a cost test: a mock vLLM
|
||||
// upstream, a group + reusable setup key, one openai_api provider pointed at the
|
||||
// mock enumerating exactly the given models (with the operator's per-1k prices),
|
||||
// a policy whose token limit switches on usage metering, and a connected proxy +
|
||||
// client. The provider is created with the given models so the router dispatches
|
||||
// them to this provider and the cost meter bills at these rates.
|
||||
//
|
||||
// Passing nil models makes it a gateway-style catch-all: the router claims every
|
||||
// model, and since the synthesizer ships no per-provider-record pricing entry
|
||||
// for a provider that enumerates nothing, the shipped defaults table is the only
|
||||
// thing that can price the request. The policy sets no model guardrail, so the
|
||||
// proxy's per-provider allowlist backstop stays empty and any model routes.
|
||||
func provisionPricedProvider(t *testing.T, ctx context.Context, name string, models []api.AgentNetworkProviderModel) pricedEnv {
|
||||
t.Helper()
|
||||
|
||||
vllm, err := harness.StartVLLM(ctx, srv)
|
||||
require.NoError(t, err, "start mock vLLM upstream")
|
||||
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
|
||||
|
||||
grp, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-price-" + name})
|
||||
require.NoError(t, err, "create group")
|
||||
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grp.Id) })
|
||||
|
||||
ephemeral := false
|
||||
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
|
||||
Name: "e2e-price-" + name + "-client",
|
||||
Type: "reusable",
|
||||
ExpiresIn: 86400,
|
||||
UsageLimit: 0,
|
||||
AutoGroups: []string{grp.Id},
|
||||
Ephemeral: &ephemeral,
|
||||
})
|
||||
require.NoError(t, err, "mint setup key")
|
||||
require.NotEmpty(t, sk.Key, "setup key plaintext")
|
||||
|
||||
// The mock ignores auth, so a dummy key satisfies the "Bearer ${API_KEY}"
|
||||
// template. openai_api is a known catalog provider; the enumerated model id
|
||||
// need NOT be in the catalog — the operator names it and prices it here.
|
||||
dummyKey := "sk-price-e2e"
|
||||
prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: name,
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &dummyKey,
|
||||
Enabled: ptr(true),
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
Models: &models,
|
||||
})
|
||||
require.NoError(t, err, "create provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
|
||||
|
||||
// Uncapped token limit: never blocks the handful of tokens driven here, but
|
||||
// switches on usage metering — the switch that makes consumption rows record.
|
||||
enabled := true
|
||||
pol, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-price-" + name,
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{grp.Id},
|
||||
DestinationProviderIds: []string{prov.Id},
|
||||
Limits: &api.AgentNetworkPolicyLimits{
|
||||
TokenLimit: api.AgentNetworkPolicyTokenLimit{
|
||||
Enabled: true,
|
||||
GroupCap: 10_000_000,
|
||||
UserCap: 10_000_000,
|
||||
WindowSeconds: 60,
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err, "create policy")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), pol.Id) })
|
||||
|
||||
settings, err := srv.GetSettings(ctx)
|
||||
require.NoError(t, err, "read settings")
|
||||
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
|
||||
|
||||
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-price-"+name+"-proxy")
|
||||
require.NoError(t, err, "mint proxy token")
|
||||
px, err := harness.StartProxy(ctx, srv, proxyToken)
|
||||
require.NoError(t, err, "start proxy")
|
||||
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
|
||||
|
||||
cl, err := harness.StartClient(ctx, srv, sk.Key)
|
||||
require.NoError(t, err, "start client")
|
||||
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
|
||||
|
||||
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
|
||||
// Probe first: the GET resolves the endpoint and its first packet wakes the
|
||||
// lazy proxy peer, so WaitProxyPeer then observes it connected.
|
||||
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
|
||||
require.NoError(t, err, "resolve endpoint to proxy IP")
|
||||
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
|
||||
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
|
||||
}
|
||||
|
||||
return pricedEnv{
|
||||
providerID: prov.Id,
|
||||
groupID: grp.Id,
|
||||
policyID: pol.Id,
|
||||
upstream: vllm.URL,
|
||||
endpoint: settings.Endpoint,
|
||||
proxyIP: proxyIP,
|
||||
client: cl,
|
||||
proxy: px,
|
||||
}
|
||||
}
|
||||
|
||||
// chatOnce drives one OpenAI-shaped chat for model through the tunnel, retrying
|
||||
// to absorb first-call tunnel/DNS jitter, and returns the response body.
|
||||
func chatOnce(t *testing.T, ctx context.Context, env pricedEnv, model, sessionID string) string {
|
||||
t.Helper()
|
||||
var code int
|
||||
var body string
|
||||
deadline := time.Now().Add(90 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
c, b, cerr := env.client.Chat(ctx, env.endpoint, env.proxyIP, harness.WireChat, model, "Reply with exactly: pong", sessionID)
|
||||
if cerr == nil {
|
||||
code, body = c, b
|
||||
if code == 200 {
|
||||
break
|
||||
}
|
||||
}
|
||||
time.Sleep(5 * time.Second)
|
||||
}
|
||||
require.Equal(t, 200, code,
|
||||
"chat for %s must return 200; body: %s\n=== proxy logs ===\n%s", model, body, env.proxy.Logs(context.Background()))
|
||||
return body
|
||||
}
|
||||
|
||||
// findAccessLogBySession polls the access-log page for the row carrying sessionID.
|
||||
func findAccessLogBySession(t *testing.T, ctx context.Context, sessionID string) api.AgentNetworkAccessLog {
|
||||
t.Helper()
|
||||
var row api.AgentNetworkAccessLog
|
||||
require.Eventually(t, func() bool {
|
||||
logs, lerr := srv.ListAccessLogs(ctx)
|
||||
if lerr != nil {
|
||||
return false
|
||||
}
|
||||
for _, r := range logs.Data {
|
||||
if r.SessionId != nil && *r.SessionId == sessionID {
|
||||
row = r
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}, 30*time.Second, 2*time.Second, "session id %q must be recorded in an access-log row", sessionID)
|
||||
return row
|
||||
}
|
||||
|
||||
// assertOpenAICostAtRates asserts an access-log row's token counts and every cost
|
||||
// bucket match the mock's fixed usage priced at the given operator rates. The
|
||||
// openai surface has no cache-write bucket and the mock reports no cache tokens,
|
||||
// so the whole cost is input + output; cache costs must be exactly zero.
|
||||
func assertOpenAICostAtRates(t *testing.T, row api.AgentNetworkAccessLog, inRate, outRate float64) {
|
||||
t.Helper()
|
||||
wantInput := float64(vllmPromptTokens) / 1000 * inRate
|
||||
wantOutput := float64(vllmCompletionTokens) / 1000 * outRate
|
||||
wantTotal := wantInput + wantOutput
|
||||
|
||||
model := ""
|
||||
if row.Model != nil {
|
||||
model = *row.Model
|
||||
}
|
||||
t.Logf("[cost] model=%s in=%d out=%d rates in/out=%.4f/%.4f stored input/output/total=$%.6f/$%.6f/$%.6f expected input/output/total=$%.6f/$%.6f/$%.6f",
|
||||
model, row.InputTokens, row.OutputTokens, inRate, outRate,
|
||||
row.InputCostUsd, row.OutputCostUsd, row.CostUsd, wantInput, wantOutput, wantTotal)
|
||||
|
||||
assert.EqualValues(t, vllmPromptTokens, row.InputTokens, "prompt tokens from the mock usage block")
|
||||
assert.EqualValues(t, vllmCompletionTokens, row.OutputTokens, "completion tokens from the mock usage block")
|
||||
assert.InDeltaf(t, wantInput, row.InputCostUsd, 1e-6, "input_cost_usd must be prompt tokens at the operator input rate")
|
||||
assert.InDeltaf(t, wantOutput, row.OutputCostUsd, 1e-6, "output_cost_usd must be completion tokens at the operator output rate")
|
||||
assert.InDeltaf(t, wantTotal, row.CostUsd, 1e-6, "cost_usd must be the sum of the priced buckets")
|
||||
assert.Zerof(t, row.CachedInputCostUsd, "no cache-read tokens, so cached_input_cost_usd must be 0")
|
||||
assert.Zerof(t, row.CacheCreationCostUsd, "openai surface has no cache-write bucket, so cache_creation_cost_usd must be 0")
|
||||
assert.Zerof(t, row.CacheCostUsd, "no cache usage, so cache_cost_usd must be 0")
|
||||
assert.InDeltaf(t, row.InputCostUsd+row.OutputCostUsd, row.CostUsd, 1e-9, "stored buckets must sum to cost_usd")
|
||||
}
|
||||
|
||||
// verifyUsageRowForSession re-checks the persisted usage row for a session
|
||||
// directly in the management sqlite store — the same audit an operator runs on a
|
||||
// production store.db — asserting its cost buckets match the operator rates.
|
||||
func verifyUsageRowForSession(t *testing.T, sessionID string, inRate, outRate float64) {
|
||||
t.Helper()
|
||||
dbPath, err := srv.SnapshotStoreDB(t.TempDir())
|
||||
require.NoError(t, err, "snapshot management sqlite store")
|
||||
|
||||
db, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{})
|
||||
require.NoError(t, err, "open store snapshot")
|
||||
sqlDB, err := db.DB()
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = sqlDB.Close() }()
|
||||
|
||||
var provider, model string
|
||||
var inTok, outTok, cachedTok, cacheCreateTok int64
|
||||
var inCost, cachedInCost, cacheCreateCost, outCost float64
|
||||
row := db.Raw(`SELECT provider, model, input_tokens, output_tokens, cached_input_tokens, cache_creation_tokens,
|
||||
input_cost_usd, cached_input_cost_usd, cache_creation_cost_usd, output_cost_usd
|
||||
FROM agent_network_request_usage WHERE session_id = ? ORDER BY timestamp DESC LIMIT 1`, sessionID).Row()
|
||||
require.NoError(t, row.Scan(&provider, &model, &inTok, &outTok, &cachedTok, &cacheCreateTok,
|
||||
&inCost, &cachedInCost, &cacheCreateCost, &outCost),
|
||||
"a usage row must exist for session %q", sessionID)
|
||||
|
||||
wantInput := float64(inTok) / 1000 * inRate
|
||||
wantOutput := float64(outTok) / 1000 * outRate
|
||||
t.Logf("[sql] session=%s %s/%s in=%d out=%d stored input/cached/create/output=$%.6f/$%.6f/$%.6f/$%.6f",
|
||||
sessionID, provider, model, inTok, outTok, inCost, cachedInCost, cacheCreateCost, outCost)
|
||||
assert.EqualValues(t, vllmPromptTokens, inTok, "usage row prompt tokens")
|
||||
assert.EqualValues(t, vllmCompletionTokens, outTok, "usage row completion tokens")
|
||||
assert.InDeltaf(t, wantInput, inCost, 1e-6, "usage input_cost_usd must be prompt tokens at the operator input rate")
|
||||
assert.InDeltaf(t, wantOutput, outCost, 1e-6, "usage output_cost_usd must be completion tokens at the operator output rate")
|
||||
assert.Zerof(t, cachedInCost, "usage cached_input_cost_usd must be 0 (no cache usage)")
|
||||
assert.Zerof(t, cacheCreateCost, "usage cache_creation_cost_usd must be 0 (no cache usage)")
|
||||
}
|
||||
|
||||
// TestCustomModelPricing proves an operator can serve a model that is NOT in
|
||||
// NetBird's compiled catalog, at prices they type themselves, and that those
|
||||
// operator prices drive the recorded cost end to end — access log AND usage
|
||||
// ledger. The provider enumerates one made-up model id at deliberately odd rates
|
||||
// (no default entry could supply them), the client requests it, and every cost
|
||||
// bucket must equal the mock's fixed token counts multiplied by those rates.
|
||||
func TestCustomModelPricing(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
const (
|
||||
customModel = "e2e-custom-model" // absent from the compiled catalog
|
||||
inRate = 0.037 // odd rates so a stray default can't match
|
||||
outRate = 0.089
|
||||
)
|
||||
|
||||
env := provisionPricedProvider(t, ctx, "custommodel", []api.AgentNetworkProviderModel{
|
||||
{Id: customModel, InputPer1k: inRate, OutputPer1k: outRate},
|
||||
})
|
||||
|
||||
sessionID := "e2e-session-custommodel"
|
||||
body := chatOnce(t, ctx, env, customModel, sessionID)
|
||||
require.Contains(t, body, "chat.completion", "body should be an OpenAI-compatible completion; got: %s", body)
|
||||
|
||||
row := findAccessLogBySession(t, ctx, sessionID)
|
||||
require.NotNil(t, row.Model, "access-log row must carry the requested model")
|
||||
assert.Equal(t, customModel, *row.Model, "the row must be stamped with the requested (custom) model, not the mock's response model")
|
||||
assertOpenAICostAtRates(t, row, inRate, outRate)
|
||||
|
||||
// Metering: the uncapped token limit switches on usage recording, so the
|
||||
// request must surface as a consumption row with positive tokens and cost.
|
||||
require.Eventually(t, func() bool {
|
||||
rows, lerr := srv.ListConsumption(ctx)
|
||||
if lerr != nil {
|
||||
return false
|
||||
}
|
||||
for _, r := range rows {
|
||||
if r.TokensInput > 0 && r.TokensOutput > 0 && r.CostUsd > 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}, 60*time.Second, 3*time.Second, "custom-model usage must be metered into a consumption row with positive cost")
|
||||
|
||||
// Final raw-SQL audit: bypass the API and re-verify the persisted usage row.
|
||||
verifyUsageRowForSession(t, sessionID, inRate, outRate)
|
||||
}
|
||||
|
||||
// TestPriceChangeUpdatesRecordedCost proves that changing a provider's model
|
||||
// price is reflected in the cost recorded for subsequent requests — in both the
|
||||
// access log and the usage ledger — while requests already priced at the old
|
||||
// rate keep their original cost. The update propagates to the connected proxy
|
||||
// live (a mapping push rebuilds the cost_meter chain with the new table), so no
|
||||
// reconnect or restart is needed; the test polls a fresh request until the new
|
||||
// rate lands.
|
||||
func TestPriceChangeUpdatesRecordedCost(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
const (
|
||||
customModel = "e2e-repriced-model"
|
||||
inRateA = 0.010
|
||||
outRateA = 0.020
|
||||
inRateB = 0.050 // 5x / 4x the original, so a repriced row is unmistakable
|
||||
outRateB = 0.080
|
||||
)
|
||||
|
||||
env := provisionPricedProvider(t, ctx, "reprice", []api.AgentNetworkProviderModel{
|
||||
{Id: customModel, InputPer1k: inRateA, OutputPer1k: outRateA},
|
||||
})
|
||||
|
||||
// Phase 1 — request priced at the original rate A.
|
||||
sessionA := "e2e-session-reprice-a"
|
||||
chatOnce(t, ctx, env, customModel, sessionA)
|
||||
rowA := findAccessLogBySession(t, ctx, sessionA)
|
||||
assertOpenAICostAtRates(t, rowA, inRateA, outRateA)
|
||||
verifyUsageRowForSession(t, sessionA, inRateA, outRateA)
|
||||
|
||||
// Change the model's price. The API key is omitted so the stored one is kept;
|
||||
// the models array is re-sent with the new rates (PUT replaces the list).
|
||||
// This reconciles synchronously and pushes a fresh cost_meter table to the
|
||||
// already-connected proxy — no reconnect.
|
||||
_, err := srv.UpdateProvider(ctx, env.providerID, api.AgentNetworkProviderRequest{
|
||||
Name: "reprice",
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: env.upstream,
|
||||
Enabled: ptr(true),
|
||||
Models: &[]api.AgentNetworkProviderModel{
|
||||
{Id: customModel, InputPer1k: inRateB, OutputPer1k: outRateB},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err, "update provider price")
|
||||
|
||||
// Phase 2 — the push + chain rebuild is async, so drive fresh requests (each
|
||||
// under its own session) until one is priced at the new rate B. Each iteration
|
||||
// fires one request and waits for that session's row to be ingested before
|
||||
// reading its cost, so an un-ingested row is never mistaken for "still rate A".
|
||||
// The expected new input cost is unmistakably higher than rate A, so a
|
||||
// lingering old-rate row can't satisfy the check.
|
||||
wantInputB := float64(vllmPromptTokens) / 1000 * inRateB
|
||||
var repriced api.AgentNetworkAccessLog
|
||||
var lastSession string
|
||||
deadline := time.Now().Add(90 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
lastSession = fmt.Sprintf("e2e-session-reprice-b-%d", time.Now().UnixNano())
|
||||
code, _, cerr := env.client.Chat(ctx, env.endpoint, env.proxyIP, harness.WireChat, customModel, "Reply with exactly: pong", lastSession)
|
||||
if cerr != nil || code != 200 {
|
||||
time.Sleep(5 * time.Second)
|
||||
continue
|
||||
}
|
||||
row := findAccessLogBySession(t, ctx, lastSession)
|
||||
if inDelta(row.InputCostUsd, wantInputB, 1e-6) {
|
||||
repriced = row
|
||||
break
|
||||
}
|
||||
// Still priced at the old rate — the push hasn't landed yet; retry.
|
||||
time.Sleep(5 * time.Second)
|
||||
}
|
||||
require.NotEmpty(t, repriced.Id, "a request after the price change must be priced at the new rate B; last input_cost_usd=$%.6f, wanted $%.6f\n=== proxy logs ===\n%s",
|
||||
repriced.InputCostUsd, wantInputB, env.proxy.Logs(context.Background()))
|
||||
|
||||
assertOpenAICostAtRates(t, repriced, inRateB, outRateB)
|
||||
verifyUsageRowForSession(t, lastSession, inRateB, outRateB)
|
||||
|
||||
// The original request keeps its original cost: repricing is not retroactive.
|
||||
rowAStill := findAccessLogBySession(t, ctx, sessionA)
|
||||
assertOpenAICostAtRates(t, rowAStill, inRateA, outRateA)
|
||||
verifyUsageRowForSession(t, sessionA, inRateA, outRateA)
|
||||
}
|
||||
|
||||
// TestPricingDefaultsFileDrivesCost proves the operator-supplied pricing
|
||||
// defaults file is what the proxy bills with. The harness configures
|
||||
// server.agentNetwork.pricingDefaultsFile as a BARE FILENAME and writes that
|
||||
// file into the bind-mounted datadir (see harness.PricingDefaultsFileName), so a
|
||||
// pass exercises the whole chain: combined yaml → ToManagementConfig →
|
||||
// pricing.LoadFile (relative path resolved against datadir) → DefaultTable →
|
||||
// the synthesizer's cost_meter defaults payload → the proxy's lookup.
|
||||
//
|
||||
// The provider enumerates NO models, so it is a catch-all route with no
|
||||
// per-provider-record pricing entry at all — the only rates that can price the
|
||||
// request are the shipped defaults. The model is a real catalog model whose
|
||||
// built-in rates the file replaces with deliberately odd values, so billing at
|
||||
// the compiled-in rates (i.e. the file never loaded) fails the assertions.
|
||||
func TestPricingDefaultsFileDrivesCost(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
// nil models: a gateway-style provider claiming every model. The synthesizer
|
||||
// ships no per-record entry for it, so the defaults table is its price list.
|
||||
env := provisionPricedProvider(t, ctx, "defaultsfile", nil)
|
||||
|
||||
sessionID := "e2e-session-defaultsfile"
|
||||
body := chatOnce(t, ctx, env, harness.PricedDefaultModel, sessionID)
|
||||
require.Contains(t, body, "chat.completion", "body should be an OpenAI-compatible completion; got: %s", body)
|
||||
|
||||
row := findAccessLogBySession(t, ctx, sessionID)
|
||||
require.NotNil(t, row.Model, "access-log row must carry the requested model")
|
||||
assert.Equal(t, harness.PricedDefaultModel, *row.Model, "the row must be stamped with the requested model")
|
||||
|
||||
// The file's rates, not the compiled-in catalog rates for this model.
|
||||
assertOpenAICostAtRates(t, row, harness.PricedDefaultInputPer1k, harness.PricedDefaultOutputPer1k)
|
||||
verifyUsageRowForSession(t, sessionID, harness.PricedDefaultInputPer1k, harness.PricedDefaultOutputPer1k)
|
||||
}
|
||||
|
||||
// TestPricingDefaultsFileLeavesOtherModelsAlone proves the defaults file merges
|
||||
// per entry rather than replacing the whole table: the file names exactly one
|
||||
// model, so a DIFFERENT catalog model must still bill at its compiled-in rates.
|
||||
// Without this, a file that shipped as a wholesale replacement would silently
|
||||
// zero-cost every model the operator didn't list, and TestPricingDefaultsFile-
|
||||
// DrivesCost alone would not notice.
|
||||
func TestPricingDefaultsFileLeavesOtherModelsAlone(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
// gpt-4o-mini is a catalog model the pricing file does NOT mention, so it must
|
||||
// keep its built-in rates. Pinned here independently of the catalog source so
|
||||
// a rate change in either place surfaces as a failure to reconcile rather
|
||||
// than passing silently.
|
||||
const (
|
||||
untouchedModel = "gpt-4o-mini"
|
||||
builtinInRate = 0.00015
|
||||
builtinOutRate = 0.0006
|
||||
)
|
||||
|
||||
env := provisionPricedProvider(t, ctx, "defaultsfileother", nil)
|
||||
|
||||
sessionID := "e2e-session-defaultsfile-other"
|
||||
chatOnce(t, ctx, env, untouchedModel, sessionID)
|
||||
|
||||
row := findAccessLogBySession(t, ctx, sessionID)
|
||||
assertOpenAICostAtRates(t, row, builtinInRate, builtinOutRate)
|
||||
verifyUsageRowForSession(t, sessionID, builtinInRate, builtinOutRate)
|
||||
}
|
||||
|
||||
// TestCustomModelAccessLogAttribution proves a custom (non-catalog) model is
|
||||
// handled correctly in the ACCESS LOG, not just in the cost columns. The other
|
||||
// tests here assert money; this one asserts the row's identity and attribution
|
||||
// dimensions — the columns the dashboard filters, groups and drills down on.
|
||||
//
|
||||
// A custom model id is the interesting case precisely because nothing in
|
||||
// NetBird's catalog describes it. Its provider vendor, parser surface, cost
|
||||
// buckets, and dashboard filterability all have to come from the operator's
|
||||
// provider record rather than from a compiled-in entry. So this checks:
|
||||
//
|
||||
// - the row is stamped with the REQUESTED model id verbatim, not the mock
|
||||
// upstream's response model (Qwen/Qwen2.5-0.5B-Instruct) and not a
|
||||
// normalized or catalog-substituted id;
|
||||
// - provider is the vendor SURFACE ("openai", from the catalog entry's
|
||||
// ParserID) — a custom model does not change which wire shape was spoken;
|
||||
// - resolved_provider_id / selected_policy_id / group_ids attribute the row to
|
||||
// the operator's provider record, the authorising policy, and the caller's
|
||||
// group, so spend on a custom model is attributable;
|
||||
// - decision is "allow" with no deny reason, and the request dimensions
|
||||
// (status 200, POST, the OpenAI chat path, non-stream, source IP, duration)
|
||||
// are recorded;
|
||||
// - management's SERVER-SIDE model filter finds the row by its custom id, so
|
||||
// the model column is genuinely indexed and queryable rather than merely
|
||||
// stored;
|
||||
// - prompt/completion capture stays empty, since prompt collection is off by
|
||||
// default and a custom model must not bypass that gate.
|
||||
func TestCustomModelAccessLogAttribution(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
// A model id no catalog entry carries, at odd rates so its cost cannot come
|
||||
// from anywhere but the provider record.
|
||||
const (
|
||||
customModel = "e2e-attribution-model-v9"
|
||||
inRate = 0.0271
|
||||
outRate = 0.0913
|
||||
)
|
||||
|
||||
env := provisionPricedProvider(t, ctx, "attribution", []api.AgentNetworkProviderModel{
|
||||
{Id: customModel, InputPer1k: inRate, OutputPer1k: outRate},
|
||||
})
|
||||
|
||||
sessionID := "e2e-session-attribution"
|
||||
body := chatOnce(t, ctx, env, customModel, sessionID)
|
||||
require.Contains(t, body, "chat.completion", "body should be an OpenAI-compatible completion; got: %s", body)
|
||||
|
||||
row := findAccessLogBySession(t, ctx, sessionID)
|
||||
|
||||
// Identity: the requested model verbatim. The mock answers with its own
|
||||
// served model id, so a row carrying that instead means the log is sourced
|
||||
// from the response body rather than the parsed request.
|
||||
require.NotNil(t, row.Model, "access-log row must carry the requested model")
|
||||
assert.Equal(t, customModel, *row.Model,
|
||||
"the row must be stamped with the requested custom model id verbatim, not the mock upstream's response model (%s)", harness.VLLMModel)
|
||||
|
||||
// Surface: a custom model id does not change the wire shape that was spoken.
|
||||
// provider is the vendor surface from the catalog entry's parser, which is
|
||||
// also the key the cost meter's cache formula switches on.
|
||||
require.NotNil(t, row.Provider, "access-log row must carry the vendor surface")
|
||||
assert.Equal(t, "openai", *row.Provider,
|
||||
"openai_api's parser surface is openai, regardless of how exotic the model id is")
|
||||
|
||||
// Attribution: which provider record served it, which policy authorised it,
|
||||
// and which group the authorisation came through. Without these, spend on a
|
||||
// custom model can be seen but not attributed.
|
||||
require.NotNil(t, row.ResolvedProviderId, "row must name the provider record that served the request")
|
||||
assert.Equal(t, env.providerID, *row.ResolvedProviderId,
|
||||
"the router stamps the operator's provider record id; a custom model must attribute to the record that enumerated it")
|
||||
require.NotNil(t, row.SelectedPolicyId, "row must name the policy that authorised the request")
|
||||
assert.Equal(t, env.policyID, *row.SelectedPolicyId,
|
||||
"the policy carrying the token limit is the one that paid for the request")
|
||||
require.NotNil(t, row.GroupIds, "row must carry the authorising group ids")
|
||||
assert.Contains(t, *row.GroupIds, env.groupID,
|
||||
"the caller's group is the policy's source group, so it must be the authorising group")
|
||||
|
||||
// Decision + request dimensions.
|
||||
require.NotNil(t, row.Decision, "row must carry the policy decision")
|
||||
assert.Equal(t, "allow", *row.Decision, "the uncapped policy allows this request")
|
||||
if row.DenyReason != nil {
|
||||
assert.Empty(t, *row.DenyReason, "an allowed request must carry no deny reason")
|
||||
}
|
||||
assert.Equal(t, 200, row.StatusCode, "the mock upstream answers 200")
|
||||
if row.Method != nil {
|
||||
assert.Equal(t, "POST", *row.Method, "a chat completion is a POST")
|
||||
}
|
||||
require.NotNil(t, row.Path, "row must record the request path")
|
||||
assert.Equal(t, "/v1/chat/completions", *row.Path,
|
||||
"the OpenAI chat path the client called, as seen by the proxy")
|
||||
require.NotNil(t, row.Host, "row must record the host the client addressed")
|
||||
assert.Equal(t, env.endpoint, *row.Host, "the agent-network endpoint the client resolved")
|
||||
if row.Stream != nil {
|
||||
assert.False(t, *row.Stream, "the harness sends a non-streaming request")
|
||||
}
|
||||
require.NotNil(t, row.SourceIp, "row must record the caller's tunnel IP")
|
||||
assert.NotEmpty(t, *row.SourceIp, "the request arrived over the tunnel, so a source IP is known")
|
||||
|
||||
// Tokens and cost, so the attribution above is anchored to a real priced row
|
||||
// rather than an empty shell that happens to carry the right ids.
|
||||
assertOpenAICostAtRates(t, row, inRate, outRate)
|
||||
assert.EqualValues(t, vllmPromptTokens+vllmCompletionTokens, row.TotalTokens,
|
||||
"total_tokens is the mock's reported total")
|
||||
|
||||
// Prompt capture is off by default (account master switch), and a custom
|
||||
// model must not bypass that gate.
|
||||
if row.RequestPrompt != nil {
|
||||
assert.Empty(t, *row.RequestPrompt, "prompt collection is off by default, so no prompt may be stored")
|
||||
}
|
||||
if row.ResponseCompletion != nil {
|
||||
assert.Empty(t, *row.ResponseCompletion, "prompt collection is off by default, so no completion may be stored")
|
||||
}
|
||||
|
||||
// Queryability: management's SERVER-SIDE model filter must find the row by
|
||||
// its custom id. findAccessLogBySession above scans a page client-side, so
|
||||
// this is the check that the model column is actually indexed and filterable
|
||||
// — the dashboard's per-model drill-down on a custom model depends on it.
|
||||
filtered, err := srv.ListAccessLogsFiltered(ctx, url.Values{"model": []string{customModel}})
|
||||
require.NoError(t, err, "filter access logs by the custom model id")
|
||||
require.Positive(t, filtered.TotalRecords, "the custom model must be findable via the server-side model filter")
|
||||
foundSession := false
|
||||
for _, r := range filtered.Data {
|
||||
require.NotNil(t, r.Model, "filtered row must carry a model")
|
||||
assert.Equal(t, customModel, *r.Model, "the model filter must not return rows for other models")
|
||||
if r.SessionId != nil && *r.SessionId == sessionID {
|
||||
foundSession = true
|
||||
}
|
||||
}
|
||||
assert.True(t, foundSession, "the filtered page must include this test's request")
|
||||
|
||||
// Final raw-SQL audit of the parallel usage row: the ledger must carry the
|
||||
// same custom model, surface, and provider-record attribution as the log.
|
||||
verifyUsageAttributionForSession(t, sessionID, customModel, "openai", env.providerID, env.groupID)
|
||||
}
|
||||
|
||||
// verifyUsageAttributionForSession checks the usage ledger's attribution columns
|
||||
// for a session directly in the management sqlite store — including the group
|
||||
// child row, which the API renders but which only exists if the proxy's
|
||||
// authorising-group CSV was parsed into normalised rows. The usage table is
|
||||
// written unconditionally (independent of the log-collection toggle), so this is
|
||||
// the record that must attribute spend even for accounts with logs off.
|
||||
func verifyUsageAttributionForSession(t *testing.T, sessionID, wantModel, wantProvider, wantProviderID, wantGroupID string) {
|
||||
t.Helper()
|
||||
dbPath, err := srv.SnapshotStoreDB(t.TempDir())
|
||||
require.NoError(t, err, "snapshot management sqlite store")
|
||||
|
||||
db, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{})
|
||||
require.NoError(t, err, "open store snapshot")
|
||||
sqlDB, err := db.DB()
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = sqlDB.Close() }()
|
||||
|
||||
var id, provider, model, resolvedProviderID, userID string
|
||||
require.NoError(t, db.Raw(
|
||||
`SELECT id, provider, model, resolved_provider_id, user_id
|
||||
FROM agent_network_request_usage WHERE session_id = ? ORDER BY timestamp DESC LIMIT 1`, sessionID).
|
||||
Row().Scan(&id, &provider, &model, &resolvedProviderID, &userID),
|
||||
"a usage row must exist for session %q", sessionID)
|
||||
|
||||
t.Logf("[sql] usage attribution session=%s id=%s provider=%s model=%s resolved_provider_id=%s user_id=%s",
|
||||
sessionID, id, provider, model, resolvedProviderID, userID)
|
||||
assert.Equal(t, wantModel, model, "usage row must carry the requested custom model")
|
||||
assert.Equal(t, wantProvider, provider, "usage row must carry the vendor surface")
|
||||
assert.Equal(t, wantProviderID, resolvedProviderID, "usage row must attribute to the operator's provider record")
|
||||
assert.NotEmpty(t, userID, "the tunnel peer resolves to a principal, so the usage row must be attributable to it")
|
||||
|
||||
// The authorising group lands in the normalised child table, which is what
|
||||
// the usage overview joins on to break spend down by group.
|
||||
var groupIDs []string
|
||||
require.NoError(t, db.Raw(
|
||||
`SELECT group_id FROM agent_network_request_usage_group WHERE usage_id = ?`, id).
|
||||
Scan(&groupIDs).Error, "read usage group child rows")
|
||||
assert.Contains(t, groupIDs, wantGroupID,
|
||||
"the authorising group must be normalised into a usage_group row so spend can be grouped by it")
|
||||
}
|
||||
|
||||
// inDelta reports whether a and b are within tol of each other.
|
||||
func inDelta(a, b, tol float64) bool {
|
||||
d := a - b
|
||||
if d < 0 {
|
||||
d = -d
|
||||
}
|
||||
return d <= tol
|
||||
}
|
||||
@@ -9,7 +9,6 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
@@ -75,13 +74,6 @@ func (c *Combined) DeleteProvider(ctx context.Context, id string) error {
|
||||
return anDelete(ctx, c, "/api/agent-network/providers/"+id)
|
||||
}
|
||||
|
||||
// UpdateProvider replaces a provider by id (PUT). The API key may be omitted on
|
||||
// the request to keep the stored one; Models replaces the enumerated list, so
|
||||
// this is the path a test uses to change a model's price mid-run.
|
||||
func (c *Combined) UpdateProvider(ctx context.Context, id string, req api.AgentNetworkProviderRequest) (api.AgentNetworkProvider, error) {
|
||||
return anRequest[api.AgentNetworkProvider](ctx, c, http.MethodPut, "/api/agent-network/providers/"+id, req)
|
||||
}
|
||||
|
||||
// SetProviderEnabled toggles a provider's enabled flag, preserving its other
|
||||
// fields (the API key is omitted, which keeps the stored one). Used to run one
|
||||
// provider at a time so model→provider routing is unambiguous.
|
||||
@@ -147,16 +139,3 @@ func (c *Combined) ListConsumption(ctx context.Context) ([]api.AgentNetworkConsu
|
||||
func (c *Combined) ListAccessLogs(ctx context.Context) (api.AgentNetworkAccessLogsResponse, error) {
|
||||
return anRequest[api.AgentNetworkAccessLogsResponse](ctx, c, http.MethodGet, "/api/agent-network/access-logs", nil)
|
||||
}
|
||||
|
||||
// ListAccessLogsFiltered returns the access-log page narrowed by the given
|
||||
// query parameters (e.g. model=..., session_id=..., provider_id=...). This
|
||||
// exercises management's server-side filtering rather than filtering client
|
||||
// side, so a row that is ingested but not indexed under the filtered column
|
||||
// surfaces as an empty page.
|
||||
func (c *Combined) ListAccessLogsFiltered(ctx context.Context, query url.Values) (api.AgentNetworkAccessLogsResponse, error) {
|
||||
path := "/api/agent-network/access-logs"
|
||||
if encoded := query.Encode(); encoded != "" {
|
||||
path += "?" + encoded
|
||||
}
|
||||
return anRequest[api.AgentNetworkAccessLogsResponse](ctx, c, http.MethodGet, path, nil)
|
||||
}
|
||||
|
||||
@@ -93,19 +93,10 @@ func StartCombined(ctx context.Context) (*Combined, error) {
|
||||
_ = net.Remove(ctx)
|
||||
return nil, fmt.Errorf("write combined config: %w", err)
|
||||
}
|
||||
dataDir := filepath.Join(workDir, "data")
|
||||
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Join(workDir, "data"), 0o755); err != nil {
|
||||
_ = net.Remove(ctx)
|
||||
return nil, fmt.Errorf("create datadir: %w", err)
|
||||
}
|
||||
// The config's agentNetwork.pricingDefaultsFile is a bare filename, so the
|
||||
// server resolves it against the datadir; write it there. It is an explicitly
|
||||
// configured path, so a failure to load fails the server's startup — which
|
||||
// surfaces here as the /api/instance readiness wait timing out.
|
||||
if err := os.WriteFile(filepath.Join(dataDir, PricingDefaultsFileName), []byte(pricingDefaultsYAML), 0o644); err != nil { //nolint:gosec // non-secret config, bind-mounted and read by the container
|
||||
_ = net.Remove(ctx)
|
||||
return nil, fmt.Errorf("write pricing defaults: %w", err)
|
||||
}
|
||||
|
||||
req := testcontainers.ContainerRequest{
|
||||
Image: combinedImage,
|
||||
|
||||
@@ -8,13 +8,6 @@ package harness
|
||||
// embedded IdP, local signal/relay/STUN, and a sqlite store under the mounted
|
||||
// data dir. exposedAddress is the address peers use to reach this container; it
|
||||
// is overridden per-run so the value matches the container's network alias.
|
||||
//
|
||||
// pricingDefaultsFile is deliberately a BARE FILENAME, not an absolute path: it
|
||||
// must resolve against dataDir (→ /nb/data/<name>), which is the resolution rule
|
||||
// the combined server applies. It is also an EXPLICITLY configured path, so the
|
||||
// server is required to load it — a broken path or malformed file fails startup
|
||||
// rather than silently falling back to the compiled-in rates, and TestMain then
|
||||
// fails with the container logs.
|
||||
const combinedConfigYAML = `server:
|
||||
listenAddress: ":8080"
|
||||
exposedAddress: "%s"
|
||||
@@ -30,36 +23,4 @@ const combinedConfigYAML = `server:
|
||||
issuer: "%s"
|
||||
store:
|
||||
engine: "sqlite"
|
||||
agentNetwork:
|
||||
pricingDefaultsFile: "` + PricingDefaultsFileName + `"
|
||||
`
|
||||
|
||||
const (
|
||||
// PricingDefaultsFileName is the basename of the operator-supplied LLM
|
||||
// pricing defaults file the combined server is configured to load. Written
|
||||
// into the bind-mounted datadir by StartCombined.
|
||||
PricingDefaultsFileName = "e2e_llm_pricing.yaml"
|
||||
|
||||
// PricedDefaultModel is a real catalog model (openai surface) whose rates the
|
||||
// defaults file below REPLACES. Tests drive it against the mock vLLM upstream
|
||||
// and assert the file's rates were billed, which is only true if the file
|
||||
// travelled: config → LoadFile → DefaultTable → synthesizer → the proxy's
|
||||
// cost_meter defaults table.
|
||||
PricedDefaultModel = "gpt-4.1-mini"
|
||||
// PricedDefaultInputPer1k / PricedDefaultOutputPer1k are deliberately odd
|
||||
// values that no compiled-in catalog entry carries (gpt-4.1-mini ships as
|
||||
// 0.0004 / 0.0016), so a test asserting them cannot pass on the built-in
|
||||
// table.
|
||||
PricedDefaultInputPer1k = 0.0123
|
||||
PricedDefaultOutputPer1k = 0.0456
|
||||
)
|
||||
|
||||
// pricingDefaultsYAML is the operator-supplied pricing defaults file. Its schema
|
||||
// is surface -> model -> per-1k rates. Entries replace the compiled-in entry for
|
||||
// the same surface+model whole; every other model keeps its built-in rates, so
|
||||
// this file overriding one model must not disturb the rest of the table.
|
||||
const pricingDefaultsYAML = `openai:
|
||||
gpt-4.1-mini:
|
||||
input_per_1k: 0.0123
|
||||
output_per_1k: 0.0456
|
||||
`
|
||||
|
||||
@@ -23,7 +23,6 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
|
||||
"github.com/netbirdio/netbird/formatter/hook"
|
||||
agentnetworkpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
"github.com/netbirdio/netbird/management/internals/server"
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
nbdomain "github.com/netbirdio/netbird/shared/management/domain"
|
||||
@@ -113,29 +112,6 @@ var (
|
||||
mgmtSingleAccModeDomain = ""
|
||||
}
|
||||
|
||||
// Load the management-side LLM pricing defaults file: an
|
||||
// explicitly configured path is required to load (a typo must
|
||||
// fail startup — the operator believes those rates are live);
|
||||
// otherwise <datadir>/defaults_llm_pricing.yaml is probed and
|
||||
// may be absent (compiled-in defaults serve). A relative path
|
||||
// is resolved against the datadir so a bare filename lands
|
||||
// alongside the store. Either way the path stays watched: the
|
||||
// reloader picks up edits — and the file appearing later —
|
||||
// without a restart.
|
||||
pricingPath := config.AgentNetwork.PricingDefaultsFile
|
||||
pricingRequired := pricingPath != ""
|
||||
if !pricingRequired {
|
||||
pricingPath = agentnetworkpricing.DefaultFileName
|
||||
}
|
||||
if !filepath.IsAbs(pricingPath) {
|
||||
pricingPath = filepath.Join(config.Datadir, pricingPath)
|
||||
}
|
||||
log.Infof("loading agent-network pricing defaults from %s (required: %v)", pricingPath, pricingRequired)
|
||||
if err := agentnetworkpricing.LoadFile(pricingPath, pricingRequired); err != nil {
|
||||
return fmt.Errorf("load agent-network pricing defaults: %v", err)
|
||||
}
|
||||
agentnetworkpricing.StartReloader(ctx, agentnetworkpricing.ReloadInterval)
|
||||
|
||||
srv := newServer(&server.Config{
|
||||
NbConfig: config,
|
||||
DNSDomain: dnsDomain,
|
||||
|
||||
@@ -7,28 +7,12 @@ package catalog
|
||||
import "github.com/netbirdio/netbird/shared/management/http/api"
|
||||
|
||||
// Model is the in-memory representation of a catalog model.
|
||||
//
|
||||
// The three cache rates mirror the proxy cost meter's Entry semantics
|
||||
// (USD per 1k tokens; 0 = no rate configured, that bucket bills at
|
||||
// InputPer1k):
|
||||
// - CachedInputPer1k: OpenAI-shape rate for cached prompt tokens
|
||||
// (a SUBSET of input tokens). Typically 0.1-0.5x input.
|
||||
// - CacheReadPer1k / CacheCreationPer1k: Anthropic-shape rates for
|
||||
// the two ADDITIVE prompt-cache buckets. Typically 0.1x / 1.25x
|
||||
// input.
|
||||
//
|
||||
// The catalog is the single default-pricing source: the agentnetwork
|
||||
// pricing package folds these models into per-surface tables that the
|
||||
// synthesizer ships to the proxy's cost_meter.
|
||||
type Model struct {
|
||||
ID string
|
||||
Label string
|
||||
InputPer1k float64
|
||||
OutputPer1k float64
|
||||
CachedInputPer1k float64
|
||||
CacheReadPer1k float64
|
||||
CacheCreationPer1k float64
|
||||
ContextWindow int
|
||||
ID string
|
||||
Label string
|
||||
InputPer1k float64
|
||||
OutputPer1k float64
|
||||
ContextWindow int
|
||||
}
|
||||
|
||||
// ProviderKind groups catalog entries for UI presentation. The split
|
||||
@@ -81,17 +65,6 @@ type Provider struct {
|
||||
// surface — the proxy middleware then falls back to URL sniffing
|
||||
// or skips request-side enrichment.
|
||||
ParserID string
|
||||
// PricingSurfaces names the cost-meter pricing surfaces this
|
||||
// provider's Models are priced under ("openai", "anthropic",
|
||||
// "bedrock" — the llm.Parser surface the request parser stamps as
|
||||
// llm.provider at billing time). NOT derivable from ParserID:
|
||||
// bedrock_api and vertex_ai_api leave ParserID empty (URL-sniffed)
|
||||
// yet price under "bedrock" / "anthropic", and kimi_api serves two
|
||||
// body shapes so it prices under both. Nil for gateway/custom
|
||||
// entries, which declare no models. Same (surface, model) pair
|
||||
// contributed by two providers must carry identical rates — the
|
||||
// pricing package's tests enforce that.
|
||||
PricingSurfaces []string
|
||||
// IdentityInjection, when non-nil, instructs the proxy to stamp
|
||||
// the caller's NetBird identity onto upstream requests under the
|
||||
// configured header names. Used for gateways like LiteLLM that
|
||||
@@ -246,7 +219,6 @@ var providers = []Provider{
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#10A37F",
|
||||
ParserID: "openai",
|
||||
PricingSurfaces: []string{"openai"},
|
||||
// Pricing + context windows cross-checked against LiteLLM's
|
||||
// model_prices_and_context_window.json. Notable corrections from
|
||||
// earlier values: o4-mini repriced from $4/$16 to $1.10/$4.40
|
||||
@@ -254,20 +226,20 @@ var providers = []Provider{
|
||||
// family context windows split between 1.05M for full-size
|
||||
// models and 272K for mini/nano/codex variants.
|
||||
Models: []Model{
|
||||
{ID: "gpt-5.5", Label: "GPT-5.5", InputPer1k: 0.005, OutputPer1k: 0.030, CachedInputPer1k: 0.0005, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.5-pro", Label: "GPT-5.5 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, CachedInputPer1k: 0.003, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4", Label: "GPT-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015, CachedInputPer1k: 0.00025, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-pro", Label: "GPT-5.4 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, CachedInputPer1k: 0.003, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini", InputPer1k: 0.00075, OutputPer1k: 0.0045, CachedInputPer1k: 0.000075, ContextWindow: 272000},
|
||||
{ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano", InputPer1k: 0.0002, OutputPer1k: 0.00125, CachedInputPer1k: 0.00002, ContextWindow: 272000},
|
||||
{ID: "gpt-5.3-codex", Label: "GPT-5.3 Codex", InputPer1k: 0.00175, OutputPer1k: 0.014, CachedInputPer1k: 0.000175, ContextWindow: 272000},
|
||||
{ID: "gpt-5.3-chat-latest", Label: "GPT-5.3 Chat", InputPer1k: 0.00175, OutputPer1k: 0.014, CachedInputPer1k: 0.000175, ContextWindow: 128000},
|
||||
{ID: "o4-mini", Label: "o4-mini", InputPer1k: 0.0011, OutputPer1k: 0.0044, CachedInputPer1k: 0.000275, ContextWindow: 200000},
|
||||
{ID: "gpt-4.1", Label: "GPT-4.1", InputPer1k: 0.002, OutputPer1k: 0.008, CachedInputPer1k: 0.0005, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-mini", Label: "GPT-4.1 mini", InputPer1k: 0.0004, OutputPer1k: 0.0016, CachedInputPer1k: 0.0001, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-nano", Label: "GPT-4.1 nano", InputPer1k: 0.0001, OutputPer1k: 0.0004, CachedInputPer1k: 0.000025, ContextWindow: 1047576},
|
||||
{ID: "gpt-4o", Label: "GPT-4o", InputPer1k: 0.0025, OutputPer1k: 0.010, CachedInputPer1k: 0.00125, ContextWindow: 128000},
|
||||
{ID: "gpt-4o-mini", Label: "GPT-4o mini", InputPer1k: 0.00015, OutputPer1k: 0.0006, CachedInputPer1k: 0.000075, ContextWindow: 128000},
|
||||
{ID: "gpt-5.5", Label: "GPT-5.5", InputPer1k: 0.005, OutputPer1k: 0.030, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.5-pro", Label: "GPT-5.5 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4", Label: "GPT-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-pro", Label: "GPT-5.4 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini", InputPer1k: 0.00075, OutputPer1k: 0.0045, ContextWindow: 272000},
|
||||
{ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano", InputPer1k: 0.0002, OutputPer1k: 0.00125, ContextWindow: 272000},
|
||||
{ID: "gpt-5.3-codex", Label: "GPT-5.3 Codex", InputPer1k: 0.00175, OutputPer1k: 0.014, ContextWindow: 272000},
|
||||
{ID: "gpt-5.3-chat-latest", Label: "GPT-5.3 Chat", InputPer1k: 0.00175, OutputPer1k: 0.014, ContextWindow: 128000},
|
||||
{ID: "o4-mini", Label: "o4-mini", InputPer1k: 0.0011, OutputPer1k: 0.0044, ContextWindow: 200000},
|
||||
{ID: "gpt-4.1", Label: "GPT-4.1", InputPer1k: 0.002, OutputPer1k: 0.008, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-mini", Label: "GPT-4.1 mini", InputPer1k: 0.0004, OutputPer1k: 0.0016, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-nano", Label: "GPT-4.1 nano", InputPer1k: 0.0001, OutputPer1k: 0.0004, ContextWindow: 1047576},
|
||||
{ID: "gpt-4o", Label: "GPT-4o", InputPer1k: 0.0025, OutputPer1k: 0.010, ContextWindow: 128000},
|
||||
{ID: "gpt-4o-mini", Label: "GPT-4o mini", InputPer1k: 0.00015, OutputPer1k: 0.0006, ContextWindow: 128000},
|
||||
{ID: "gpt-4-turbo", Label: "GPT-4 Turbo", InputPer1k: 0.01, OutputPer1k: 0.03, ContextWindow: 128000},
|
||||
{ID: "gpt-3.5-turbo", Label: "GPT-3.5 Turbo", InputPer1k: 0.0005, OutputPer1k: 0.0015, ContextWindow: 16385},
|
||||
{ID: "text-embedding-3-large", Label: "text-embedding-3-large", InputPer1k: 0.00013, OutputPer1k: 0, ContextWindow: 8191},
|
||||
@@ -285,7 +257,6 @@ var providers = []Provider{
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#D97757",
|
||||
ParserID: "anthropic",
|
||||
PricingSurfaces: []string{"anthropic"},
|
||||
// Per Anthropic's current model lineup. Pricing in USD per 1k
|
||||
// tokens. Context windows: 4.6+ family is 1M; Haiku 4.5 stays at
|
||||
// 200K. claude-3-7-sonnet and claude-3-5-haiku retired
|
||||
@@ -296,14 +267,14 @@ var providers = []Provider{
|
||||
// account to be on >= 30-day data retention or all requests
|
||||
// 400.
|
||||
Models: []Model{
|
||||
{ID: "claude-fable-5", Label: "Claude Fable 5", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-6", Label: "Claude Opus 4.6", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (deprecated, retires 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000},
|
||||
{ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
|
||||
{ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000},
|
||||
{ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000},
|
||||
{ID: "claude-fable-5", Label: "Claude Fable 5", InputPer1k: 0.010, OutputPer1k: 0.050, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-6", Label: "Claude Opus 4.6", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (deprecated, retires 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000},
|
||||
{ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000},
|
||||
{ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000},
|
||||
{ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000},
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -317,19 +288,18 @@ var providers = []Provider{
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#0078D4",
|
||||
ParserID: "openai",
|
||||
PricingSurfaces: []string{"openai"},
|
||||
// Mirrors openai_api pricing — Azure resells OpenAI models at the
|
||||
// same per-token rates, just under different deployment names.
|
||||
Models: []Model{
|
||||
{ID: "gpt-5.5", Label: "GPT-5.5 (Azure)", InputPer1k: 0.005, OutputPer1k: 0.030, CachedInputPer1k: 0.0005, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4", Label: "GPT-5.4 (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.015, CachedInputPer1k: 0.00025, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini (Azure)", InputPer1k: 0.00075, OutputPer1k: 0.0045, CachedInputPer1k: 0.000075, ContextWindow: 272000},
|
||||
{ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano (Azure)", InputPer1k: 0.0002, OutputPer1k: 0.00125, CachedInputPer1k: 0.00002, ContextWindow: 272000},
|
||||
{ID: "o4-mini", Label: "o4-mini (Azure)", InputPer1k: 0.0011, OutputPer1k: 0.0044, CachedInputPer1k: 0.000275, ContextWindow: 200000},
|
||||
{ID: "gpt-4.1", Label: "GPT-4.1 (Azure)", InputPer1k: 0.002, OutputPer1k: 0.008, CachedInputPer1k: 0.0005, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-mini", Label: "GPT-4.1 mini (Azure)", InputPer1k: 0.0004, OutputPer1k: 0.0016, CachedInputPer1k: 0.0001, ContextWindow: 1047576},
|
||||
{ID: "gpt-4o", Label: "GPT-4o (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.010, CachedInputPer1k: 0.00125, ContextWindow: 128000},
|
||||
{ID: "gpt-4o-mini", Label: "GPT-4o mini (Azure)", InputPer1k: 0.00015, OutputPer1k: 0.0006, CachedInputPer1k: 0.000075, ContextWindow: 128000},
|
||||
{ID: "gpt-5.5", Label: "GPT-5.5 (Azure)", InputPer1k: 0.005, OutputPer1k: 0.030, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4", Label: "GPT-5.4 (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.015, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini (Azure)", InputPer1k: 0.00075, OutputPer1k: 0.0045, ContextWindow: 272000},
|
||||
{ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano (Azure)", InputPer1k: 0.0002, OutputPer1k: 0.00125, ContextWindow: 272000},
|
||||
{ID: "o4-mini", Label: "o4-mini (Azure)", InputPer1k: 0.0011, OutputPer1k: 0.0044, ContextWindow: 200000},
|
||||
{ID: "gpt-4.1", Label: "GPT-4.1 (Azure)", InputPer1k: 0.002, OutputPer1k: 0.008, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-mini", Label: "GPT-4.1 mini (Azure)", InputPer1k: 0.0004, OutputPer1k: 0.0016, ContextWindow: 1047576},
|
||||
{ID: "gpt-4o", Label: "GPT-4o (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.010, ContextWindow: 128000},
|
||||
{ID: "gpt-4o-mini", Label: "GPT-4o mini (Azure)", InputPer1k: 0.00015, OutputPer1k: 0.0006, ContextWindow: 128000},
|
||||
{ID: "gpt-35-turbo", Label: "GPT-3.5 Turbo (Azure)", InputPer1k: 0.0005, OutputPer1k: 0.0015, ContextWindow: 16385},
|
||||
},
|
||||
},
|
||||
@@ -343,9 +313,6 @@ var providers = []Provider{
|
||||
AuthHeaderTemplate: "Bearer ${API_KEY}",
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#FF9900",
|
||||
// ParserID stays empty (path-style dispatch via IsBedrockPathStyle);
|
||||
// the request parser meters these under the "bedrock" surface.
|
||||
PricingSurfaces: []string{"bedrock"},
|
||||
// Anthropic models on Bedrock take the anthropic.* prefix and
|
||||
// follow the same lineup / pricing as the first-party Anthropic
|
||||
// catalog entry above. claude-3-7-sonnet and claude-3-5-haiku
|
||||
@@ -355,13 +322,13 @@ var providers = []Provider{
|
||||
// Llama 3.3 70B entry kept unchanged — LiteLLM tracks only
|
||||
// per-region Llama 3 entries; standalone 3.3 not yet listed.
|
||||
Models: []Model{
|
||||
{ID: "anthropic.claude-opus-4-8", Label: "Claude Opus 4.8 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-7", Label: "Claude Opus 4.7 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-6", Label: "Claude Opus 4.6 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-1", Label: "Claude Opus 4.1 (Bedrock, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000},
|
||||
{ID: "anthropic.claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000},
|
||||
{ID: "anthropic.claude-haiku-4-5", Label: "Claude Haiku 4.5 (Bedrock)", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000},
|
||||
{ID: "anthropic.claude-opus-4-8", Label: "Claude Opus 4.8 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-7", Label: "Claude Opus 4.7 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-6", Label: "Claude Opus 4.6 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-1", Label: "Claude Opus 4.1 (Bedrock, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000},
|
||||
{ID: "anthropic.claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000},
|
||||
{ID: "anthropic.claude-haiku-4-5", Label: "Claude Haiku 4.5 (Bedrock)", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000},
|
||||
{ID: "meta.llama3-3-70b-instruct", Label: "Llama 3.3 70B (Bedrock)", InputPer1k: 0.00072, OutputPer1k: 0.00072, ContextWindow: 128000},
|
||||
{ID: "amazon.nova-2-lite", Label: "Amazon Nova 2 Lite (Bedrock, preview)", InputPer1k: 0.0003, OutputPer1k: 0.0025, ContextWindow: 1000000},
|
||||
{ID: "amazon.nova-pro", Label: "Amazon Nova Pro (Bedrock)", InputPer1k: 0.0008, OutputPer1k: 0.0032, ContextWindow: 300000},
|
||||
@@ -391,10 +358,6 @@ var providers = []Provider{
|
||||
AuthHeaderTemplate: "Bearer ${API_KEY}",
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#4285F4",
|
||||
// ParserID stays empty (path-style dispatch via IsVertexPathStyle);
|
||||
// Anthropic-on-Vertex requests are metered under the "anthropic"
|
||||
// surface with the bare, unversioned model id.
|
||||
PricingSurfaces: []string{"anthropic"},
|
||||
// Vertex carries the model in the URL path and authenticates with a
|
||||
// service-account-minted OAuth token (api_key = "keyfile::<base64 SA>").
|
||||
// Only Anthropic-on-Vertex is metered today: the request parser maps the
|
||||
@@ -406,14 +369,14 @@ var providers = []Provider{
|
||||
// exists — the router denies unmeterable publishers rather than forward
|
||||
// them uncounted.
|
||||
Models: []Model{
|
||||
{ID: "claude-fable-5", Label: "Claude Fable 5 (Vertex)", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-6", Label: "Claude Opus 4.6 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (Vertex, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000},
|
||||
{ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
|
||||
{ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000},
|
||||
{ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5 (Vertex)", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000},
|
||||
{ID: "claude-fable-5", Label: "Claude Fable 5 (Vertex)", InputPer1k: 0.010, OutputPer1k: 0.050, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-6", Label: "Claude Opus 4.6 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (Vertex, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000},
|
||||
{ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000},
|
||||
{ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000},
|
||||
{ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5 (Vertex)", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000},
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -427,7 +390,6 @@ var providers = []Provider{
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#FF7000",
|
||||
ParserID: "openai",
|
||||
PricingSurfaces: []string{"openai"},
|
||||
// Pricing + context windows cross-checked against LiteLLM. Key
|
||||
// gotchas the marketing page hides:
|
||||
// - `mistral-medium-latest` aliases to Medium 3.1 ($0.40/$2),
|
||||
@@ -486,10 +448,6 @@ var providers = []Provider{
|
||||
// model id "k3") is account-bound seat licensing rather than a
|
||||
// meterable platform key, so it's deliberately not the default.
|
||||
ParserID: "",
|
||||
// Both body shapes are metered: /v1/chat/completions under
|
||||
// "openai", /anthropic/v1/messages under "anthropic" — so the
|
||||
// K3 entry is priced on both surfaces.
|
||||
PricingSurfaces: []string{"openai", "anthropic"},
|
||||
// Pricing per Moonshot's platform rates at K3 launch (July 2026):
|
||||
// $3/$15 per MTok with $0.30 cached input, flat across the 1M-token
|
||||
// window. kimi-k3 is the ONLY model the platform serves newer
|
||||
@@ -500,14 +458,7 @@ var providers = []Provider{
|
||||
// The consumer app's "K3 Swarm Max" mode is not an API SKU, so it
|
||||
// doesn't appear here.
|
||||
Models: []Model{
|
||||
// Carries both cache shapes: Moonshot reports cache hits
|
||||
// OpenAI-style on /v1/chat/completions (CachedInputPer1k)
|
||||
// and Anthropic-style on /anthropic/v1/messages
|
||||
// (CacheReadPer1k) — $0.30/MTok either way. Each surface's
|
||||
// cost formula reads only its own field, so the superset
|
||||
// entry prices both endpoints correctly. No cache-creation
|
||||
// rate published; writes bill at the input rate.
|
||||
{ID: "kimi-k3", Label: "Kimi K3", InputPer1k: 0.003, OutputPer1k: 0.015, CachedInputPer1k: 0.0003, CacheReadPer1k: 0.0003, ContextWindow: 1000000},
|
||||
{ID: "kimi-k3", Label: "Kimi K3", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000},
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -807,28 +758,13 @@ func IsBedrockPathStyle(providerID string) bool {
|
||||
func (p Provider) ToAPIResponse() api.AgentNetworkCatalogProvider {
|
||||
models := make([]api.AgentNetworkCatalogModel, 0, len(p.Models))
|
||||
for _, m := range p.Models {
|
||||
am := api.AgentNetworkCatalogModel{
|
||||
models = append(models, api.AgentNetworkCatalogModel{
|
||||
Id: m.ID,
|
||||
Label: m.Label,
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
ContextWindow: m.ContextWindow,
|
||||
}
|
||||
// Cache rates are emitted only when configured so the dashboard
|
||||
// can prefill them; 0 stays off the wire (absent = no rate).
|
||||
if m.CachedInputPer1k > 0 {
|
||||
v := m.CachedInputPer1k
|
||||
am.CachedInputPer1k = &v
|
||||
}
|
||||
if m.CacheReadPer1k > 0 {
|
||||
v := m.CacheReadPer1k
|
||||
am.CacheReadPer1k = &v
|
||||
}
|
||||
if m.CacheCreationPer1k > 0 {
|
||||
v := m.CacheCreationPer1k
|
||||
am.CacheCreationPer1k = &v
|
||||
}
|
||||
models = append(models, am)
|
||||
})
|
||||
}
|
||||
kind := api.AgentNetworkCatalogProviderKindProvider
|
||||
switch p.Kind {
|
||||
@@ -848,10 +784,6 @@ func (p Provider) ToAPIResponse() api.AgentNetworkCatalogProvider {
|
||||
BrandColor: p.BrandColor,
|
||||
Models: models,
|
||||
}
|
||||
if len(p.PricingSurfaces) > 0 {
|
||||
surfaces := append([]string(nil), p.PricingSurfaces...)
|
||||
resp.PricingSurfaces = &surfaces
|
||||
}
|
||||
if len(p.ExtraHeaders) > 0 {
|
||||
extras := make([]api.AgentNetworkCatalogExtraHeader, 0, len(p.ExtraHeaders))
|
||||
for _, h := range p.ExtraHeaders {
|
||||
|
||||
@@ -7,7 +7,6 @@ package handlers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"math"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
@@ -16,7 +15,6 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
@@ -54,45 +52,11 @@ func (h *handler) getCatalogProviders(w http.ResponseWriter, r *http.Request) {
|
||||
entries := catalog.All()
|
||||
out := make([]api.AgentNetworkCatalogProvider, 0, len(entries))
|
||||
for _, e := range entries {
|
||||
resp := e.ToAPIResponse()
|
||||
applyDefaultPricing(e, &resp)
|
||||
out = append(out, resp)
|
||||
out = append(out, e.ToAPIResponse())
|
||||
}
|
||||
util.WriteJSONObject(r.Context(), w, out)
|
||||
}
|
||||
|
||||
// applyDefaultPricing overwrites the catalog response's model rates with
|
||||
// the LIVE default pricing table, which may differ from the compiled-in
|
||||
// catalog rates when the operator provides a defaults_llm_pricing.yaml.
|
||||
// This keeps the dashboard's model-row prefill identical to what the
|
||||
// proxy will actually bill — the same table the synthesizer ships.
|
||||
func applyDefaultPricing(cp catalog.Provider, resp *api.AgentNetworkCatalogProvider) {
|
||||
if len(cp.PricingSurfaces) == 0 {
|
||||
return
|
||||
}
|
||||
for i := range resp.Models {
|
||||
m := &resp.Models[i]
|
||||
e, ok := pricing.LookupDefault(cp.PricingSurfaces, m.Id)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
m.InputPer1k = e.InputPer1k
|
||||
m.OutputPer1k = e.OutputPer1k
|
||||
m.CachedInputPer1k = positiveRatePtr(e.CachedInputPer1k)
|
||||
m.CacheReadPer1k = positiveRatePtr(e.CacheReadPer1k)
|
||||
m.CacheCreationPer1k = positiveRatePtr(e.CacheCreationPer1k)
|
||||
}
|
||||
}
|
||||
|
||||
// positiveRatePtr renders a cache rate for the API: absent (nil) when
|
||||
// unset, matching the catalog response convention.
|
||||
func positiveRatePtr(v float64) *float64 {
|
||||
if v <= 0 {
|
||||
return nil
|
||||
}
|
||||
return &v
|
||||
}
|
||||
|
||||
func (h *handler) getAllProviders(w http.ResponseWriter, r *http.Request) {
|
||||
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
|
||||
if err != nil {
|
||||
@@ -249,38 +213,5 @@ func validate(req *api.AgentNetworkProviderRequest, requireAPIKey bool) error {
|
||||
if requireAPIKey && (req.ApiKey == nil || strings.TrimSpace(*req.ApiKey) == "") {
|
||||
return status.Errorf(status.InvalidArgument, "api_key is required")
|
||||
}
|
||||
if req.Models != nil {
|
||||
for i, m := range *req.Models {
|
||||
if err := validateModel(i, m); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateModel is the single ingress guard for operator-entered pricing:
|
||||
// these rates are synthesized into the proxy's cost_meter config verbatim,
|
||||
// and a negative or non-finite rate there would poison every cost the
|
||||
// proxy records, so reject at the API boundary.
|
||||
func validateModel(i int, m api.AgentNetworkProviderModel) error {
|
||||
if strings.TrimSpace(m.Id) == "" {
|
||||
return status.Errorf(status.InvalidArgument, "models[%d]: id is required", i)
|
||||
}
|
||||
rates := map[string]*float64{
|
||||
"input_per_1k": &m.InputPer1k,
|
||||
"output_per_1k": &m.OutputPer1k,
|
||||
"cached_input_per_1k": m.CachedInputPer1k,
|
||||
"cache_read_per_1k": m.CacheReadPer1k,
|
||||
"cache_creation_per_1k": m.CacheCreationPer1k,
|
||||
}
|
||||
for field, v := range rates {
|
||||
if v == nil {
|
||||
continue
|
||||
}
|
||||
if *v < 0 || math.IsNaN(*v) || math.IsInf(*v, 0) {
|
||||
return status.Errorf(status.InvalidArgument, "models[%d] (%s): %s must be a finite, non-negative USD rate", i, m.Id, field)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,53 +0,0 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
func f(v float64) *float64 { return &v }
|
||||
|
||||
// TestValidate_ModelRates guards the single ingress point for operator-entered
|
||||
// pricing. These rates flow verbatim into the proxy's cost_meter config at
|
||||
// synthesis time; the proxy treats a bad rate as a chain-build failure, so
|
||||
// rejecting here is what keeps an account's gateway from going down.
|
||||
func TestValidate_ModelRates(t *testing.T) {
|
||||
base := func(models ...api.AgentNetworkProviderModel) *api.AgentNetworkProviderRequest {
|
||||
key := "sk-test"
|
||||
return &api.AgentNetworkProviderRequest{
|
||||
ProviderId: "openai_api",
|
||||
Name: "OpenAI",
|
||||
UpstreamUrl: "https://api.openai.com",
|
||||
ApiKey: &key,
|
||||
Models: &models,
|
||||
}
|
||||
}
|
||||
|
||||
valid := api.AgentNetworkProviderModel{
|
||||
Id: "gpt-4o", InputPer1k: 0.0025, OutputPer1k: 0.01,
|
||||
CachedInputPer1k: f(0.00125),
|
||||
}
|
||||
require.NoError(t, validate(base(valid), true), "finite non-negative rates must pass")
|
||||
|
||||
zeroRates := api.AgentNetworkProviderModel{Id: "self-hosted-llama", InputPer1k: 0, OutputPer1k: 0}
|
||||
require.NoError(t, validate(base(zeroRates), true), "explicit zero prices are allowed (free / self-hosted models)")
|
||||
|
||||
cases := map[string]api.AgentNetworkProviderModel{
|
||||
"empty id": {Id: " ", InputPer1k: 0.001, OutputPer1k: 0.002},
|
||||
"negative input": {Id: "m", InputPer1k: -0.001, OutputPer1k: 0.002},
|
||||
"negative output": {Id: "m", InputPer1k: 0.001, OutputPer1k: -0.002},
|
||||
"NaN input": {Id: "m", InputPer1k: math.NaN(), OutputPer1k: 0.002},
|
||||
"Inf output": {Id: "m", InputPer1k: 0.001, OutputPer1k: math.Inf(1)},
|
||||
"negative cached": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CachedInputPer1k: f(-1)},
|
||||
"NaN cache read": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CacheReadPer1k: f(math.NaN())},
|
||||
"Inf cache creation": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CacheCreationPer1k: f(math.Inf(-1))},
|
||||
}
|
||||
for name, m := range cases {
|
||||
assert.Error(t, validate(base(m), true), "case %q must be rejected", name)
|
||||
}
|
||||
}
|
||||
@@ -1,156 +0,0 @@
|
||||
// Package pricing builds the default LLM pricing table the synthesizer
|
||||
// ships to the proxy's cost_meter middleware. The catalog is the single
|
||||
// source of default rates: every catalog provider's models are folded
|
||||
// into the pricing surfaces the provider declares (PricingSurfaces),
|
||||
// then a small supplemental list adds priced-but-not-operator-selectable
|
||||
// entries. Management is the sole pricing authority — the proxy carries
|
||||
// no embedded price list and bills exclusively from the table it is
|
||||
// sent.
|
||||
//go:generate go run gen.go
|
||||
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
)
|
||||
|
||||
// Entry is a single model's pricing in USD per 1k tokens. This struct IS
|
||||
// the wire shape: the synthesizer marshals it verbatim into cost_meter's
|
||||
// ConfigJSON, and the proxy unmarshals the same field names.
|
||||
//
|
||||
// A zero rate means "no rate configured" — the proxy bills that cache
|
||||
// bucket at InputPer1k (identical semantics to the retired proxy-embedded
|
||||
// table). CachedInputPer1k is the OpenAI shape (cached prompt tokens are
|
||||
// a subset of input); CacheReadPer1k / CacheCreationPer1k are the
|
||||
// Anthropic shape (additive buckets).
|
||||
type Entry struct {
|
||||
InputPer1k float64 `json:"input_per_1k"`
|
||||
OutputPer1k float64 `json:"output_per_1k"`
|
||||
CachedInputPer1k float64 `json:"cached_input_per_1k,omitempty"`
|
||||
CacheReadPer1k float64 `json:"cache_read_per_1k,omitempty"`
|
||||
CacheCreationPer1k float64 `json:"cache_creation_per_1k,omitempty"`
|
||||
}
|
||||
|
||||
// supplementalDefaults are (surface, model) entries that are priced but
|
||||
// deliberately not operator-selectable in the catalog. Each carries a
|
||||
// reason; when one of these models joins a catalog lineup, delete the
|
||||
// row here — the collision test fails loudly if the rates ever disagree.
|
||||
var supplementalDefaults = map[string]map[string]Entry{
|
||||
"openai": {
|
||||
// GPT-5 (2025) family — kept for gateway requests using the
|
||||
// unsuffixed ids; the dashboard offers only the 5.x lineup.
|
||||
"gpt-5": {InputPer1k: 0.00125, OutputPer1k: 0.01, CachedInputPer1k: 0.000125},
|
||||
"gpt-5-mini": {InputPer1k: 0.00025, OutputPer1k: 0.002, CachedInputPer1k: 0.000025},
|
||||
"gpt-5-nano": {InputPer1k: 0.00005, OutputPer1k: 0.0004, CachedInputPer1k: 0.000005},
|
||||
},
|
||||
"anthropic": {
|
||||
// claude-opus-5 is not yet in the catalog lineup but gateway /
|
||||
// grandfathered traffic uses it; priced so it isn't skipped.
|
||||
"claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625},
|
||||
// "kimi-k3[1m]" is the 1M-context alias some Claude Code guides
|
||||
// configure against Moonshot's Anthropic-compatible endpoint;
|
||||
// priced identically to kimi-k3 so those requests aren't skipped.
|
||||
"kimi-k3[1m]": {InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003},
|
||||
},
|
||||
"bedrock": {
|
||||
"anthropic.claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625},
|
||||
},
|
||||
}
|
||||
|
||||
var (
|
||||
compiledOnce sync.Once
|
||||
compiledTable map[string]map[string]Entry
|
||||
// mergedTable holds the current live table when a pricing defaults
|
||||
// file is loaded: the file merged entry-whole over the compiled-in
|
||||
// base. Nil while no file is loaded (or after the file is removed),
|
||||
// in which case the compiled-in table serves. Swapped atomically by
|
||||
// the file loader/reloader; readers never block.
|
||||
mergedTable atomic.Pointer[map[string]map[string]Entry]
|
||||
)
|
||||
|
||||
// DefaultTable returns the current default pricing table keyed
|
||||
// surface -> model -> Entry: the management-side defaults file (see
|
||||
// LoadFile / StartReloader) when one is loaded, merged over the
|
||||
// compiled-in catalog table, which alone serves as the fallback when no
|
||||
// file exists. The snapshot may change between calls as the file is
|
||||
// re-read — consumers (the synthesizer on every reconcile, the catalog
|
||||
// endpoint on every request) pick up fresh rates automatically. Callers
|
||||
// must not mutate the returned maps.
|
||||
func DefaultTable() map[string]map[string]Entry {
|
||||
if t := mergedTable.Load(); t != nil {
|
||||
return *t
|
||||
}
|
||||
return compiledBase()
|
||||
}
|
||||
|
||||
// compiledBase returns the compiled-in table (catalog + supplementals),
|
||||
// built once.
|
||||
func compiledBase() map[string]map[string]Entry {
|
||||
compiledOnce.Do(func() {
|
||||
compiledTable = buildDefaultTable()
|
||||
})
|
||||
return compiledTable
|
||||
}
|
||||
|
||||
func buildDefaultTable() map[string]map[string]Entry {
|
||||
out := make(map[string]map[string]Entry)
|
||||
for _, p := range catalog.All() {
|
||||
for _, surface := range p.PricingSurfaces {
|
||||
inner, ok := out[surface]
|
||||
if !ok {
|
||||
inner = make(map[string]Entry)
|
||||
out[surface] = inner
|
||||
}
|
||||
for _, m := range p.Models {
|
||||
// First writer wins; providers contributing the same
|
||||
// (surface, model) must agree on rates — enforced by
|
||||
// TestDefaultTable_NoConflictingContributions.
|
||||
if _, dup := inner[m.ID]; dup {
|
||||
continue
|
||||
}
|
||||
inner[m.ID] = entryFromCatalogModel(m)
|
||||
}
|
||||
}
|
||||
}
|
||||
for surface, models := range supplementalDefaults {
|
||||
inner, ok := out[surface]
|
||||
if !ok {
|
||||
inner = make(map[string]Entry)
|
||||
out[surface] = inner
|
||||
}
|
||||
for id, e := range models {
|
||||
if _, dup := inner[id]; dup {
|
||||
continue
|
||||
}
|
||||
inner[id] = e
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func entryFromCatalogModel(m catalog.Model) Entry {
|
||||
return Entry{
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
CachedInputPer1k: m.CachedInputPer1k,
|
||||
CacheReadPer1k: m.CacheReadPer1k,
|
||||
CacheCreationPer1k: m.CacheCreationPer1k,
|
||||
}
|
||||
}
|
||||
|
||||
// LookupDefault returns the default entry for model on the first of the
|
||||
// given surfaces that prices it. Used by the synthesizer to seed a
|
||||
// per-provider entry with default cache rates before overlaying the
|
||||
// operator's stored prices.
|
||||
func LookupDefault(surfaces []string, model string) (Entry, bool) {
|
||||
table := DefaultTable()
|
||||
for _, s := range surfaces {
|
||||
if e, ok := table[s][model]; ok {
|
||||
return e, true
|
||||
}
|
||||
}
|
||||
return Entry{}, false
|
||||
}
|
||||
@@ -1,148 +0,0 @@
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
)
|
||||
|
||||
// TestDefaultTable_CoversEveryCatalogModel replaces the proxy's old
|
||||
// hand-maintained coverage list: because the table is built FROM the
|
||||
// catalog, drift is impossible by construction — this test guards the
|
||||
// fold itself (every catalog model of every surfaced provider resolves,
|
||||
// with exactly the catalog's rates).
|
||||
func TestDefaultTable_CoversEveryCatalogModel(t *testing.T) {
|
||||
table := DefaultTable()
|
||||
for _, p := range catalog.All() {
|
||||
if len(p.PricingSurfaces) == 0 {
|
||||
assert.Empty(t, p.Models, "catalog entry %s declares models but no pricing surfaces — those models would never be priced", p.ID)
|
||||
continue
|
||||
}
|
||||
for _, surface := range p.PricingSurfaces {
|
||||
byModel, ok := table[surface]
|
||||
require.True(t, ok, "surface %q (provider %s) missing from default table", surface, p.ID)
|
||||
for _, m := range p.Models {
|
||||
e, ok := byModel[m.ID]
|
||||
require.True(t, ok, "%s/%s (provider %s) missing from default table", surface, m.ID, p.ID)
|
||||
assert.Equal(t, m.InputPer1k, e.InputPer1k, "%s/%s input rate", surface, m.ID)
|
||||
assert.Equal(t, m.OutputPer1k, e.OutputPer1k, "%s/%s output rate", surface, m.ID)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestDefaultTable_NoConflictingContributions enforces the collision rule
|
||||
// documented on catalog.Provider.PricingSurfaces: when two catalog
|
||||
// providers contribute the same (surface, model) pair — azure/vertex
|
||||
// mirroring openai/anthropic, kimi on both surfaces — their rates must be
|
||||
// identical, because the surface-keyed table can only hold one entry.
|
||||
// If a provider ever diverges (e.g. Azure reprices a model), this fails
|
||||
// and the divergence must move to per-provider-record pricing.
|
||||
func TestDefaultTable_NoConflictingContributions(t *testing.T) {
|
||||
type contribution struct {
|
||||
providerID string
|
||||
entry Entry
|
||||
}
|
||||
seen := map[string]map[string]contribution{}
|
||||
for _, p := range catalog.All() {
|
||||
for _, surface := range p.PricingSurfaces {
|
||||
if seen[surface] == nil {
|
||||
seen[surface] = map[string]contribution{}
|
||||
}
|
||||
for _, m := range p.Models {
|
||||
e := entryFromCatalogModel(m)
|
||||
if prev, dup := seen[surface][m.ID]; dup {
|
||||
assert.Equal(t, prev.entry, e,
|
||||
"%s/%s: %s and %s contribute different rates", surface, m.ID, prev.providerID, p.ID)
|
||||
continue
|
||||
}
|
||||
seen[surface][m.ID] = contribution{providerID: p.ID, entry: e}
|
||||
}
|
||||
}
|
||||
}
|
||||
// Supplemental entries must never shadow a catalog-contributed model —
|
||||
// they exist precisely because the catalog does NOT list them.
|
||||
for surface, models := range supplementalDefaults {
|
||||
for id := range models {
|
||||
_, fromCatalog := seen[surface][id]
|
||||
assert.False(t, fromCatalog, "supplemental %s/%s is now in the catalog — delete the supplemental row", surface, id)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestDefaultTable_AllRatesFiniteNonNegative mirrors the proxy-side
|
||||
// NewTable validation so a bad catalog edit is caught here, at unit-test
|
||||
// time, rather than as a chain-build failure on every proxy.
|
||||
func TestDefaultTable_AllRatesFiniteNonNegative(t *testing.T) {
|
||||
for surface, models := range DefaultTable() {
|
||||
for id, e := range models {
|
||||
for field, v := range map[string]float64{
|
||||
"input": e.InputPer1k,
|
||||
"output": e.OutputPer1k,
|
||||
"cached_input": e.CachedInputPer1k,
|
||||
"cache_read": e.CacheReadPer1k,
|
||||
"cache_creation": e.CacheCreationPer1k,
|
||||
} {
|
||||
assert.False(t, v < 0 || math.IsNaN(v) || math.IsInf(v, 0),
|
||||
"%s/%s: %s rate %v must be finite and non-negative", surface, id, field, v)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestDefaultTable_PinnedRates pins rates that previously drifted or are
|
||||
// easy to mis-enter (carried over from the proxy's retired
|
||||
// defaults_coverage_test), plus the supplemental entries.
|
||||
func TestDefaultTable_PinnedRates(t *testing.T) {
|
||||
table := DefaultTable()
|
||||
|
||||
gpt54 := table["openai"]["gpt-5.4"]
|
||||
assert.InDelta(t, 0.0025, gpt54.InputPer1k, 1e-9, "gpt-5.4 input")
|
||||
assert.InDelta(t, 0.015, gpt54.OutputPer1k, 1e-9, "gpt-5.4 output")
|
||||
assert.InDelta(t, 0.00025, gpt54.CachedInputPer1k, 1e-9, "gpt-5.4 cached input")
|
||||
|
||||
sonnet := table["bedrock"]["anthropic.claude-sonnet-4-5"]
|
||||
assert.InDelta(t, 0.003, sonnet.InputPer1k, 1e-9, "bedrock sonnet-4-5 input")
|
||||
assert.InDelta(t, 0.015, sonnet.OutputPer1k, 1e-9, "bedrock sonnet-4-5 output")
|
||||
assert.InDelta(t, 0.0003, sonnet.CacheReadPer1k, 1e-9, "bedrock sonnet-4-5 cache read")
|
||||
assert.InDelta(t, 0.00375, sonnet.CacheCreationPer1k, 1e-9, "bedrock sonnet-4-5 cache creation")
|
||||
|
||||
// Vertex Claude prices under "anthropic" with the bare id.
|
||||
fable := table["anthropic"]["claude-fable-5"]
|
||||
assert.InDelta(t, 0.010, fable.InputPer1k, 1e-9, "claude-fable-5 input")
|
||||
assert.InDelta(t, 0.0125, fable.CacheCreationPer1k, 1e-9, "claude-fable-5 cache creation")
|
||||
|
||||
// Supplementals present on their surfaces.
|
||||
for surface, ids := range map[string][]string{
|
||||
"openai": {"gpt-5", "gpt-5-mini", "gpt-5-nano"},
|
||||
"anthropic": {"claude-opus-5", "kimi-k3[1m]", "kimi-k3"},
|
||||
"bedrock": {"anthropic.claude-opus-5"},
|
||||
} {
|
||||
for _, id := range ids {
|
||||
_, ok := table[surface][id]
|
||||
assert.True(t, ok, "%s/%s must be priced", surface, id)
|
||||
}
|
||||
}
|
||||
|
||||
// Embeddings bill input-only — output stays zero.
|
||||
emb := table["openai"]["text-embedding-3-large"]
|
||||
assert.Zero(t, emb.OutputPer1k, "embedding output rate must be zero")
|
||||
assert.Positive(t, emb.InputPer1k, "embedding input rate must be set")
|
||||
}
|
||||
|
||||
func TestLookupDefault_SurfaceOrder(t *testing.T) {
|
||||
// kimi-k3 exists on both surfaces; first surface in the slice wins.
|
||||
e, ok := LookupDefault([]string{"openai", "anthropic"}, "kimi-k3")
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.003, e.InputPer1k, 1e-9)
|
||||
|
||||
_, ok = LookupDefault([]string{"bedrock"}, "gpt-4o")
|
||||
assert.False(t, ok, "gpt-4o is not a bedrock model")
|
||||
|
||||
_, ok = LookupDefault(nil, "gpt-4o")
|
||||
assert.False(t, ok, "no surfaces, no match")
|
||||
}
|
||||
@@ -1,109 +0,0 @@
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// MarshalDefaultsYAML renders the built-in default pricing table (catalog
|
||||
// + supplementals, WITHOUT any operator override) as the YAML schema
|
||||
// LoadOverrideFile consumes. It backs the generated
|
||||
// defaults_llm_pricing.example.yaml so operators start from a file that
|
||||
// matches the compiled-in rates exactly; a golden test keeps the two in
|
||||
// sync. Output is deterministic (sorted surfaces and models).
|
||||
func MarshalDefaultsYAML() []byte {
|
||||
var b bytes.Buffer
|
||||
b.WriteString(exampleHeader)
|
||||
|
||||
table := buildDefaultTable()
|
||||
surfaces := make([]string, 0, len(table))
|
||||
for s := range table {
|
||||
surfaces = append(surfaces, s)
|
||||
}
|
||||
sort.Strings(surfaces)
|
||||
|
||||
for _, surface := range surfaces {
|
||||
fmt.Fprintf(&b, "\n%s:\n", surface)
|
||||
models := table[surface]
|
||||
ids := make([]string, 0, len(models))
|
||||
for id := range models {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Strings(ids)
|
||||
for _, id := range ids {
|
||||
e := models[id]
|
||||
fmt.Fprintf(&b, " %s:\n", yamlKey(id))
|
||||
fmt.Fprintf(&b, " input_per_1k: %s\n", rate(e.InputPer1k))
|
||||
fmt.Fprintf(&b, " output_per_1k: %s\n", rate(e.OutputPer1k))
|
||||
if e.CachedInputPer1k > 0 {
|
||||
fmt.Fprintf(&b, " cached_input_per_1k: %s\n", rate(e.CachedInputPer1k))
|
||||
}
|
||||
if e.CacheReadPer1k > 0 {
|
||||
fmt.Fprintf(&b, " cache_read_per_1k: %s\n", rate(e.CacheReadPer1k))
|
||||
}
|
||||
if e.CacheCreationPer1k > 0 {
|
||||
fmt.Fprintf(&b, " cache_creation_per_1k: %s\n", rate(e.CacheCreationPer1k))
|
||||
}
|
||||
}
|
||||
}
|
||||
return b.Bytes()
|
||||
}
|
||||
|
||||
// rate renders a USD-per-1k rate without float noise ("0.00015", not
|
||||
// "0.000150000000...").
|
||||
func rate(v float64) string {
|
||||
return strconv.FormatFloat(v, 'f', -1, 64)
|
||||
}
|
||||
|
||||
// yamlKey quotes model ids that YAML would otherwise misparse (e.g.
|
||||
// "kimi-k3[1m]" starts a flow sequence unquoted).
|
||||
func yamlKey(id string) string {
|
||||
for _, r := range id {
|
||||
switch r {
|
||||
case '[', ']', '{', '}', ':', '#', ',', '&', '*', '!', '|', '>', '\'', '"', '%', '@', '`':
|
||||
return strconv.Quote(id)
|
||||
}
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
const exampleHeader = `# Default LLM pricing used by NetBird's Agent Network cost metering.
|
||||
# GENERATED from the management catalog — do not edit this file in the
|
||||
# repository; regenerate with:
|
||||
#
|
||||
# go generate ./management/internals/modules/agentnetwork/pricing
|
||||
#
|
||||
# Operators: copy this file to <datadir>/defaults_llm_pricing.yaml (or
|
||||
# any path configured via management.json:
|
||||
#
|
||||
# { "AgentNetwork": { "PricingDefaultsFile": "/path/defaults_llm_pricing.yaml" } }
|
||||
#
|
||||
# ) and adjust the entries you want to change. Management re-reads the
|
||||
# file periodically (mtime poll, every minute): the live table feeds the
|
||||
# proxies' cost metering and the dashboard's model-price prefill, so
|
||||
# edits apply without a restart. Your file only needs the entries you
|
||||
# want to change — but each entry REPLACES the built-in entry for that
|
||||
# surface+model whole, so repeat the cache rates you want to keep.
|
||||
# Unknown fields and negative or non-finite rates are rejected: at
|
||||
# startup that fails boot (for an explicitly configured path); at
|
||||
# runtime the previous table is kept and a warning is logged. Deleting
|
||||
# the file reverts to the built-in defaults below.
|
||||
#
|
||||
# Top-level keys are pricing surfaces — the parser shape requests are
|
||||
# metered under: "openai" (also Azure, Mistral, and OpenAI-compatible
|
||||
# gateways), "anthropic" (also Anthropic-on-Vertex), "bedrock"
|
||||
# (normalized ids, e.g. anthropic.claude-sonnet-4-5). Model keys must be
|
||||
# the normalized id the proxy meters (version/region suffixes stripped).
|
||||
#
|
||||
# Values are USD per 1_000 tokens. Optional cache fields:
|
||||
# cached_input_per_1k OpenAI shape: rate for cached prompt tokens
|
||||
# (a SUBSET of input tokens). Absent -> cached
|
||||
# portion bills at input_per_1k.
|
||||
# cache_read_per_1k Anthropic shape: rate for cache_read tokens
|
||||
# (ADDITIVE to input). Absent -> input rate.
|
||||
# cache_creation_per_1k Anthropic shape: rate for cache_creation
|
||||
# tokens (ADDITIVE to input). Absent -> input
|
||||
# rate.
|
||||
`
|
||||
@@ -1,20 +0,0 @@
|
||||
//go:build ignore
|
||||
|
||||
// Regenerates defaults_llm_pricing.example.yaml from the compiled-in
|
||||
// default pricing table. Run via:
|
||||
//
|
||||
// go generate ./management/internals/modules/agentnetwork/pricing
|
||||
package main
|
||||
|
||||
import (
|
||||
"log"
|
||||
"os"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
)
|
||||
|
||||
func main() {
|
||||
if err := os.WriteFile("defaults_llm_pricing.example.yaml", pricing.MarshalDefaultsYAML(), 0o644); err != nil {
|
||||
log.Fatalf("write defaults_llm_pricing.example.yaml: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -1,249 +0,0 @@
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"math"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
// DefaultFileName is the basename probed under management's datadir when
|
||||
// AgentNetwork.PricingDefaultsFile doesn't configure an explicit path.
|
||||
const DefaultFileName = "defaults_llm_pricing.yaml"
|
||||
|
||||
// ReloadInterval is the cadence at which the pricing file's mtime is
|
||||
// polled for changes.
|
||||
const ReloadInterval = time.Minute
|
||||
|
||||
// maxFileBytes bounds the pricing file read so a misconfigured path
|
||||
// (pointed at a huge file) cannot exhaust process memory.
|
||||
const maxFileBytes = 4 << 20
|
||||
|
||||
// pricingFile mirrors the on-disk YAML schema — the same schema the
|
||||
// proxy's retired embedded defaults_pricing.yaml used, so files written
|
||||
// for it keep working. Keys are pricing surfaces ("openai", "anthropic",
|
||||
// "bedrock"); nested keys are normalized model ids.
|
||||
type pricingFile map[string]map[string]struct {
|
||||
InputPer1k float64 `yaml:"input_per_1k"`
|
||||
OutputPer1k float64 `yaml:"output_per_1k"`
|
||||
CachedInputPer1k float64 `yaml:"cached_input_per_1k"`
|
||||
CacheReadPer1k float64 `yaml:"cache_read_per_1k"`
|
||||
CacheCreationPer1k float64 `yaml:"cache_creation_per_1k"`
|
||||
}
|
||||
|
||||
// fileState tracks the watched pricing file across reloads.
|
||||
var fileState struct {
|
||||
mu sync.Mutex
|
||||
path string
|
||||
mtime int64
|
||||
}
|
||||
|
||||
// LoadFile loads the management-side pricing defaults file and makes it
|
||||
// the live table (merged entry-whole over the compiled-in fallback; see
|
||||
// DefaultTable). The path stays registered for the periodic reloader, so
|
||||
// later edits — or the file (re)appearing after deletion — are picked up
|
||||
// without a restart.
|
||||
//
|
||||
// required governs the missing-file case: true for an explicitly
|
||||
// configured path (a typo must fail startup rather than silently bill
|
||||
// with built-ins the operator believes they replaced), false for the
|
||||
// conventional <datadir>/defaults_llm_pricing.yaml probe (absent file =
|
||||
// compiled-in defaults, still watched in case it appears). A file that
|
||||
// exists but is malformed is always an error at load time.
|
||||
func LoadFile(path string, required bool) error {
|
||||
if path == "" {
|
||||
return nil
|
||||
}
|
||||
fileState.mu.Lock()
|
||||
fileState.path = path
|
||||
fileState.mu.Unlock()
|
||||
|
||||
table, mtime, err := readFile(path)
|
||||
if err != nil {
|
||||
if errors.Is(err, fs.ErrNotExist) && !required {
|
||||
log.Infof("agent-network pricing defaults file %s not present; serving built-in defaults", path)
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
storeFileTable(table, mtime)
|
||||
log.Infof("agent-network pricing defaults loaded from %s", path)
|
||||
return nil
|
||||
}
|
||||
|
||||
// StartReloader launches the periodic mtime-poll goroutine for the file
|
||||
// registered by LoadFile. Runtime failures are lenient — a parse error
|
||||
// keeps the previously loaded table and logs a warning; a deleted file
|
||||
// reverts to the compiled-in defaults — so a mid-edit save can never
|
||||
// take pricing down. Returns immediately when no path was registered.
|
||||
func StartReloader(ctx context.Context, interval time.Duration) {
|
||||
fileState.mu.Lock()
|
||||
path := fileState.path
|
||||
fileState.mu.Unlock()
|
||||
if path == "" {
|
||||
return
|
||||
}
|
||||
if interval <= 0 {
|
||||
interval = ReloadInterval
|
||||
}
|
||||
go func() {
|
||||
t := time.NewTicker(interval)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-t.C:
|
||||
reload()
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// reload performs one mtime check + reload cycle.
|
||||
func reload() {
|
||||
fileState.mu.Lock()
|
||||
path, lastMtime := fileState.path, fileState.mtime
|
||||
fileState.mu.Unlock()
|
||||
|
||||
log.Debugf("agent-network pricing defaults reload: checking %s for changes", path)
|
||||
|
||||
st, err := os.Stat(path)
|
||||
if err != nil {
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
// File removed (or not yet created): serve compiled-in
|
||||
// defaults and reset mtime so a future (re)appearance loads.
|
||||
if mergedTable.Swap(nil) != nil {
|
||||
log.Warnf("agent-network pricing defaults file %s removed; reverting to built-in defaults", path)
|
||||
}
|
||||
setMtime(0)
|
||||
return
|
||||
}
|
||||
log.Warnf("agent-network pricing defaults reload: stat %s: %v", path, err)
|
||||
return
|
||||
}
|
||||
if st.ModTime().UnixNano() == lastMtime {
|
||||
log.Debugf("agent-network pricing defaults %s unchanged since last check", path)
|
||||
return
|
||||
}
|
||||
|
||||
table, mtime, err := readFile(path)
|
||||
if err != nil {
|
||||
// Keep the previously loaded table — never blank prices because
|
||||
// an operator saved mid-edit.
|
||||
log.Warnf("agent-network pricing defaults reload failed for %s (keeping previous table): %v", path, err)
|
||||
return
|
||||
}
|
||||
storeFileTable(table, mtime)
|
||||
log.Infof("agent-network pricing defaults reloaded from %s", path)
|
||||
}
|
||||
|
||||
func readFile(path string) (map[string]map[string]Entry, int64, error) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("open pricing defaults %s: %w", path, err)
|
||||
}
|
||||
defer func() { _ = f.Close() }()
|
||||
|
||||
st, err := f.Stat()
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("stat pricing defaults %s: %w", path, err)
|
||||
}
|
||||
data, err := io.ReadAll(io.LimitReader(f, maxFileBytes+1))
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("read pricing defaults %s: %w", path, err)
|
||||
}
|
||||
if len(data) > maxFileBytes {
|
||||
return nil, 0, fmt.Errorf("pricing defaults %s exceeds %d bytes", path, maxFileBytes)
|
||||
}
|
||||
table, err := parsePricingYAML(data)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("parse pricing defaults %s: %w", path, err)
|
||||
}
|
||||
return table, st.ModTime().UnixNano(), nil
|
||||
}
|
||||
|
||||
// storeFileTable merges the parsed file over the compiled-in base and
|
||||
// publishes the result as the live table. File entries replace the
|
||||
// built-in entry for the same (surface, model) whole — they are not
|
||||
// field-merged — and surfaces/models the file doesn't mention keep the
|
||||
// built-in rates, so a partial file only needs the entries it changes.
|
||||
func storeFileTable(table map[string]map[string]Entry, mtime int64) {
|
||||
base := compiledBase()
|
||||
merged := make(map[string]map[string]Entry, len(base)+len(table))
|
||||
for surface, models := range base {
|
||||
inner := make(map[string]Entry, len(models))
|
||||
for id, e := range models {
|
||||
inner[id] = e
|
||||
}
|
||||
merged[surface] = inner
|
||||
}
|
||||
for surface, models := range table {
|
||||
inner, ok := merged[surface]
|
||||
if !ok {
|
||||
inner = make(map[string]Entry, len(models))
|
||||
merged[surface] = inner
|
||||
}
|
||||
for id, e := range models {
|
||||
inner[id] = e
|
||||
}
|
||||
}
|
||||
mergedTable.Store(&merged)
|
||||
setMtime(mtime)
|
||||
}
|
||||
|
||||
func setMtime(v int64) {
|
||||
fileState.mu.Lock()
|
||||
fileState.mtime = v
|
||||
fileState.mu.Unlock()
|
||||
}
|
||||
|
||||
// parsePricingYAML decodes and validates the pricing YAML. Unknown
|
||||
// fields are rejected (typos surface instead of silently pricing at 0)
|
||||
// and every rate must be a finite, non-negative USD amount — the same
|
||||
// constraints the HTTP API enforces on operator per-provider prices.
|
||||
func parsePricingYAML(data []byte) (map[string]map[string]Entry, error) {
|
||||
dec := yaml.NewDecoder(bytes.NewReader(data))
|
||||
dec.KnownFields(true)
|
||||
|
||||
var raw pricingFile
|
||||
if err := dec.Decode(&raw); err != nil && !errors.Is(err, io.EOF) {
|
||||
return nil, fmt.Errorf("decode yaml: %w", err)
|
||||
}
|
||||
|
||||
out := make(map[string]map[string]Entry, len(raw))
|
||||
for surface, models := range raw {
|
||||
inner := make(map[string]Entry, len(models))
|
||||
for model, e := range models {
|
||||
for field, v := range map[string]float64{
|
||||
"input_per_1k": e.InputPer1k,
|
||||
"output_per_1k": e.OutputPer1k,
|
||||
"cached_input_per_1k": e.CachedInputPer1k,
|
||||
"cache_read_per_1k": e.CacheReadPer1k,
|
||||
"cache_creation_per_1k": e.CacheCreationPer1k,
|
||||
} {
|
||||
if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) {
|
||||
return nil, fmt.Errorf("%s/%s: %s must be a finite, non-negative rate, got %v", surface, model, field, v)
|
||||
}
|
||||
}
|
||||
inner[model] = Entry{
|
||||
InputPer1k: e.InputPer1k,
|
||||
OutputPer1k: e.OutputPer1k,
|
||||
CachedInputPer1k: e.CachedInputPer1k,
|
||||
CacheReadPer1k: e.CacheReadPer1k,
|
||||
CacheCreationPer1k: e.CacheCreationPer1k,
|
||||
}
|
||||
}
|
||||
out[surface] = inner
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
@@ -1,173 +0,0 @@
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// resetFileState snapshots and restores the package-level file state so
|
||||
// tests stay order-independent.
|
||||
func resetFileState(t *testing.T) {
|
||||
t.Helper()
|
||||
prevMerged := mergedTable.Load()
|
||||
fileState.mu.Lock()
|
||||
prevPath, prevMtime := fileState.path, fileState.mtime
|
||||
fileState.mu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
mergedTable.Store(prevMerged)
|
||||
fileState.mu.Lock()
|
||||
fileState.path, fileState.mtime = prevPath, prevMtime
|
||||
fileState.mu.Unlock()
|
||||
})
|
||||
}
|
||||
|
||||
func writePricing(t *testing.T, path, yml string) {
|
||||
t.Helper()
|
||||
require.NoError(t, os.WriteFile(path, []byte(yml), 0o600))
|
||||
}
|
||||
|
||||
func TestLoadFile_MergesOverCompiledDefaults(t *testing.T) {
|
||||
resetFileState(t)
|
||||
path := filepath.Join(t.TempDir(), DefaultFileName)
|
||||
writePricing(t, path, `
|
||||
openai:
|
||||
# Reprice a built-in model. The entry replaces the built-in WHOLE:
|
||||
# omitting the cache rate here drops the built-in 0.00125 discount.
|
||||
gpt-4o:
|
||||
input_per_1k: 0.9
|
||||
output_per_1k: 1.8
|
||||
# A model NetBird doesn't know at all.
|
||||
my-private-ft:
|
||||
input_per_1k: 0.01
|
||||
output_per_1k: 0.02
|
||||
cached_input_per_1k: 0.005
|
||||
gemini:
|
||||
gemini-pro:
|
||||
input_per_1k: 0.00125
|
||||
output_per_1k: 0.005
|
||||
`)
|
||||
require.NoError(t, LoadFile(path, true))
|
||||
table := DefaultTable()
|
||||
|
||||
gpt4o := table["openai"]["gpt-4o"]
|
||||
assert.InDelta(t, 0.9, gpt4o.InputPer1k, 1e-9, "file rate replaces the compiled-in rate")
|
||||
assert.Zero(t, gpt4o.CachedInputPer1k, "entries replace whole — omitted cache rate is dropped, not inherited")
|
||||
|
||||
ft := table["openai"]["my-private-ft"]
|
||||
assert.InDelta(t, 0.005, ft.CachedInputPer1k, 1e-9, "unknown models are added to the surface")
|
||||
_, ok := table["gemini"]["gemini-pro"]
|
||||
assert.True(t, ok, "a surface the catalog doesn't declare can be added")
|
||||
|
||||
// Untouched entries keep compiled-in rates (catalog, other surface,
|
||||
// supplemental).
|
||||
assert.InDelta(t, 0.00015, table["openai"]["gpt-4o-mini"].InputPer1k, 1e-9, "unlisted model keeps compiled rate")
|
||||
assert.InDelta(t, 0.003, table["anthropic"]["claude-sonnet-4-5"].InputPer1k, 1e-9, "unlisted surface untouched")
|
||||
assert.InDelta(t, 0.00125, table["openai"]["gpt-5"].InputPer1k, 1e-9, "supplemental entries untouched")
|
||||
|
||||
// The synthesizer-facing lookup reads the live table too.
|
||||
e, ok := LookupDefault([]string{"openai"}, "gpt-4o")
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.9, e.InputPer1k, 1e-9, "LookupDefault serves the file-backed rate")
|
||||
}
|
||||
|
||||
func TestLoadFile_MissingPath(t *testing.T) {
|
||||
resetFileState(t)
|
||||
missing := filepath.Join(t.TempDir(), DefaultFileName)
|
||||
|
||||
require.Error(t, LoadFile(missing, true),
|
||||
"explicitly configured path that doesn't exist must fail startup")
|
||||
|
||||
require.NoError(t, LoadFile(missing, false),
|
||||
"conventional datadir probe tolerates an absent file (compiled-in defaults serve)")
|
||||
assert.Nil(t, mergedTable.Load(), "no file, no merged table")
|
||||
fileState.mu.Lock()
|
||||
path := fileState.path
|
||||
fileState.mu.Unlock()
|
||||
assert.Equal(t, missing, path, "the path stays registered so the reloader picks the file up when it appears")
|
||||
}
|
||||
|
||||
func TestLoadFile_RejectsInvalid(t *testing.T) {
|
||||
resetFileState(t)
|
||||
dir := t.TempDir()
|
||||
cases := map[string]string{
|
||||
"unknown field (typo)": "openai:\n gpt-4o:\n input_per1k: 0.1\n",
|
||||
"negative rate": "openai:\n gpt-4o:\n input_per_1k: -0.1\n",
|
||||
"non-numeric rate": "openai:\n gpt-4o:\n input_per_1k: cheap\n",
|
||||
"not a mapping": "- just\n- a\n- list\n",
|
||||
}
|
||||
for name, yml := range cases {
|
||||
path := filepath.Join(dir, DefaultFileName)
|
||||
writePricing(t, path, yml)
|
||||
assert.Error(t, LoadFile(path, true), "case %q must be rejected", name)
|
||||
}
|
||||
}
|
||||
|
||||
// TestReload_LifeCycle drives the reloader's single-shot reload through
|
||||
// its full lifecycle: file edit picked up on mtime change, a broken save
|
||||
// keeps the previous table, and file removal reverts to the compiled-in
|
||||
// defaults (then a re-created file loads again).
|
||||
func TestReload_LifeCycle(t *testing.T) {
|
||||
resetFileState(t)
|
||||
path := filepath.Join(t.TempDir(), DefaultFileName)
|
||||
writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.5\n output_per_1k: 1\n")
|
||||
require.NoError(t, LoadFile(path, true))
|
||||
require.InDelta(t, 0.5, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9)
|
||||
|
||||
// Edit: new mtime, new rates.
|
||||
writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.7\n output_per_1k: 1.4\n")
|
||||
bumpMtime(t, path)
|
||||
reload()
|
||||
assert.InDelta(t, 0.7, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9, "edit must be picked up")
|
||||
|
||||
// Broken save: previous table survives.
|
||||
writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: -1\n")
|
||||
bumpMtime(t, path)
|
||||
reload()
|
||||
assert.InDelta(t, 0.7, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9,
|
||||
"a malformed save must keep the previously loaded table, never blank prices")
|
||||
|
||||
// Removal: compiled-in defaults serve again.
|
||||
require.NoError(t, os.Remove(path))
|
||||
reload()
|
||||
assert.InDelta(t, 0.0025, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9,
|
||||
"file removal reverts to the compiled-in rate")
|
||||
|
||||
// Re-created file loads without a restart.
|
||||
writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.9\n output_per_1k: 1.8\n")
|
||||
bumpMtime(t, path)
|
||||
reload()
|
||||
assert.InDelta(t, 0.9, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9,
|
||||
"a file appearing after removal (or after a missing-probe boot) must load")
|
||||
}
|
||||
|
||||
// bumpMtime pushes the file's mtime forward past the previously recorded
|
||||
// value — timestamps can otherwise collide within the test's timescale.
|
||||
func bumpMtime(t *testing.T, path string) {
|
||||
t.Helper()
|
||||
st, err := os.Stat(path)
|
||||
require.NoError(t, err)
|
||||
next := st.ModTime().Add(2 * 1e9)
|
||||
require.NoError(t, os.Chtimes(path, next, next))
|
||||
}
|
||||
|
||||
// TestExampleYAML_InSyncWithBuiltins is the golden guard for
|
||||
// defaults_llm_pricing.example.yaml: the shipped example must stay
|
||||
// byte-identical to what the compiled-in table renders (catalog edits
|
||||
// require `go generate ./management/internals/modules/agentnetwork/pricing`)
|
||||
// and must round-trip through the same parser operators' files go
|
||||
// through, reproducing the compiled-in table exactly.
|
||||
func TestExampleYAML_InSyncWithBuiltins(t *testing.T) {
|
||||
onDisk, err := os.ReadFile("defaults_llm_pricing.example.yaml")
|
||||
require.NoError(t, err, "example file must exist next to the package")
|
||||
require.Equal(t, string(MarshalDefaultsYAML()), string(onDisk),
|
||||
"defaults_llm_pricing.example.yaml is stale — run: go generate ./management/internals/modules/agentnetwork/pricing")
|
||||
|
||||
parsed, err := parsePricingYAML(onDisk)
|
||||
require.NoError(t, err, "the example must be a valid pricing defaults file")
|
||||
assert.Equal(t, buildDefaultTable(), parsed,
|
||||
"parsing the example must reproduce the compiled-in table exactly")
|
||||
}
|
||||
@@ -233,11 +233,6 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
|
||||
return nil, err
|
||||
}
|
||||
|
||||
costMeterJSON, err := buildCostMeterConfigJSON(enabledProviders, groupIndex)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID)
|
||||
applyAccountCollectionControls(&mergedGuardrails, settings)
|
||||
// The proxy guardrail is a per-provider fail-closed backstop; the
|
||||
@@ -253,7 +248,7 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
|
||||
// Use the merged decision (account settings OR policy-required redaction),
|
||||
// not the raw account flag, so a policy that mandates PII redaction is
|
||||
// honored by the capture parsers even when the account toggle is off.
|
||||
middlewares := buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, costMeterJSON, mergedGuardrails.PromptCapture.RedactPii, mergedGuardrails.PromptCapture.Enabled)
|
||||
middlewares := buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, mergedGuardrails.PromptCapture.RedactPii, mergedGuardrails.PromptCapture.Enabled)
|
||||
|
||||
priv, pub, err := pickServiceSessionKeys(enabledProviders)
|
||||
if err != nil {
|
||||
@@ -705,7 +700,7 @@ func buildIdentityExtraHeaders(p *types.Provider, extras []catalog.ExtraHeader)
|
||||
// requests bound for gateways like LiteLLM that key budgets and
|
||||
// attribution off request headers. CanMutate is required so its
|
||||
// HeadersAdd / HeadersRemove pass the framework's mutation gate.
|
||||
func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, costMeterJSON []byte, redactPii, capturePromptContent bool) []rpservice.MiddlewareConfig {
|
||||
func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byte, redactPii, capturePromptContent bool) []rpservice.MiddlewareConfig {
|
||||
// Both parsers receive an explicit capture flag derived from the account's
|
||||
// enable_prompt_collection toggle; nil/unset would default to the legacy
|
||||
// "always emit" behavior in the middleware, which is precisely what we
|
||||
@@ -774,13 +769,10 @@ func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, cost
|
||||
ConfigJSON: []byte("{}"),
|
||||
},
|
||||
{
|
||||
// Carries the full pricing table (defaults + per-provider
|
||||
// operator prices) so the proxy bills without an embedded
|
||||
// price list; see buildCostMeterConfigJSON.
|
||||
ID: middlewareIDCostMeter,
|
||||
Enabled: true,
|
||||
Slot: rpservice.MiddlewareSlotOnResponse,
|
||||
ConfigJSON: costMeterJSON,
|
||||
ConfigJSON: []byte("{}"),
|
||||
},
|
||||
{
|
||||
ID: middlewareIDLLMResponseParser,
|
||||
|
||||
@@ -1,131 +0,0 @@
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
sharedllm "github.com/netbirdio/netbird/shared/llm"
|
||||
)
|
||||
|
||||
// costMeterConfig is the JSON shape the proxy-side cost_meter middleware
|
||||
// expects (mirror-type pattern, same as routerConfig). The top-level
|
||||
// "pricing" wrapper is the feature-detection signal: an old proxy's config
|
||||
// struct ignores it as an unknown field, and a new proxy treats its
|
||||
// absence as "old management" (skips every cost computation and warns).
|
||||
type costMeterConfig struct {
|
||||
Pricing *costMeterPricing `json:"pricing,omitempty"`
|
||||
}
|
||||
|
||||
// costMeterPricing carries the full pricing table:
|
||||
// - Defaults: surface ("openai"/"anthropic"/"bedrock") -> normalized
|
||||
// model id -> rates. The full default table ships to every account —
|
||||
// it is small (~10 KB) and keeps gateway-style providers (which
|
||||
// enumerate no models) priced for every catalog model.
|
||||
// - Providers: provider record id (matched against the
|
||||
// llm.resolved_provider_id metadata llm_router stamps) -> normalized
|
||||
// model id -> rates. Entries are fully materialized here at synth
|
||||
// time — default cache rates already folded in — so the proxy does
|
||||
// two map lookups and no merging.
|
||||
type costMeterPricing struct {
|
||||
Defaults map[string]map[string]pricing.Entry `json:"defaults,omitempty"`
|
||||
Providers map[string]map[string]pricing.Entry `json:"providers,omitempty"`
|
||||
}
|
||||
|
||||
// buildCostMeterConfigJSON assembles the cost_meter middleware config
|
||||
// from the default pricing table plus the operator's stored per-provider
|
||||
// model prices. Same orphan rule as the router: a provider no enabled
|
||||
// policy authorises is unreachable, so its prices are not shipped.
|
||||
//
|
||||
// Overlay semantics per model row:
|
||||
// - The row's model id is normalized exactly the way the proxy's
|
||||
// request parser normalizes the ids it meters (bedrock ARN/region/
|
||||
// version stripping, vertex "@version" stripping), so the per-record
|
||||
// lookup key compares equal to llm.model at billing time.
|
||||
// - The entry starts from the default entry for that model (when one
|
||||
// exists) to inherit cache rates the operator didn't state.
|
||||
// - Operator input/output overlay verbatim — including an explicit 0,
|
||||
// which prices the model as free (self-hosted / internal endpoints)
|
||||
// rather than silently reverting to list price.
|
||||
// - Cache-rate pointers overlay only when non-nil: nil means "inherit
|
||||
// the default", an explicit 0 means "no discount, bill this bucket
|
||||
// at the input rate".
|
||||
func buildCostMeterConfigJSON(providers []*types.Provider, groupIndex map[string][]string) ([]byte, error) {
|
||||
cfg := costMeterConfig{Pricing: &costMeterPricing{
|
||||
Defaults: pricing.DefaultTable(),
|
||||
}}
|
||||
|
||||
perRecord := make(map[string]map[string]pricing.Entry)
|
||||
for _, p := range providers {
|
||||
if _, hasPolicy := groupIndex[p.ID]; !hasPolicy {
|
||||
// Orphan: unreachable via the router, so unpriceable.
|
||||
continue
|
||||
}
|
||||
if len(p.Models) == 0 {
|
||||
// Gateway-style "claim every model" provider: the defaults
|
||||
// table is its price list.
|
||||
continue
|
||||
}
|
||||
entry, _ := catalog.Lookup(p.ProviderID)
|
||||
models := make(map[string]pricing.Entry, len(p.Models))
|
||||
for _, m := range p.Models {
|
||||
id := normalizePricingModelID(p.ProviderID, m.ID)
|
||||
if id == "" {
|
||||
continue
|
||||
}
|
||||
if _, dup := models[id]; dup {
|
||||
// First occurrence wins on post-normalization duplicates,
|
||||
// matching providerModelIDs' dedup order for routing.
|
||||
continue
|
||||
}
|
||||
models[id] = materializeEntry(entry.PricingSurfaces, id, m)
|
||||
}
|
||||
if len(models) > 0 {
|
||||
perRecord[p.ID] = models
|
||||
}
|
||||
}
|
||||
if len(perRecord) > 0 {
|
||||
cfg.Pricing.Providers = perRecord
|
||||
}
|
||||
|
||||
out, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal cost_meter middleware config: %w", err)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// normalizePricingModelID maps an operator-entered model id onto the
|
||||
// normalized id the proxy's request parser emits as llm.model — the key
|
||||
// the cost meter looks up at billing time.
|
||||
func normalizePricingModelID(catalogProviderID, modelID string) string {
|
||||
switch {
|
||||
case catalog.IsBedrockPathStyle(catalogProviderID):
|
||||
return sharedllm.NormalizeBedrockModel(modelID)
|
||||
case catalog.IsVertexPathStyle(catalogProviderID):
|
||||
return sharedllm.NormalizeVertexModel(modelID)
|
||||
default:
|
||||
return modelID
|
||||
}
|
||||
}
|
||||
|
||||
// materializeEntry folds the default entry for (surfaces, model) — when
|
||||
// one exists — under the operator's stored prices, producing the fully
|
||||
// materialized wire entry.
|
||||
func materializeEntry(surfaces []string, normalizedID string, m types.ProviderModel) pricing.Entry {
|
||||
e, _ := pricing.LookupDefault(surfaces, normalizedID) // zero Entry on miss
|
||||
e.InputPer1k = m.InputPer1k
|
||||
e.OutputPer1k = m.OutputPer1k
|
||||
if m.CachedInputPer1k != nil {
|
||||
e.CachedInputPer1k = *m.CachedInputPer1k
|
||||
}
|
||||
if m.CacheReadPer1k != nil {
|
||||
e.CacheReadPer1k = *m.CacheReadPer1k
|
||||
}
|
||||
if m.CacheCreationPer1k != nil {
|
||||
e.CacheCreationPer1k = *m.CacheCreationPer1k
|
||||
}
|
||||
return e
|
||||
}
|
||||
@@ -1,105 +0,0 @@
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
)
|
||||
|
||||
func fptr(v float64) *float64 { return &v }
|
||||
|
||||
func decodeCostMeterConfig(t *testing.T, raw []byte) costMeterConfig {
|
||||
t.Helper()
|
||||
var cfg costMeterConfig
|
||||
require.NoError(t, json.Unmarshal(raw, &cfg), "cost meter config must round-trip")
|
||||
require.NotNil(t, cfg.Pricing, "pricing wrapper must be present")
|
||||
return cfg
|
||||
}
|
||||
|
||||
// TestBuildCostMeterConfig_BedrockModelNormalization: the operator may
|
||||
// paste region-prefixed, versioned, or ARN-wrapped Bedrock ids; the
|
||||
// per-record map must be keyed by the normalized id the request parser
|
||||
// emits as llm.model, or the lookup never hits at billing time.
|
||||
func TestBuildCostMeterConfig_BedrockModelNormalization(t *testing.T) {
|
||||
bedrock := &types.Provider{
|
||||
ID: "prov-bedrock",
|
||||
ProviderID: "bedrock_api",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{
|
||||
{ID: "eu.anthropic.claude-sonnet-4-5-20250929-v1:0", InputPer1k: 0.0033, OutputPer1k: 0.0165},
|
||||
// Post-normalization duplicate of the row above under a
|
||||
// different regional spelling — first occurrence wins.
|
||||
{ID: "us.anthropic.claude-sonnet-4-5-20250929-v1:0", InputPer1k: 9.9, OutputPer1k: 9.9},
|
||||
},
|
||||
}
|
||||
raw, err := buildCostMeterConfigJSON([]*types.Provider{bedrock}, map[string][]string{"prov-bedrock": {"grp"}})
|
||||
require.NoError(t, err)
|
||||
cfg := decodeCostMeterConfig(t, raw)
|
||||
|
||||
models := cfg.Pricing.Providers["prov-bedrock"]
|
||||
require.Len(t, models, 1, "both spellings normalize to one model; first row wins")
|
||||
e, ok := models["anthropic.claude-sonnet-4-5"]
|
||||
require.True(t, ok, "key must be the normalized id the parser emits, not the operator's raw spelling")
|
||||
assert.InDelta(t, 0.0033, e.InputPer1k, 1e-9, "first row's rate wins the dedup")
|
||||
assert.InDelta(t, 0.0003, e.CacheReadPer1k, 1e-9, "cache read inherited from the bedrock default entry")
|
||||
assert.InDelta(t, 0.00375, e.CacheCreationPer1k, 1e-9, "cache creation inherited from the bedrock default entry")
|
||||
}
|
||||
|
||||
// TestBuildCostMeterConfig_CacheRateNilVsZero pins the pointer semantics:
|
||||
// nil inherits the default cache rate, explicit 0 clears it (that bucket
|
||||
// bills at the input rate on the proxy).
|
||||
func TestBuildCostMeterConfig_CacheRateNilVsZero(t *testing.T) {
|
||||
p := &types.Provider{
|
||||
ID: "prov-oai",
|
||||
ProviderID: "openai_api",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{
|
||||
{ID: "gpt-4o", InputPer1k: 0.002, OutputPer1k: 0.008}, // nil → inherit 0.00125
|
||||
{ID: "gpt-4o-mini", InputPer1k: 0.0001, OutputPer1k: 0.0005, CachedInputPer1k: fptr(0)}, // explicit 0 → no discount
|
||||
{ID: "my-custom-ft", InputPer1k: 0.01, OutputPer1k: 0.02, CachedInputPer1k: fptr(0.005)}, // unknown model, explicit rate
|
||||
},
|
||||
}
|
||||
raw, err := buildCostMeterConfigJSON([]*types.Provider{p}, map[string][]string{"prov-oai": {"grp"}})
|
||||
require.NoError(t, err)
|
||||
cfg := decodeCostMeterConfig(t, raw)
|
||||
models := cfg.Pricing.Providers["prov-oai"]
|
||||
|
||||
assert.InDelta(t, 0.00125, models["gpt-4o"].CachedInputPer1k, 1e-9, "nil cache pointer inherits the default rate")
|
||||
assert.Zero(t, models["gpt-4o-mini"].CachedInputPer1k, "explicit 0 overrides the default (0.000075) — bucket bills at input rate")
|
||||
custom := models["my-custom-ft"]
|
||||
assert.InDelta(t, 0.005, custom.CachedInputPer1k, 1e-9, "unknown model keeps the operator's explicit cache rate")
|
||||
assert.Zero(t, custom.CacheReadPer1k, "no default to inherit for a model outside the catalog")
|
||||
}
|
||||
|
||||
// TestBuildCostMeterConfig_OrphanAndGatewayProviders: an orphan (no
|
||||
// authorising policy) is unreachable so its prices must not ship; a
|
||||
// gateway with no model rows relies on the defaults table and gets no
|
||||
// per-record entry.
|
||||
func TestBuildCostMeterConfig_OrphanAndGatewayProviders(t *testing.T) {
|
||||
orphan := &types.Provider{
|
||||
ID: "prov-orphan",
|
||||
ProviderID: "openai_api",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{{ID: "gpt-4o", InputPer1k: 1, OutputPer1k: 1}},
|
||||
}
|
||||
gateway := &types.Provider{
|
||||
ID: "prov-litellm",
|
||||
ProviderID: "litellm_proxy",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{},
|
||||
}
|
||||
raw, err := buildCostMeterConfigJSON(
|
||||
[]*types.Provider{orphan, gateway},
|
||||
map[string][]string{"prov-litellm": {"grp"}}, // orphan has no policy
|
||||
)
|
||||
require.NoError(t, err)
|
||||
cfg := decodeCostMeterConfig(t, raw)
|
||||
|
||||
assert.NotContains(t, cfg.Pricing.Providers, "prov-orphan", "orphan provider prices must not ship")
|
||||
assert.NotContains(t, cfg.Pricing.Providers, "prov-litellm", "empty-models gateway needs no per-record entry")
|
||||
assert.NotEmpty(t, cfg.Pricing.Defaults["openai"], "defaults still ship so the gateway's catalog-model traffic is priced")
|
||||
}
|
||||
@@ -33,17 +33,14 @@ func newSynthTestSettings() *types.Settings {
|
||||
|
||||
func newSynthTestProvider() *types.Provider {
|
||||
return &types.Provider{
|
||||
ID: "prov-1",
|
||||
AccountID: testAccountID,
|
||||
ProviderID: "openai_api",
|
||||
Name: "OpenAI",
|
||||
UpstreamURL: "https://api.openai.com",
|
||||
APIKey: "sk-test-key",
|
||||
Enabled: true,
|
||||
// Prices deliberately differ from the catalog's gpt-5.4 rates
|
||||
// (0.0025/0.015) so pricing tests can prove the operator's
|
||||
// stored price overlays the catalog default.
|
||||
Models: []types.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.004, OutputPer1k: 0.02}},
|
||||
ID: "prov-1",
|
||||
AccountID: testAccountID,
|
||||
ProviderID: "openai_api",
|
||||
Name: "OpenAI",
|
||||
UpstreamURL: "https://api.openai.com",
|
||||
APIKey: "sk-test-key",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015}},
|
||||
SessionPrivateKey: "test-priv-key",
|
||||
SessionPublicKey: "test-pub-key",
|
||||
CreatedAt: time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC),
|
||||
@@ -217,27 +214,7 @@ func TestSynthesizeServices_HappyPath(t *testing.T) {
|
||||
|
||||
assert.Equal(t, middlewareIDCostMeter, mws[6].ID, "seventh middleware is the cost meter")
|
||||
assert.Equal(t, rpservice.MiddlewareSlotOnResponse, mws[6].Slot, "cost meter runs on_response")
|
||||
|
||||
var costCfg costMeterConfig
|
||||
require.NoError(t, json.Unmarshal(mws[6].ConfigJSON, &costCfg), "cost meter config must unmarshal")
|
||||
require.NotNil(t, costCfg.Pricing, "cost meter config must carry the pricing table — its absence tells the proxy management predates config-delivered pricing")
|
||||
|
||||
gpt4o, ok := costCfg.Pricing.Defaults["openai"]["gpt-4o"]
|
||||
require.True(t, ok, "the full default table ships regardless of the account's providers")
|
||||
assert.InDelta(t, 0.0025, gpt4o.InputPer1k, 1e-9, "default gpt-4o input rate comes from the catalog")
|
||||
|
||||
openaiPrices, ok := costCfg.Pricing.Providers[openai.ID]
|
||||
require.True(t, ok, "operator-priced provider must have a per-record entry")
|
||||
gpt54, ok := openaiPrices["gpt-5.4"]
|
||||
require.True(t, ok, "operator's model row keys the per-record map")
|
||||
assert.InDelta(t, 0.004, gpt54.InputPer1k, 1e-9, "operator input price overlays the catalog default (0.0025)")
|
||||
assert.InDelta(t, 0.02, gpt54.OutputPer1k, 1e-9, "operator output price overlays the catalog default (0.015)")
|
||||
assert.InDelta(t, 0.00025, gpt54.CachedInputPer1k, 1e-9, "cache rate the operator didn't state is inherited from the default entry")
|
||||
|
||||
opus, ok := costCfg.Pricing.Providers[anthropic.ID]["claude-opus-4-7"]
|
||||
require.True(t, ok, "anthropic's model row keys its per-record map")
|
||||
assert.Zero(t, opus.InputPer1k, "operator-stored zero prices ship verbatim — an explicit $0 model bills as free, it does not revert to list price")
|
||||
assert.InDelta(t, 0.0005, opus.CacheReadPer1k, 1e-9, "cache rates still inherit from the default entry")
|
||||
assert.Equal(t, []byte("{}"), mws[6].ConfigJSON, "cost meter carries an explicit empty config")
|
||||
|
||||
assert.Equal(t, middlewareIDLLMResponseParser, mws[7].ID, "eighth middleware is the response parser")
|
||||
assert.Equal(t, rpservice.MiddlewareSlotOnResponse, mws[7].Slot, "response parser runs on_response")
|
||||
|
||||
@@ -14,24 +14,10 @@ import (
|
||||
// ProviderModel is one row in the provider's models list. The operator
|
||||
// pins the per-1k input/output price for cost tracking; ID is the
|
||||
// model identifier the upstream provider expects on the wire.
|
||||
//
|
||||
// The three cache rates are pointers because absence is meaningful: nil
|
||||
// means "inherit NetBird's default rate for this model" (folded in at
|
||||
// synthesis time), while an explicit 0 means "no discount — bill this
|
||||
// cache bucket at the input rate".
|
||||
type ProviderModel struct {
|
||||
ID string `json:"id"`
|
||||
InputPer1k float64 `json:"input_per_1k"`
|
||||
OutputPer1k float64 `json:"output_per_1k"`
|
||||
// CachedInputPer1k is the OpenAI-shape rate for cached prompt tokens
|
||||
// (a subset of input tokens).
|
||||
CachedInputPer1k *float64 `json:"cached_input_per_1k,omitempty"`
|
||||
// CacheReadPer1k is the Anthropic-shape rate for cache-read tokens
|
||||
// (additive to input tokens).
|
||||
CacheReadPer1k *float64 `json:"cache_read_per_1k,omitempty"`
|
||||
// CacheCreationPer1k is the Anthropic-shape rate for cache-creation
|
||||
// tokens (additive to input tokens).
|
||||
CacheCreationPer1k *float64 `json:"cache_creation_per_1k,omitempty"`
|
||||
}
|
||||
|
||||
// Provider is an Agent Network AI provider record persisted per account.
|
||||
@@ -142,12 +128,9 @@ func (p *Provider) FromAPIRequest(req *api.AgentNetworkProviderRequest) {
|
||||
if req.Models != nil {
|
||||
for _, m := range *req.Models {
|
||||
p.Models = append(p.Models, ProviderModel{
|
||||
ID: m.Id,
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
CachedInputPer1k: copyFloatPtr(m.CachedInputPer1k),
|
||||
CacheReadPer1k: copyFloatPtr(m.CacheReadPer1k),
|
||||
CacheCreationPer1k: copyFloatPtr(m.CacheCreationPer1k),
|
||||
ID: m.Id,
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -181,12 +164,9 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider {
|
||||
models := make([]api.AgentNetworkProviderModel, 0, len(p.Models))
|
||||
for _, m := range p.Models {
|
||||
models = append(models, api.AgentNetworkProviderModel{
|
||||
Id: m.ID,
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
CachedInputPer1k: copyFloatPtr(m.CachedInputPer1k),
|
||||
CacheReadPer1k: copyFloatPtr(m.CacheReadPer1k),
|
||||
CacheCreationPer1k: copyFloatPtr(m.CacheCreationPer1k),
|
||||
Id: m.ID,
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
})
|
||||
}
|
||||
created := p.CreatedAt
|
||||
@@ -221,27 +201,11 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider {
|
||||
return resp
|
||||
}
|
||||
|
||||
// copyFloatPtr returns a fresh pointer to the same value, or nil. Keeps
|
||||
// stored models and API payloads from aliasing each other's rate fields.
|
||||
func copyFloatPtr(v *float64) *float64 {
|
||||
if v == nil {
|
||||
return nil
|
||||
}
|
||||
out := *v
|
||||
return &out
|
||||
}
|
||||
|
||||
// Copy returns a deep copy of the provider.
|
||||
func (p *Provider) Copy() *Provider {
|
||||
clone := *p
|
||||
if p.Models != nil {
|
||||
clone.Models = make([]ProviderModel, len(p.Models))
|
||||
for i, m := range p.Models {
|
||||
m.CachedInputPer1k = copyFloatPtr(m.CachedInputPer1k)
|
||||
m.CacheReadPer1k = copyFloatPtr(m.CacheReadPer1k)
|
||||
m.CacheCreationPer1k = copyFloatPtr(m.CacheCreationPer1k)
|
||||
clone.Models[i] = m
|
||||
}
|
||||
clone.Models = append([]ProviderModel(nil), p.Models...)
|
||||
}
|
||||
if p.ExtraValues != nil {
|
||||
clone.ExtraValues = make(map[string]string, len(p.ExtraValues))
|
||||
|
||||
@@ -103,12 +103,6 @@ func TestSynthesizedService_WireShape(t *testing.T) {
|
||||
|
||||
assert.Equal(t, middlewareIDCostMeter, mws[6].GetId(), "seventh middleware id")
|
||||
assert.Equal(t, proto.MiddlewareSlot_MIDDLEWARE_SLOT_ON_RESPONSE, mws[6].GetSlot(), "cost meter slot")
|
||||
var costCfg costMeterConfig
|
||||
require.NoError(t, json.Unmarshal(mws[6].GetConfigJson(), &costCfg), "cost meter config JSON must decode from the wire")
|
||||
require.NotNil(t, costCfg.Pricing, "the pricing table must travel on the wire — the proxy has no embedded price list to fall back to")
|
||||
assert.NotEmpty(t, costCfg.Pricing.Defaults["openai"], "default table rides in every mapping")
|
||||
assert.NotEmpty(t, costCfg.Pricing.Defaults["anthropic"], "default table covers all surfaces")
|
||||
assert.NotEmpty(t, costCfg.Pricing.Defaults["bedrock"], "default table covers all surfaces")
|
||||
|
||||
assert.Equal(t, middlewareIDLLMResponseParser, mws[7].GetId(), "eighth middleware id")
|
||||
assert.Equal(t, proto.MiddlewareSlot_MIDDLEWARE_SLOT_ON_RESPONSE, mws[7].GetSlot(), "response parser slot")
|
||||
|
||||
@@ -55,8 +55,6 @@ type Config struct {
|
||||
|
||||
ReverseProxy ReverseProxy
|
||||
|
||||
AgentNetwork AgentNetwork
|
||||
|
||||
// disable default all-to-all policy
|
||||
DisableDefaultPolicy bool
|
||||
|
||||
@@ -187,25 +185,6 @@ type StoreConfig struct {
|
||||
Engine types.Engine
|
||||
}
|
||||
|
||||
// AgentNetwork contains agent-network (LLM gateway) configuration.
|
||||
type AgentNetwork struct {
|
||||
// PricingDefaultsFile is the path to the YAML file holding the default
|
||||
// LLM pricing table (defaults_llm_pricing.yaml). A relative path is
|
||||
// resolved against <Datadir>, so a bare filename lands alongside the
|
||||
// store. Empty falls back to probing <Datadir>/defaults_llm_pricing.yaml;
|
||||
// with no file present the compiled-in defaults serve. Schema: surface ("openai"/"anthropic"/
|
||||
// "bedrock") -> model -> rates in USD per 1k tokens (input_per_1k,
|
||||
// output_per_1k, and the optional cached_input_per_1k /
|
||||
// cache_read_per_1k / cache_creation_per_1k). File entries replace the
|
||||
// compiled-in entry for the same surface+model whole; everything else
|
||||
// keeps the compiled-in rates. The file is re-read periodically (mtime
|
||||
// poll), and the live table feeds both the synthesizer (what proxies
|
||||
// bill with) and the dashboard's catalog endpoint (what model rows
|
||||
// prefill with). An explicitly configured path that fails to load
|
||||
// fails startup; runtime reload errors keep the previous table.
|
||||
PricingDefaultsFile string
|
||||
}
|
||||
|
||||
// ReverseProxy contains reverse proxy configuration in front of management.
|
||||
type ReverseProxy struct {
|
||||
// TrustedHTTPProxies represents a list of trusted HTTP proxies by their IP prefixes.
|
||||
|
||||
38
proxy/internal/llm/bedrock_model.go
Normal file
38
proxy/internal/llm/bedrock_model.go
Normal file
@@ -0,0 +1,38 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// bedrockRegionPrefixes are the cross-region inference-profile prefixes that
|
||||
// front a Bedrock model id (e.g. "eu.anthropic.claude-...").
|
||||
var bedrockRegionPrefixes = []string{"us.", "eu.", "apac.", "global."}
|
||||
|
||||
// bedrockVersionSuffix matches the trailing "-vN[:N]" or "-YYYYMMDD-vN[:N]"
|
||||
// version/throughput suffix of a Bedrock model id.
|
||||
var bedrockVersionSuffix = regexp.MustCompile(`-(\d{8}-)?v\d+(:\d+)?$`)
|
||||
|
||||
// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile
|
||||
// prefix, and the version/throughput suffix from a Bedrock model id so it
|
||||
// matches the catalog/pricing key, e.g.
|
||||
// "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -> "anthropic.claude-sonnet-4-5"
|
||||
// and the inference-profile ARN's last segment likewise. It is the single
|
||||
// source of truth shared by the request parser (which normalizes the request
|
||||
// model from the URL path) and the router (which normalizes the operator's
|
||||
// registered Bedrock model ids so both sides compare equal).
|
||||
func NormalizeBedrockModel(modelID string) string {
|
||||
m := modelID
|
||||
if strings.HasPrefix(m, "arn:") {
|
||||
if i := strings.LastIndex(m, "/"); i >= 0 {
|
||||
m = m[i+1:]
|
||||
}
|
||||
}
|
||||
for _, p := range bedrockRegionPrefixes {
|
||||
if strings.HasPrefix(m, p) {
|
||||
m = m[len(p):]
|
||||
break
|
||||
}
|
||||
}
|
||||
return bedrockVersionSuffix.ReplaceAllString(m, "")
|
||||
}
|
||||
@@ -11,8 +11,6 @@ func TestNormalizeBedrockModel(t *testing.T) {
|
||||
"eu.anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5",
|
||||
"us.anthropic.claude-haiku-4-5": "anthropic.claude-haiku-4-5",
|
||||
"us.anthropic.claude-opus-4-8-20250101-v1:0": "anthropic.claude-opus-4-8",
|
||||
"apac.anthropic.claude-haiku-4-5-v1:0": "anthropic.claude-haiku-4-5",
|
||||
"amazon.nova-2-lite-v1:0": "amazon.nova-2-lite",
|
||||
"anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5",
|
||||
"meta.llama3-3-70b-instruct-v1:0": "meta.llama3-3-70b-instruct",
|
||||
"amazon.nova-pro-v1:0": "amazon.nova-pro",
|
||||
@@ -23,14 +21,3 @@ func TestNormalizeBedrockModel(t *testing.T) {
|
||||
require.Equal(t, want, NormalizeBedrockModel(in), "normalize %q", in)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeVertexModel(t *testing.T) {
|
||||
cases := map[string]string{
|
||||
"claude-sonnet-4-5@20250929": "claude-sonnet-4-5",
|
||||
"claude-haiku-4-5": "claude-haiku-4-5",
|
||||
"gpt-4o@2024-08-06": "gpt-4o",
|
||||
}
|
||||
for in, want := range cases {
|
||||
require.Equal(t, want, NormalizeVertexModel(in), "normalize %q", in)
|
||||
}
|
||||
}
|
||||
59
proxy/internal/llm/fixtures/pricing.yaml
Normal file
59
proxy/internal/llm/fixtures/pricing.yaml
Normal file
@@ -0,0 +1,59 @@
|
||||
# Realistic-pricing starter for llm_observability. Drop this into the
|
||||
# directory you point the proxy at via --plugin-data-dir, then reference it
|
||||
# from the target's plugin config:
|
||||
#
|
||||
# plugins:
|
||||
# - id: llm_observability
|
||||
# enabled: true
|
||||
# params:
|
||||
# pricing_path: pricing.yaml
|
||||
#
|
||||
# Values are USD per 1_000 tokens. Public list prices drift; treat this as a
|
||||
# starting point and keep your production copy current.
|
||||
|
||||
openai:
|
||||
# GPT-5 family
|
||||
gpt-5:
|
||||
input_per_1k: 0.00125
|
||||
output_per_1k: 0.01
|
||||
gpt-5-mini:
|
||||
input_per_1k: 0.00025
|
||||
output_per_1k: 0.002
|
||||
gpt-5-nano:
|
||||
input_per_1k: 0.00005
|
||||
output_per_1k: 0.0004
|
||||
gpt-5.4:
|
||||
input_per_1k: 0.00125
|
||||
output_per_1k: 0.01
|
||||
# GPT-4o family
|
||||
gpt-4o:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.01
|
||||
gpt-4o-mini:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.0006
|
||||
# Embeddings
|
||||
text-embedding-3-large:
|
||||
input_per_1k: 0.00013
|
||||
output_per_1k: 0
|
||||
text-embedding-3-small:
|
||||
input_per_1k: 0.00002
|
||||
output_per_1k: 0
|
||||
|
||||
anthropic:
|
||||
# Claude 4.x family
|
||||
claude-opus-4-7:
|
||||
input_per_1k: 0.015
|
||||
output_per_1k: 0.075
|
||||
claude-sonnet-4-7:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
claude-sonnet-4-6:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
claude-haiku-4-5:
|
||||
input_per_1k: 0.0008
|
||||
output_per_1k: 0.004
|
||||
@@ -1,21 +0,0 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
sharedllm "github.com/netbirdio/netbird/shared/llm"
|
||||
)
|
||||
|
||||
// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile
|
||||
// prefix, and the version/throughput suffix from a Bedrock model id so it
|
||||
// matches the catalog/pricing key. Thin delegate to the shared implementation
|
||||
// (shared/llm), which management also uses at synthesis time so both sides of
|
||||
// the pricing / routing contract normalize identically.
|
||||
func NormalizeBedrockModel(modelID string) string {
|
||||
return sharedllm.NormalizeBedrockModel(modelID)
|
||||
}
|
||||
|
||||
// NormalizeVertexModel strips the "@version" suffix from a Vertex AI model id
|
||||
// so it matches the catalog/pricing key. Thin delegate to shared/llm, kept
|
||||
// beside NormalizeBedrockModel for the same contract reason.
|
||||
func NormalizeVertexModel(modelID string) string {
|
||||
return sharedllm.NormalizeVertexModel(modelID)
|
||||
}
|
||||
65
proxy/internal/llm/pricing/defaults_coverage_test.go
Normal file
65
proxy/internal/llm/pricing/defaults_coverage_test.go
Normal file
@@ -0,0 +1,65 @@
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestDefaultTable_FirstPartyModelCoverage guards the embedded defaults against
|
||||
// silent drift/gaps: every metered first-party model the management catalog
|
||||
// enumerates must resolve to a price, and a few rates that previously drifted
|
||||
// are pinned to their LiteLLM-validated values. Keep this list in step with the
|
||||
// catalog (management/server/agentnetwork/catalog) when adding models.
|
||||
func TestDefaultTable_FirstPartyModelCoverage(t *testing.T) {
|
||||
tbl := DefaultTable()
|
||||
require.NotNil(t, tbl, "embedded default pricing table must load")
|
||||
|
||||
mustPrice := map[string][]string{
|
||||
// openai parser covers openai_api, azure_openai_api, and mistral_api.
|
||||
"openai": {
|
||||
"gpt-5.5", "gpt-5.5-pro", "gpt-5.4", "gpt-5.4-mini", "gpt-5.4-nano",
|
||||
"gpt-5.3-codex", "gpt-5.3-chat-latest", "o4-mini",
|
||||
"gpt-4.1", "gpt-4.1-mini", "gpt-4.1-nano", "gpt-4o", "gpt-4o-mini",
|
||||
"gpt-4-turbo", "gpt-3.5-turbo", "gpt-35-turbo",
|
||||
"text-embedding-3-large", "text-embedding-3-small",
|
||||
"mistral-large-latest", "mistral-medium-3-5", "codestral-2508",
|
||||
"ministral-8b-latest", "mistral-embed",
|
||||
},
|
||||
"anthropic": {
|
||||
"claude-fable-5", "claude-opus-5", "claude-opus-4-8", "claude-opus-4-7", "claude-opus-4-6",
|
||||
"claude-opus-4-1", "claude-sonnet-4-6", "claude-sonnet-4-5", "claude-haiku-4-5",
|
||||
},
|
||||
// bedrock keys are the normalized ids the request parser emits.
|
||||
"bedrock": {
|
||||
"anthropic.claude-opus-5", "anthropic.claude-opus-4-8", "anthropic.claude-opus-4-7", "anthropic.claude-opus-4-6",
|
||||
"anthropic.claude-opus-4-1", "anthropic.claude-sonnet-4-6", "anthropic.claude-sonnet-4-5",
|
||||
"anthropic.claude-haiku-4-5", "meta.llama3-3-70b-instruct",
|
||||
"amazon.nova-pro", "amazon.nova-lite", "amazon.nova-micro", "amazon.nova-2-lite",
|
||||
},
|
||||
}
|
||||
for provider, models := range mustPrice {
|
||||
for _, m := range models {
|
||||
_, ok := tbl.Cost(provider, m, 1000, 1000, 0, 0)
|
||||
assert.True(t, ok, "%s/%s must be priced in the embedded defaults", provider, m)
|
||||
}
|
||||
}
|
||||
|
||||
// Pin per-direction rates independently (input-only then output-only) so a
|
||||
// swap or skew of input<->output that preserves the combined total is still
|
||||
// caught — these are rates that previously drifted or are easy to mis-enter.
|
||||
in, ok := tbl.Cost("openai", "gpt-5.4", 1000, 0, 0, 0)
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.0025, in, 1e-9, "gpt-5.4 input = 0.0025 per 1k")
|
||||
out, ok := tbl.Cost("openai", "gpt-5.4", 0, 1000, 0, 0)
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.015, out, 1e-9, "gpt-5.4 output = 0.015 per 1k")
|
||||
|
||||
in, ok = tbl.Cost("bedrock", "anthropic.claude-sonnet-4-5", 1000, 0, 0, 0)
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.003, in, 1e-9, "bedrock sonnet-4-5 input = 0.003 per 1k")
|
||||
out, ok = tbl.Cost("bedrock", "anthropic.claude-sonnet-4-5", 0, 1000, 0, 0)
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.015, out, 1e-9, "bedrock sonnet-4-5 output = 0.015 per 1k")
|
||||
}
|
||||
@@ -1,196 +1,66 @@
|
||||
# Default LLM pricing used by NetBird's Agent Network cost metering.
|
||||
# GENERATED from the management catalog — do not edit this file in the
|
||||
# repository; regenerate with:
|
||||
# Embedded default pricing for llm_observability. Compiled into the proxy
|
||||
# binary via go:embed in pricing.go; cost annotation works out of the box
|
||||
# without any operator action.
|
||||
#
|
||||
# go generate ./management/internals/modules/agentnetwork/pricing
|
||||
# Operators override entries by dropping a pricing.yaml into --plugin-data-dir
|
||||
# (or whichever basename is given via params.pricing_path). The override file
|
||||
# only needs entries the operator wants to change; missing entries fall
|
||||
# through to these defaults.
|
||||
#
|
||||
# Operators: copy this file to <datadir>/defaults_llm_pricing.yaml (or
|
||||
# any path configured via management.json:
|
||||
# Values are USD per 1_000 tokens. Public list prices drift; ship a fresh
|
||||
# binary or override individual entries via the override file as needed.
|
||||
#
|
||||
# { "AgentNetwork": { "PricingDefaultsFile": "/path/defaults_llm_pricing.yaml" } }
|
||||
#
|
||||
# ) and adjust the entries you want to change. Management re-reads the
|
||||
# file periodically (mtime poll, every minute): the live table feeds the
|
||||
# proxies' cost metering and the dashboard's model-price prefill, so
|
||||
# edits apply without a restart. Your file only needs the entries you
|
||||
# want to change — but each entry REPLACES the built-in entry for that
|
||||
# surface+model whole, so repeat the cache rates you want to keep.
|
||||
# Unknown fields and negative or non-finite rates are rejected: at
|
||||
# startup that fails boot (for an explicitly configured path); at
|
||||
# runtime the previous table is kept and a warning is logged. Deleting
|
||||
# the file reverts to the built-in defaults below.
|
||||
#
|
||||
# Top-level keys are pricing surfaces — the parser shape requests are
|
||||
# metered under: "openai" (also Azure, Mistral, and OpenAI-compatible
|
||||
# gateways), "anthropic" (also Anthropic-on-Vertex), "bedrock"
|
||||
# (normalized ids, e.g. anthropic.claude-sonnet-4-5). Model keys must be
|
||||
# the normalized id the proxy meters (version/region suffixes stripped).
|
||||
#
|
||||
# Values are USD per 1_000 tokens. Optional cache fields:
|
||||
# cached_input_per_1k OpenAI shape: rate for cached prompt tokens
|
||||
# (a SUBSET of input tokens). Absent -> cached
|
||||
# portion bills at input_per_1k.
|
||||
# cache_read_per_1k Anthropic shape: rate for cache_read tokens
|
||||
# (ADDITIVE to input). Absent -> input rate.
|
||||
# cache_creation_per_1k Anthropic shape: rate for cache_creation
|
||||
# tokens (ADDITIVE to input). Absent -> input
|
||||
# rate.
|
||||
|
||||
anthropic:
|
||||
claude-fable-5:
|
||||
input_per_1k: 0.01
|
||||
output_per_1k: 0.05
|
||||
cache_read_per_1k: 0.001
|
||||
cache_creation_per_1k: 0.0125
|
||||
claude-haiku-4-5:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.005
|
||||
cache_read_per_1k: 0.0001
|
||||
cache_creation_per_1k: 0.00125
|
||||
claude-opus-4-1:
|
||||
input_per_1k: 0.015
|
||||
output_per_1k: 0.075
|
||||
cache_read_per_1k: 0.0015
|
||||
cache_creation_per_1k: 0.01875
|
||||
claude-opus-4-6:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-7:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-8:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
claude-sonnet-4-6:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
kimi-k3:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cached_input_per_1k: 0.0003
|
||||
cache_read_per_1k: 0.0003
|
||||
"kimi-k3[1m]":
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
|
||||
bedrock:
|
||||
amazon.nova-2-lite:
|
||||
input_per_1k: 0.0003
|
||||
output_per_1k: 0.0025
|
||||
amazon.nova-lite:
|
||||
input_per_1k: 0.00006
|
||||
output_per_1k: 0.00024
|
||||
amazon.nova-micro:
|
||||
input_per_1k: 0.000035
|
||||
output_per_1k: 0.00014
|
||||
amazon.nova-pro:
|
||||
input_per_1k: 0.0008
|
||||
output_per_1k: 0.0032
|
||||
anthropic.claude-haiku-4-5:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.005
|
||||
cache_read_per_1k: 0.0001
|
||||
cache_creation_per_1k: 0.00125
|
||||
anthropic.claude-opus-4-1:
|
||||
input_per_1k: 0.015
|
||||
output_per_1k: 0.075
|
||||
cache_read_per_1k: 0.0015
|
||||
cache_creation_per_1k: 0.01875
|
||||
anthropic.claude-opus-4-6:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-7:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-8:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
anthropic.claude-sonnet-4-6:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
meta.llama3-3-70b-instruct:
|
||||
input_per_1k: 0.00072
|
||||
output_per_1k: 0.00072
|
||||
# Optional cache fields:
|
||||
# cached_input_per_1k OpenAI: rate for prompt_tokens_details.cached_tokens
|
||||
# (a SUBSET of prompt_tokens). Typically 0.5x input.
|
||||
# Absent → cached portion bills at input_per_1k.
|
||||
# cache_read_per_1k Anthropic: rate for cache_read_input_tokens
|
||||
# (ADDITIVE to input_tokens). Typically 0.1x input.
|
||||
# Absent → cache reads bill at input_per_1k.
|
||||
# cache_creation_per_1k Anthropic: rate for cache_creation_input_tokens
|
||||
# (ADDITIVE to input_tokens). Typically 1.25x input.
|
||||
# Absent → cache writes bill at input_per_1k.
|
||||
|
||||
openai:
|
||||
codestral-2508:
|
||||
input_per_1k: 0.0003
|
||||
output_per_1k: 0.0009
|
||||
codestral-latest:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.003
|
||||
devstral-medium-latest:
|
||||
input_per_1k: 0.0004
|
||||
output_per_1k: 0.002
|
||||
devstral-small-latest:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0.0003
|
||||
gpt-3.5-turbo:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
gpt-35-turbo:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
gpt-4-turbo:
|
||||
input_per_1k: 0.01
|
||||
# OpenAI + OpenAI-compatible providers (openai_api, azure_openai_api,
|
||||
# mistral_api, and the openai-parser gateways) all emit llm.provider="openai",
|
||||
# so their models are priced here. Kept in sync with the management catalog;
|
||||
# rates cross-checked against LiteLLM model_prices_and_context_window.json.
|
||||
|
||||
# GPT-5.x family — cache reads 10% of input (0.1x).
|
||||
gpt-5.5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.03
|
||||
gpt-4.1:
|
||||
input_per_1k: 0.002
|
||||
output_per_1k: 0.008
|
||||
cached_input_per_1k: 0.0005
|
||||
gpt-4.1-mini:
|
||||
input_per_1k: 0.0004
|
||||
output_per_1k: 0.0016
|
||||
cached_input_per_1k: 0.0001
|
||||
gpt-4.1-nano:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0.0004
|
||||
cached_input_per_1k: 0.000025
|
||||
gpt-4o:
|
||||
gpt-5.5-pro:
|
||||
input_per_1k: 0.03
|
||||
output_per_1k: 0.18
|
||||
cached_input_per_1k: 0.003
|
||||
gpt-5.4:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.01
|
||||
cached_input_per_1k: 0.00125
|
||||
gpt-4o-mini:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.0006
|
||||
output_per_1k: 0.015
|
||||
cached_input_per_1k: 0.00025
|
||||
gpt-5.4-pro:
|
||||
input_per_1k: 0.03
|
||||
output_per_1k: 0.18
|
||||
cached_input_per_1k: 0.003
|
||||
gpt-5.4-mini:
|
||||
input_per_1k: 0.00075
|
||||
output_per_1k: 0.0045
|
||||
cached_input_per_1k: 0.000075
|
||||
gpt-5.4-nano:
|
||||
input_per_1k: 0.0002
|
||||
output_per_1k: 0.00125
|
||||
cached_input_per_1k: 0.00002
|
||||
gpt-5.3-codex:
|
||||
input_per_1k: 0.00175
|
||||
output_per_1k: 0.014
|
||||
cached_input_per_1k: 0.000175
|
||||
gpt-5.3-chat-latest:
|
||||
input_per_1k: 0.00175
|
||||
output_per_1k: 0.014
|
||||
cached_input_per_1k: 0.000175
|
||||
# GPT-5 (2025) family — kept for gateway requests using the unsuffixed ids.
|
||||
gpt-5:
|
||||
input_per_1k: 0.00125
|
||||
output_per_1k: 0.01
|
||||
@@ -203,80 +73,226 @@ openai:
|
||||
input_per_1k: 0.00005
|
||||
output_per_1k: 0.0004
|
||||
cached_input_per_1k: 0.000005
|
||||
gpt-5.3-chat-latest:
|
||||
input_per_1k: 0.00175
|
||||
output_per_1k: 0.014
|
||||
cached_input_per_1k: 0.000175
|
||||
gpt-5.3-codex:
|
||||
input_per_1k: 0.00175
|
||||
output_per_1k: 0.014
|
||||
cached_input_per_1k: 0.000175
|
||||
gpt-5.4:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.015
|
||||
cached_input_per_1k: 0.00025
|
||||
gpt-5.4-mini:
|
||||
input_per_1k: 0.00075
|
||||
output_per_1k: 0.0045
|
||||
cached_input_per_1k: 0.000075
|
||||
gpt-5.4-nano:
|
||||
input_per_1k: 0.0002
|
||||
output_per_1k: 0.00125
|
||||
cached_input_per_1k: 0.00002
|
||||
gpt-5.4-pro:
|
||||
input_per_1k: 0.03
|
||||
output_per_1k: 0.18
|
||||
cached_input_per_1k: 0.003
|
||||
gpt-5.5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.03
|
||||
cached_input_per_1k: 0.0005
|
||||
gpt-5.5-pro:
|
||||
input_per_1k: 0.03
|
||||
output_per_1k: 0.18
|
||||
cached_input_per_1k: 0.003
|
||||
kimi-k3:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cached_input_per_1k: 0.0003
|
||||
cache_read_per_1k: 0.0003
|
||||
magistral-medium-latest:
|
||||
input_per_1k: 0.002
|
||||
output_per_1k: 0.005
|
||||
magistral-small-latest:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
ministral-3-14b-2512:
|
||||
input_per_1k: 0.0002
|
||||
output_per_1k: 0.0002
|
||||
ministral-3-3b-2512:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0.0001
|
||||
ministral-8b-latest:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.00015
|
||||
mistral-embed:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0
|
||||
mistral-large-latest:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
mistral-medium-3-5:
|
||||
input_per_1k: 0.0015
|
||||
output_per_1k: 0.0075
|
||||
mistral-medium-latest:
|
||||
input_per_1k: 0.0004
|
||||
output_per_1k: 0.002
|
||||
mistral-small-latest:
|
||||
input_per_1k: 0.00006
|
||||
output_per_1k: 0.00018
|
||||
o4-mini:
|
||||
input_per_1k: 0.0011
|
||||
output_per_1k: 0.0044
|
||||
cached_input_per_1k: 0.000275
|
||||
# GPT-4.1 family — cache reads 25% of input.
|
||||
gpt-4.1:
|
||||
input_per_1k: 0.002
|
||||
output_per_1k: 0.008
|
||||
cached_input_per_1k: 0.0005
|
||||
gpt-4.1-mini:
|
||||
input_per_1k: 0.0004
|
||||
output_per_1k: 0.0016
|
||||
cached_input_per_1k: 0.0001
|
||||
gpt-4.1-nano:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0.0004
|
||||
cached_input_per_1k: 0.000025
|
||||
# GPT-4o family — cache reads 50% of input (0.5x).
|
||||
gpt-4o:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.01
|
||||
cached_input_per_1k: 0.00125
|
||||
gpt-4o-mini:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.0006
|
||||
cached_input_per_1k: 0.000075
|
||||
# Older GPT — no prompt caching.
|
||||
gpt-4-turbo:
|
||||
input_per_1k: 0.01
|
||||
output_per_1k: 0.03
|
||||
gpt-3.5-turbo:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
gpt-35-turbo: # Azure deployment alias of gpt-3.5-turbo
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
# Embeddings — no caching, no output tokens.
|
||||
text-embedding-3-large:
|
||||
input_per_1k: 0.00013
|
||||
output_per_1k: 0
|
||||
text-embedding-3-small:
|
||||
input_per_1k: 0.00002
|
||||
output_per_1k: 0
|
||||
|
||||
# Mistral (mistral_api) — routed via the openai parser; no prompt caching.
|
||||
mistral-large-latest:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
mistral-medium-latest:
|
||||
input_per_1k: 0.0004
|
||||
output_per_1k: 0.002
|
||||
mistral-medium-3-5:
|
||||
input_per_1k: 0.0015
|
||||
output_per_1k: 0.0075
|
||||
mistral-small-latest:
|
||||
input_per_1k: 0.00006
|
||||
output_per_1k: 0.00018
|
||||
magistral-medium-latest:
|
||||
input_per_1k: 0.002
|
||||
output_per_1k: 0.005
|
||||
magistral-small-latest:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
devstral-medium-latest:
|
||||
input_per_1k: 0.0004
|
||||
output_per_1k: 0.002
|
||||
devstral-small-latest:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0.0003
|
||||
codestral-2508:
|
||||
input_per_1k: 0.0003
|
||||
output_per_1k: 0.0009
|
||||
codestral-latest:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.003
|
||||
ministral-3-14b-2512:
|
||||
input_per_1k: 0.0002
|
||||
output_per_1k: 0.0002
|
||||
ministral-8b-latest:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.00015
|
||||
ministral-3-3b-2512:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0.0001
|
||||
mistral-embed:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0
|
||||
|
||||
# Kimi / Moonshot AI (kimi_api) — OpenAI-compatible /v1 endpoint. Moonshot
|
||||
# reports cache hits OpenAI-style when present; cached input is 10% of
|
||||
# input ($0.30 vs $3.00 per MTok). kimi-k3 is the only model the platform
|
||||
# serves newer accounts (K2-era ids and kimi-latest 404), matching the
|
||||
# management catalog.
|
||||
kimi-k3:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cached_input_per_1k: 0.0003
|
||||
|
||||
anthropic:
|
||||
# Claude 4.x family — cache reads ≈10% of input, cache writes ≈125% of input.
|
||||
# Pricing source: Anthropic's current published rates per million tokens,
|
||||
# divided by 1000 for the per-1k figures stored here.
|
||||
claude-fable-5:
|
||||
input_per_1k: 0.010
|
||||
output_per_1k: 0.050
|
||||
cache_read_per_1k: 0.001
|
||||
cache_creation_per_1k: 0.0125
|
||||
claude-opus-5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-8:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-7:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-6:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-1:
|
||||
input_per_1k: 0.015
|
||||
output_per_1k: 0.075
|
||||
cache_read_per_1k: 0.0015
|
||||
cache_creation_per_1k: 0.01875
|
||||
claude-sonnet-4-6:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
claude-haiku-4-5:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.005
|
||||
cache_read_per_1k: 0.0001
|
||||
cache_creation_per_1k: 0.00125
|
||||
|
||||
# Kimi / Moonshot AI (kimi_api) via the Anthropic-compatible endpoint
|
||||
# (/anthropic/v1/messages — the official Claude Code setup). Same rates
|
||||
# as the OpenAI-shape entry above. "kimi-k3[1m]" is the model id some
|
||||
# Claude Code guides set for the 1M-context alias; priced identically so
|
||||
# cost metering doesn't silently skip those requests.
|
||||
kimi-k3:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
"kimi-k3[1m]":
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
|
||||
bedrock:
|
||||
# AWS Bedrock model ids, normalised by the request parser (cross-region
|
||||
# inference-profile prefix + version/throughput suffix stripped), e.g.
|
||||
# eu.anthropic.claude-sonnet-4-5-20250929-v1:0 -> anthropic.claude-sonnet-4-5.
|
||||
# Anthropic-on-Bedrock keeps the additive cache buckets (read ≈0.1x input,
|
||||
# write ≈1.25x input); Nova / Llama report no cache, so cost is input+output.
|
||||
anthropic.claude-opus-5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-8:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-7:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-6:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-1:
|
||||
input_per_1k: 0.015
|
||||
output_per_1k: 0.075
|
||||
cache_read_per_1k: 0.0015
|
||||
cache_creation_per_1k: 0.01875
|
||||
anthropic.claude-sonnet-4-6:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
anthropic.claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
anthropic.claude-haiku-4-5:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.005
|
||||
cache_read_per_1k: 0.0001
|
||||
cache_creation_per_1k: 0.00125
|
||||
meta.llama3-3-70b-instruct:
|
||||
input_per_1k: 0.00072
|
||||
output_per_1k: 0.00072
|
||||
amazon.nova-2-lite:
|
||||
input_per_1k: 0.0003
|
||||
output_per_1k: 0.0025
|
||||
amazon.nova-pro:
|
||||
input_per_1k: 0.0008
|
||||
output_per_1k: 0.0032
|
||||
amazon.nova-lite:
|
||||
input_per_1k: 0.00006
|
||||
output_per_1k: 0.00024
|
||||
amazon.nova-micro:
|
||||
input_per_1k: 0.000035
|
||||
output_per_1k: 0.00014
|
||||
@@ -1,30 +1,102 @@
|
||||
// Package pricing implements the pricing table and cost formula the
|
||||
// cost_meter middleware uses to convert LLM token usage into a USD cost
|
||||
// estimate. The table's content arrives from the management server inside
|
||||
// cost_meter's middleware config (synthesized from the catalog plus the
|
||||
// operator's stored per-provider prices) — the proxy carries no embedded
|
||||
// price list. Price updates ride the ordinary mapping push: a chain
|
||||
// rebuild constructs a fresh table, so there is nothing to reload.
|
||||
// Package pricing implements the embedded-default + override pricing table
|
||||
// shared by middleware that converts LLM token usage into a USD cost
|
||||
// estimate. The table is hot-reloadable from a basename under the proxy
|
||||
// data directory; missing override files keep the embedded defaults so
|
||||
// cost annotation works without operator action.
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
_ "embed"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"math"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"go.opentelemetry.io/otel/attribute"
|
||||
"go.opentelemetry.io/otel/metric"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
//go:embed defaults_pricing.yaml
|
||||
var defaultPricingYAML []byte
|
||||
|
||||
var (
|
||||
defaultTableOnce sync.Once
|
||||
defaultTablePtr *Table
|
||||
)
|
||||
|
||||
// DefaultTable returns the pricing table embedded in the binary. The result
|
||||
// is parsed once and shared; callers must not mutate the returned value.
|
||||
// Cost annotation works without any operator action because every loader
|
||||
// starts with this table.
|
||||
func DefaultTable() *Table {
|
||||
defaultTableOnce.Do(func() {
|
||||
t, err := parsePricingBytes(defaultPricingYAML)
|
||||
if err != nil {
|
||||
panic(fmt.Sprintf("llmobs: embedded default pricing failed to parse: %v", err))
|
||||
}
|
||||
defaultTablePtr = t
|
||||
})
|
||||
return defaultTablePtr
|
||||
}
|
||||
|
||||
// mergeOver returns a new Table containing every entry from base, with any
|
||||
// matching entry from overlay replacing the base value. Either argument may
|
||||
// be nil. Result is a fresh allocation so callers can mutate / Store safely.
|
||||
func mergeOver(base, overlay *Table) *Table {
|
||||
if overlay == nil || len(overlay.entries) == 0 {
|
||||
return base
|
||||
}
|
||||
if base == nil || len(base.entries) == 0 {
|
||||
return overlay
|
||||
}
|
||||
out := make(map[string]map[string]Entry, len(base.entries))
|
||||
for provider, models := range base.entries {
|
||||
inner := make(map[string]Entry, len(models))
|
||||
for model, e := range models {
|
||||
inner[model] = e
|
||||
}
|
||||
out[provider] = inner
|
||||
}
|
||||
for provider, models := range overlay.entries {
|
||||
inner, ok := out[provider]
|
||||
if !ok {
|
||||
inner = make(map[string]Entry, len(models))
|
||||
out[provider] = inner
|
||||
}
|
||||
for model, e := range models {
|
||||
inner[model] = e
|
||||
}
|
||||
}
|
||||
return &Table{entries: out}
|
||||
}
|
||||
|
||||
// Entry is a single model's input and output pricing, expressed in USD per
|
||||
// 1000 tokens.
|
||||
//
|
||||
// CachedInputPer1K applies to OpenAI's cached prompt tokens, which are a
|
||||
// subset of input_tokens — when set, the cached portion is billed at this
|
||||
// rate and the non-cached remainder at InputPer1K. Zero means "no discount
|
||||
// configured", and cached tokens are billed at InputPer1K.
|
||||
// configured", and cached tokens are billed at InputPer1K (matches current
|
||||
// behaviour where cached counts weren't extracted at all).
|
||||
//
|
||||
// CacheReadPer1K and CacheCreationPer1K apply to Anthropic's two prompt-
|
||||
// cache fields, which are additive to input_tokens: cache_read is the
|
||||
// cheaper read-from-cache rate, cache_creation is the more expensive
|
||||
// write-to-cache rate. Zero means "no rate configured" and the
|
||||
// corresponding token bucket is billed at InputPer1K.
|
||||
// corresponding token bucket is billed at InputPer1K. This is more
|
||||
// accurate than today's behaviour, where Anthropic's cache tokens are
|
||||
// ignored and not charged at all.
|
||||
type Entry struct {
|
||||
InputPer1K float64
|
||||
OutputPer1K float64
|
||||
@@ -33,102 +105,33 @@ type Entry struct {
|
||||
CacheCreationPer1K float64
|
||||
}
|
||||
|
||||
// EntryJSON is the wire shape of a pricing entry inside cost_meter's
|
||||
// middleware config. Field names are the management→proxy contract; the
|
||||
// management synthesizer marshals the same names (its pricing.Entry).
|
||||
type EntryJSON struct {
|
||||
InputPer1K float64 `json:"input_per_1k"`
|
||||
OutputPer1K float64 `json:"output_per_1k"`
|
||||
CachedInputPer1K float64 `json:"cached_input_per_1k"`
|
||||
CacheReadPer1K float64 `json:"cache_read_per_1k"`
|
||||
CacheCreationPer1K float64 `json:"cache_creation_per_1k"`
|
||||
}
|
||||
|
||||
// Table is a provider-surface-to-model pricing lookup. Instances are
|
||||
// immutable once built; a mapping update builds a whole new middleware
|
||||
// instance (and with it a new table) rather than mutating this one.
|
||||
// Table is a provider-to-model pricing lookup. Instances are immutable once
|
||||
// built and are swapped atomically by Loader.
|
||||
type Table struct {
|
||||
entries map[string]map[string]Entry
|
||||
}
|
||||
|
||||
// NewEntries validates and converts a wire-shape map (surface-or-record ->
|
||||
// model -> rates) into the internal representation. Every rate must be a
|
||||
// finite, non-negative USD amount; a violation is returned as an error so
|
||||
// a corrupt config fails the chain build loudly instead of mispricing.
|
||||
// Management validates the same constraints at its API boundary, so this
|
||||
// is defense-in-depth. Nil input yields an empty (never-matching) map.
|
||||
func NewEntries(raw map[string]map[string]EntryJSON) (map[string]map[string]Entry, error) {
|
||||
out := make(map[string]map[string]Entry, len(raw))
|
||||
for outer, models := range raw {
|
||||
inner := make(map[string]Entry, len(models))
|
||||
for model, e := range models {
|
||||
for field, v := range map[string]float64{
|
||||
"input_per_1k": e.InputPer1K,
|
||||
"output_per_1k": e.OutputPer1K,
|
||||
"cached_input_per_1k": e.CachedInputPer1K,
|
||||
"cache_read_per_1k": e.CacheReadPer1K,
|
||||
"cache_creation_per_1k": e.CacheCreationPer1K,
|
||||
} {
|
||||
if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) {
|
||||
return nil, fmt.Errorf("pricing %s/%s: %s must be a finite, non-negative rate, got %v", outer, model, field, v)
|
||||
}
|
||||
}
|
||||
// EntryJSON and Entry are field-identical (tags aside), so a
|
||||
// direct conversion carries all five rates.
|
||||
inner[model] = Entry(e)
|
||||
}
|
||||
out[outer] = inner
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// NewTable builds an immutable Table from the wire-shape defaults map.
|
||||
// See NewEntries for validation semantics.
|
||||
func NewTable(raw map[string]map[string]EntryJSON) (*Table, error) {
|
||||
entries, err := NewEntries(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Table{entries: entries}, nil
|
||||
}
|
||||
|
||||
// Lookup returns the entry for the given provider surface and model.
|
||||
func (t *Table) Lookup(provider, model string) (Entry, bool) {
|
||||
if t == nil {
|
||||
return Entry{}, false
|
||||
}
|
||||
byModel, ok := t.entries[provider]
|
||||
if !ok {
|
||||
return Entry{}, false
|
||||
}
|
||||
e, ok := byModel[model]
|
||||
return e, ok
|
||||
}
|
||||
|
||||
// Has reports whether the provider/model pair is present in the table.
|
||||
func (t *Table) Has(provider, model string) bool {
|
||||
_, ok := t.Lookup(provider, model)
|
||||
return ok
|
||||
}
|
||||
|
||||
// Cost returns the estimated USD cost for the given token counts. ok is
|
||||
// false when the provider or model is not present in the table; the caller
|
||||
// can still emit token metrics with a model=unknown label.
|
||||
//
|
||||
// Provider-shape semantics for cached / cache-creation counts:
|
||||
//
|
||||
// - OpenAI: cachedInput is a SUBSET of inTokens. The cached portion is
|
||||
// billed at CachedInputPer1K (or InputPer1K when no override), and the
|
||||
// non-cached remainder of inTokens at InputPer1K. cacheCreation is
|
||||
// ignored (OpenAI has no analogue).
|
||||
// - Anthropic: cachedInput (cache_read) and cacheCreation are ADDITIVE to
|
||||
// inTokens. The three buckets are billed at CacheReadPer1K,
|
||||
// CacheCreationPer1K, and InputPer1K respectively, each falling back
|
||||
// to InputPer1K when the corresponding rate is zero.
|
||||
// - Other providers: cached and cacheCreation are ignored; cost is
|
||||
// inTokens*InputPer1K + outTokens*OutputPer1K.
|
||||
func (t *Table) Cost(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (float64, bool) {
|
||||
c, ok := t.Costs(provider, model, inTokens, outTokens, cachedInput, cacheCreation)
|
||||
return c.TotalUSD, ok
|
||||
}
|
||||
|
||||
// Costs returns the estimated USD cost split for the given token counts.
|
||||
// The provider surface selects the cache formula; see EntryCosts.
|
||||
func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (Costs, bool) {
|
||||
entry, ok := t.Lookup(provider, model)
|
||||
if !ok {
|
||||
return Costs{}, false
|
||||
}
|
||||
return EntryCosts(entry, provider, inTokens, outTokens, cachedInput, cacheCreation), true
|
||||
}
|
||||
|
||||
// Costs is a per-request cost split. The four per-bucket fields are the base
|
||||
// of the breakdown — one per token bucket the provider bills separately — and
|
||||
// the two aggregates are derived from them:
|
||||
@@ -162,25 +165,9 @@ func newCosts(input, cachedInput, cacheCreation, output float64) Costs {
|
||||
}
|
||||
}
|
||||
|
||||
// EntryCosts computes the USD cost split for the given entry and token
|
||||
// counts. The surface (the llm.provider value the request parser stamps)
|
||||
// selects the cache formula; the entry may come from the surface-keyed
|
||||
// defaults table or from a per-provider-record override — the math is
|
||||
// identical either way.
|
||||
//
|
||||
// Provider-shape semantics for cached / cache-creation counts:
|
||||
//
|
||||
// - "openai": cachedInput is a SUBSET of inTokens. The cached portion is
|
||||
// billed at CachedInputPer1K (or InputPer1K when no override), and the
|
||||
// non-cached remainder of inTokens at InputPer1K. cacheCreation is
|
||||
// ignored (OpenAI has no analogue).
|
||||
// - "anthropic", "bedrock": cachedInput (cache_read) and cacheCreation are
|
||||
// ADDITIVE to inTokens. The three buckets are billed at CacheReadPer1K,
|
||||
// CacheCreationPer1K, and InputPer1K respectively, each falling back
|
||||
// to InputPer1K when the corresponding rate is zero.
|
||||
// - Other surfaces: cached and cacheCreation are ignored; cost is
|
||||
// inTokens*InputPer1K + outTokens*OutputPer1K.
|
||||
func EntryCosts(entry Entry, surface string, inTokens, outTokens, cachedInput, cacheCreation int64) Costs {
|
||||
// Costs returns the estimated USD cost split for the given token counts, with
|
||||
// the same semantics as Cost.
|
||||
func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (Costs, bool) {
|
||||
// Clamp negatives to zero before any pricing math so a malformed
|
||||
// upstream count can never produce a negative cost.
|
||||
if inTokens < 0 {
|
||||
@@ -195,8 +182,19 @@ func EntryCosts(entry Entry, surface string, inTokens, outTokens, cachedInput, c
|
||||
if cacheCreation < 0 {
|
||||
cacheCreation = 0
|
||||
}
|
||||
if t == nil {
|
||||
return Costs{}, false
|
||||
}
|
||||
byModel, ok := t.entries[provider]
|
||||
if !ok {
|
||||
return Costs{}, false
|
||||
}
|
||||
entry, ok := byModel[model]
|
||||
if !ok {
|
||||
return Costs{}, false
|
||||
}
|
||||
output := (float64(outTokens) / 1000.0) * entry.OutputPer1K
|
||||
switch surface {
|
||||
switch provider {
|
||||
case "openai":
|
||||
// cachedInput is a subset of inTokens; clamp so a malformed
|
||||
// upstream (cached > total) can't produce a negative remainder.
|
||||
@@ -210,7 +208,7 @@ func EntryCosts(entry Entry, surface string, inTokens, outTokens, cachedInput, c
|
||||
}
|
||||
nonCached := float64(inTokens-clamped) / 1000.0 * entry.InputPer1K
|
||||
cached := float64(clamped) / 1000.0 * cachedRate
|
||||
return newCosts(nonCached, cached, 0, output)
|
||||
return newCosts(nonCached, cached, 0, output), true
|
||||
case "anthropic", "bedrock":
|
||||
// Bedrock-Anthropic returns the same additive cache buckets as
|
||||
// first-party Anthropic; non-Anthropic Bedrock models simply report
|
||||
@@ -226,9 +224,266 @@ func EntryCosts(entry Entry, surface string, inTokens, outTokens, cachedInput, c
|
||||
input := float64(inTokens) / 1000.0 * entry.InputPer1K
|
||||
read := float64(cachedInput) / 1000.0 * readRate
|
||||
create := float64(cacheCreation) / 1000.0 * createRate
|
||||
return newCosts(input, read, create, output)
|
||||
return newCosts(input, read, create, output), true
|
||||
default:
|
||||
input := float64(inTokens) / 1000.0 * entry.InputPer1K
|
||||
return newCosts(input, 0, 0, output)
|
||||
return newCosts(input, 0, 0, output), true
|
||||
}
|
||||
}
|
||||
|
||||
// Has reports whether the provider/model pair is present in the table.
|
||||
func (t *Table) Has(provider, model string) bool {
|
||||
if t == nil {
|
||||
return false
|
||||
}
|
||||
byModel, ok := t.entries[provider]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
_, ok = byModel[model]
|
||||
return ok
|
||||
}
|
||||
|
||||
// pricingFile mirrors the on-disk YAML schema. Keys are provider names; the
|
||||
// nested map keys are model names.
|
||||
type pricingFile map[string]map[string]struct {
|
||||
InputPer1K float64 `yaml:"input_per_1k"`
|
||||
OutputPer1K float64 `yaml:"output_per_1k"`
|
||||
CachedInputPer1K float64 `yaml:"cached_input_per_1k"`
|
||||
CacheReadPer1K float64 `yaml:"cache_read_per_1k"`
|
||||
CacheCreationPer1K float64 `yaml:"cache_creation_per_1k"`
|
||||
}
|
||||
|
||||
const (
|
||||
// ReloadInterval is the mtime-poll cadence for the background reloader.
|
||||
ReloadInterval = 30 * time.Second
|
||||
|
||||
// errorBackoff bounds how often the loader logs a repeated parse error.
|
||||
errorBackoff = 5 * time.Minute
|
||||
)
|
||||
|
||||
var basenameRegex = regexp.MustCompile(`^[a-zA-Z0-9._-]+$`)
|
||||
|
||||
// Loader is a confined, hot-reloadable pricing table reader. Construction
|
||||
// must succeed against the target file; subsequent reload failures keep the
|
||||
// previously-loaded table so callers never observe a blank price list.
|
||||
type Loader struct {
|
||||
baseDir string
|
||||
fullPath string
|
||||
pluginID string
|
||||
table atomic.Pointer[Table]
|
||||
mtime atomic.Int64
|
||||
failures metric.Int64Counter
|
||||
interval time.Duration
|
||||
}
|
||||
|
||||
// NewLoader returns a pricing loader that overlays an optional file-based
|
||||
// table on top of the embedded defaults. Missing override file, baseDir, or
|
||||
// relPath is not an error: the loader keeps the embedded defaults so cost
|
||||
// metadata is still emitted for known models.
|
||||
//
|
||||
// Errors:
|
||||
// - bad basename, traversal segment, or absolute relPath are rejected so a
|
||||
// misconfigured target surfaces immediately.
|
||||
// - permission errors and YAML parse errors keep the defaults but log a
|
||||
// warning; cost annotation does not silently break.
|
||||
//
|
||||
// failures is optional; pass nil in tests that do not care about
|
||||
// reload-failure telemetry.
|
||||
func NewLoader(baseDir, relPath, pluginID string, failures metric.Int64Counter) (*Loader, error) {
|
||||
defaults := DefaultTable()
|
||||
l := &Loader{
|
||||
baseDir: baseDir,
|
||||
pluginID: pluginID,
|
||||
failures: failures,
|
||||
}
|
||||
l.table.Store(defaults)
|
||||
|
||||
if strings.TrimSpace(baseDir) == "" || strings.TrimSpace(relPath) == "" {
|
||||
return l, nil
|
||||
}
|
||||
|
||||
full, err := resolveMiddlewareDataPath(baseDir, relPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
l.fullPath = full
|
||||
|
||||
overlay, mtime, err := loadPricing(full)
|
||||
if err != nil {
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
// Override file is optional. Defaults already stored.
|
||||
return l, nil
|
||||
}
|
||||
// Symlink rejection, oversize file, parse failure, permission errors
|
||||
// — surface so a misconfigured operator sees the problem instead of
|
||||
// silently running with stale defaults.
|
||||
return nil, fmt.Errorf("load pricing %s: %w", full, err)
|
||||
}
|
||||
l.table.Store(mergeOver(defaults, overlay))
|
||||
l.mtime.Store(mtime.UnixNano())
|
||||
return l, nil
|
||||
}
|
||||
|
||||
// Get returns the current pricing table. The returned pointer is immutable;
|
||||
// callers must not mutate its contents.
|
||||
func (l *Loader) Get() *Table {
|
||||
if l == nil {
|
||||
return nil
|
||||
}
|
||||
return l.table.Load()
|
||||
}
|
||||
|
||||
// WatchesFile reports whether this loader is bound to an override file on
|
||||
// disk. False for defaults-only loaders (no operator override given).
|
||||
// Callers use this to decide whether to spawn the mtime-poll goroutine.
|
||||
func (l *Loader) WatchesFile() bool {
|
||||
if l == nil {
|
||||
return false
|
||||
}
|
||||
return l.fullPath != ""
|
||||
}
|
||||
|
||||
// SetReloadInterval overrides the mtime-poll cadence used by Reload. Calls
|
||||
// after Reload has started have no effect on the running loop. Intended for
|
||||
// tests; production code uses the default ReloadInterval.
|
||||
func (l *Loader) SetReloadInterval(d time.Duration) {
|
||||
if l == nil || d <= 0 {
|
||||
return
|
||||
}
|
||||
l.interval = d
|
||||
}
|
||||
|
||||
// Reload runs a polling loop that checks the pricing file mtime every
|
||||
// ReloadInterval (or the value passed to SetReloadInterval). Returns when
|
||||
// ctx is cancelled.
|
||||
func (l *Loader) Reload(ctx context.Context) {
|
||||
if l == nil {
|
||||
return
|
||||
}
|
||||
interval := l.interval
|
||||
if interval <= 0 {
|
||||
interval = ReloadInterval
|
||||
}
|
||||
t := time.NewTicker(interval)
|
||||
defer t.Stop()
|
||||
|
||||
var lastErrAt time.Time
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-t.C:
|
||||
if err := l.reload(); err != nil {
|
||||
if l.failures != nil {
|
||||
l.failures.Add(ctx, 1, metric.WithAttributes(
|
||||
attribute.String("plugin", l.pluginID),
|
||||
))
|
||||
}
|
||||
now := time.Now()
|
||||
if now.Sub(lastErrAt) >= errorBackoff {
|
||||
log.Warnf("llmobs: pricing reload failed for %s: %v", l.fullPath, err)
|
||||
lastErrAt = now
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// reload performs a single-shot mtime check and reload. The reloaded
|
||||
// override file is merged on top of the embedded defaults; missing override
|
||||
// (e.g. operator deleted the file) is not an error and reverts to defaults.
|
||||
func (l *Loader) reload() error {
|
||||
if l.fullPath == "" {
|
||||
// Defaults-only loader; nothing on disk to reload.
|
||||
return nil
|
||||
}
|
||||
mtime, err := statMtime(l.fullPath)
|
||||
if err != nil {
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
// File was removed since startup. Drop back to defaults and
|
||||
// reset mtime so a future re-creation triggers a reload.
|
||||
l.table.Store(DefaultTable())
|
||||
l.mtime.Store(0)
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
if mtime.UnixNano() == l.mtime.Load() {
|
||||
return nil
|
||||
}
|
||||
|
||||
overlay, newMtime, err := loadPricing(l.fullPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
l.table.Store(mergeOver(DefaultTable(), overlay))
|
||||
l.mtime.Store(newMtime.UnixNano())
|
||||
return nil
|
||||
}
|
||||
|
||||
// resolveMiddlewareDataPath validates relPath is a safe basename and resolves
|
||||
// it under baseDir. An additional cleaned-prefix check guards against
|
||||
// CVE-style edge cases where Join is used with trailing path segments.
|
||||
func resolveMiddlewareDataPath(baseDir, relPath string) (string, error) {
|
||||
if strings.TrimSpace(baseDir) == "" {
|
||||
return "", errors.New("middleware-data-dir is not configured")
|
||||
}
|
||||
if relPath == "" {
|
||||
return "", errors.New("pricing path is empty")
|
||||
}
|
||||
if !basenameRegex.MatchString(relPath) {
|
||||
return "", fmt.Errorf("pricing path %q is not a safe basename", relPath)
|
||||
}
|
||||
if filepath.IsAbs(relPath) {
|
||||
return "", fmt.Errorf("pricing path %q must be a basename, not absolute", relPath)
|
||||
}
|
||||
|
||||
cleanBase, err := filepath.Abs(filepath.Clean(baseDir))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("resolve middleware-data-dir: %w", err)
|
||||
}
|
||||
full := filepath.Join(cleanBase, relPath)
|
||||
cleanedFull := filepath.Clean(full)
|
||||
if !strings.HasPrefix(cleanedFull, cleanBase+string(filepath.Separator)) && cleanedFull != cleanBase {
|
||||
return "", fmt.Errorf("pricing path %q escapes middleware-data-dir", relPath)
|
||||
}
|
||||
return cleanedFull, nil
|
||||
}
|
||||
|
||||
func parsePricingBytes(data []byte) (*Table, error) {
|
||||
dec := yaml.NewDecoder(bytes.NewReader(data))
|
||||
dec.KnownFields(true)
|
||||
|
||||
var raw pricingFile
|
||||
if err := dec.Decode(&raw); err != nil && !errors.Is(err, io.EOF) {
|
||||
return nil, fmt.Errorf("decode pricing yaml: %w", err)
|
||||
}
|
||||
|
||||
out := make(map[string]map[string]Entry, len(raw))
|
||||
for provider, models := range raw {
|
||||
inner := make(map[string]Entry, len(models))
|
||||
for model, entry := range models {
|
||||
for field, v := range map[string]float64{
|
||||
"input_per_1k": entry.InputPer1K,
|
||||
"output_per_1k": entry.OutputPer1K,
|
||||
"cached_input_per_1k": entry.CachedInputPer1K,
|
||||
"cache_read_per_1k": entry.CacheReadPer1K,
|
||||
"cache_creation_per_1k": entry.CacheCreationPer1K,
|
||||
} {
|
||||
if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) {
|
||||
return nil, fmt.Errorf("pricing %s/%s: %s must be a finite, non-negative rate, got %v", provider, model, field, v)
|
||||
}
|
||||
}
|
||||
inner[model] = Entry{
|
||||
InputPer1K: entry.InputPer1K,
|
||||
OutputPer1K: entry.OutputPer1K,
|
||||
CachedInputPer1K: entry.CachedInputPer1K,
|
||||
CacheReadPer1K: entry.CacheReadPer1K,
|
||||
CacheCreationPer1K: entry.CacheCreationPer1K,
|
||||
}
|
||||
}
|
||||
out[provider] = inner
|
||||
}
|
||||
return &Table{entries: out}, nil
|
||||
}
|
||||
|
||||
20
proxy/internal/llm/pricing/pricing_other.go
Normal file
20
proxy/internal/llm/pricing/pricing_other.go
Normal file
@@ -0,0 +1,20 @@
|
||||
//go:build !unix
|
||||
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// loadPricing is unavailable on non-Unix platforms because O_NOFOLLOW and
|
||||
// fstat-from-FD are required to honour the spec's symlink-safety rules. The
|
||||
// proxy is only deployed on Linux today; a Windows port would need an
|
||||
// equivalent path-as-handle implementation.
|
||||
func loadPricing(path string) (*Table, time.Time, error) {
|
||||
return nil, time.Time{}, fmt.Errorf("llmobs pricing loader is not supported on this platform: %s", path)
|
||||
}
|
||||
|
||||
func statMtime(path string) (time.Time, error) {
|
||||
return time.Time{}, fmt.Errorf("llmobs pricing loader is not supported on this platform: %s", path)
|
||||
}
|
||||
@@ -1,13 +1,47 @@
|
||||
//go:build unix
|
||||
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"math"
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func copyFixture(t *testing.T, src, dst string) {
|
||||
t.Helper()
|
||||
data, err := os.ReadFile(src)
|
||||
require.NoError(t, err, "read source fixture")
|
||||
require.NoError(t, os.WriteFile(dst, data, 0o600), "write target fixture")
|
||||
}
|
||||
|
||||
func TestNewLoader_HappyPath(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml"))
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err, "NewLoader must succeed with a valid fixture")
|
||||
table := l.Get()
|
||||
require.NotNil(t, table, "table populated after load")
|
||||
|
||||
cost, ok := table.Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "known provider/model resolves")
|
||||
assert.InDelta(t, 0.00075, cost, 1e-9, "cost = 0.00015 + 0.0006 per 1k tokens")
|
||||
|
||||
cost, ok = table.Cost("openai", "gpt-4o", 2000, 1000, 0, 0)
|
||||
require.True(t, ok, "second known model resolves")
|
||||
assert.InDelta(t, 0.015, cost, 1e-9, "cost for gpt-4o: 2*0.0025 + 1*0.01")
|
||||
|
||||
cost, ok = table.Cost("anthropic", "claude-sonnet-4-5", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "anthropic model resolves")
|
||||
assert.InDelta(t, 0.018, cost, 1e-9, "cost for claude-sonnet-4-5: 0.003 + 0.015")
|
||||
}
|
||||
|
||||
// TestCost_OpenAICachedSubsetDiscount proves OpenAI's cached input
|
||||
// tokens are billed at the configured cached_input_per_1k rate while
|
||||
// the non-cached remainder of input_tokens is billed at the regular
|
||||
@@ -31,9 +65,11 @@ func TestCost_OpenAICachedSubsetDiscount(t *testing.T) {
|
||||
"cached subset must bill at the discount rate; non-cached remainder at regular rate")
|
||||
}
|
||||
|
||||
// TestCost_OpenAICachedFallsBackToInputRate covers the fallback
|
||||
// contract: when CachedInputPer1K is unset (zero), cached tokens bill
|
||||
// at the regular input rate.
|
||||
// TestCost_OpenAICachedFallsBackToInputRate covers the operator
|
||||
// opt-in contract: when CachedInputPer1K is unset (zero), cached
|
||||
// tokens bill at the regular input rate. This matches today's
|
||||
// behaviour (cached counts weren't extracted at all so they
|
||||
// implicitly billed at the input rate via prompt_tokens).
|
||||
func TestCost_OpenAICachedFallsBackToInputRate(t *testing.T) {
|
||||
tbl := &Table{entries: map[string]map[string]Entry{
|
||||
"openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}},
|
||||
@@ -42,7 +78,7 @@ func TestCost_OpenAICachedFallsBackToInputRate(t *testing.T) {
|
||||
require.True(t, ok)
|
||||
want := 0.0025 + (500.0/1000.0)*0.01
|
||||
assert.InDelta(t, want, cost, 1e-12,
|
||||
"absent cached_input_per_1k rate must fall back to input_per_1k")
|
||||
"absent cached_input_per_1k rate must fall back to input_per_1k — same as pre-feature behaviour")
|
||||
}
|
||||
|
||||
// TestCost_OpenAIClampsCachedToInputCount is the defensive guard
|
||||
@@ -64,7 +100,10 @@ func TestCost_OpenAIClampsCachedToInputCount(t *testing.T) {
|
||||
// TestCost_AnthropicCacheReadAndCreationAreAdditive proves the
|
||||
// Anthropic shape: cache_read and cache_creation tokens are
|
||||
// ADDITIVE to input_tokens (not subset), each billed at its own
|
||||
// configured rate.
|
||||
// configured rate. The two rates pull in opposite directions —
|
||||
// cache_read is the cheaper read-from-cache rate (≈0.1× input),
|
||||
// cache_creation is the more expensive write-to-cache rate
|
||||
// (≈1.25× input).
|
||||
func TestCost_AnthropicCacheReadAndCreationAreAdditive(t *testing.T) {
|
||||
tbl := &Table{entries: map[string]map[string]Entry{
|
||||
"anthropic": {"claude-sonnet": {
|
||||
@@ -86,9 +125,11 @@ func TestCost_AnthropicCacheReadAndCreationAreAdditive(t *testing.T) {
|
||||
"each Anthropic input bucket must bill at its own configured rate")
|
||||
}
|
||||
|
||||
// TestCost_AnthropicCacheRatesFallBackToInput covers the no-rate
|
||||
// TestCost_AnthropicCacheRatesFallBackToInput covers the no-opt-in
|
||||
// path: when neither CacheReadPer1K nor CacheCreationPer1K is set,
|
||||
// cache tokens bill at the regular input rate.
|
||||
// cache tokens bill at the regular input rate. This is more
|
||||
// accurate than today's behaviour (cache tokens ignored entirely)
|
||||
// without requiring operators to opt in via YAML.
|
||||
func TestCost_AnthropicCacheRatesFallBackToInput(t *testing.T) {
|
||||
tbl := &Table{entries: map[string]map[string]Entry{
|
||||
"anthropic": {"claude-sonnet": {InputPer1K: 0.003, OutputPer1K: 0.015}},
|
||||
@@ -98,39 +139,259 @@ func TestCost_AnthropicCacheRatesFallBackToInput(t *testing.T) {
|
||||
// Without overrides: every input bucket at input_per_1k.
|
||||
want := ((256.0+768.0+512.0)/1000.0)*0.003 + (200.0/1000.0)*0.015
|
||||
assert.InDelta(t, want, cost, 1e-12,
|
||||
"absent cache rates must fall back to input_per_1k")
|
||||
"absent cache rates must fall back to input_per_1k — Anthropic cache tokens were ignored before this change, billing at input rate is more accurate as a default")
|
||||
}
|
||||
|
||||
// TestEntryCosts_SurfaceSelectsFormula pins that the formula branches on
|
||||
// the SURFACE, not on which table the entry came from: the same entry
|
||||
// bills a subset carve-out on "openai", additive buckets on
|
||||
// "anthropic"/"bedrock", and ignores cache counts everywhere else. This
|
||||
// is what keeps per-provider-record entries (looked up by record id)
|
||||
// mathematically identical to defaults-table entries.
|
||||
func TestEntryCosts_SurfaceSelectsFormula(t *testing.T) {
|
||||
e := Entry{InputPer1K: 0.002, OutputPer1K: 0.01, CachedInputPer1K: 0.001, CacheReadPer1K: 0.0002, CacheCreationPer1K: 0.0025}
|
||||
func TestNewLoader_UnknownModel(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml"))
|
||||
|
||||
openai := EntryCosts(e, "openai", 1000, 0, 400, 300)
|
||||
assert.InDelta(t, (600.0/1000.0)*0.002+(400.0/1000.0)*0.001, openai.TotalUSD, 1e-12,
|
||||
"openai: cached is a subset, cacheCreation ignored")
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
anthropic := EntryCosts(e, "anthropic", 1000, 0, 400, 300)
|
||||
assert.InDelta(t, 0.002+(400.0/1000.0)*0.0002+(300.0/1000.0)*0.0025, anthropic.TotalUSD, 1e-12,
|
||||
"anthropic: cache buckets are additive")
|
||||
_, ok := l.Get().Cost("openai", "fantasy-model", 10, 10, 0, 0)
|
||||
assert.False(t, ok, "unknown model returns ok=false")
|
||||
|
||||
bedrock := EntryCosts(e, "bedrock", 1000, 0, 400, 300)
|
||||
assert.InDelta(t, anthropic.TotalUSD, bedrock.TotalUSD, 1e-12, "bedrock shares the anthropic formula")
|
||||
|
||||
other := EntryCosts(e, "gemini", 1000, 0, 400, 300)
|
||||
assert.InDelta(t, 0.002, other.TotalUSD, 1e-12, "unknown surface: cache counts ignored")
|
||||
_, ok = l.Get().Cost("cohere", "anything", 10, 10, 0, 0)
|
||||
assert.False(t, ok, "unknown provider returns ok=false")
|
||||
}
|
||||
|
||||
// TestEntryCosts_ClampsNegativeTokens: malformed upstream counts must
|
||||
// never produce a negative cost.
|
||||
func TestEntryCosts_ClampsNegativeTokens(t *testing.T) {
|
||||
e := Entry{InputPer1K: 0.002, OutputPer1K: 0.01}
|
||||
c := EntryCosts(e, "openai", -50, -10, -5, -3)
|
||||
assert.Zero(t, c.TotalUSD, "all-negative counts clamp to zero cost")
|
||||
func TestNewLoader_InvalidYAMLRejected(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
require.NoError(t, os.WriteFile(filepath.Join(base, "pricing.yaml"), []byte("\t- this is not: valid: yaml: :["), 0o600))
|
||||
|
||||
_, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.Error(t, err, "invalid YAML must surface as construction error")
|
||||
}
|
||||
|
||||
func TestLoader_ReloadKeepsPreviousOnParseError(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
target := filepath.Join(base, "pricing.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target)
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, l.Get(), "initial table populated")
|
||||
|
||||
// Overwrite with content that violates the strict schema (extra field)
|
||||
// plus a bumped mtime to trigger reload.
|
||||
require.NoError(t, os.WriteFile(target, []byte("openai:\n gpt-4o:\n input_per_1k: 1.0\n output_per_1k: 2.0\n bogus_field: nope\n"), 0o600))
|
||||
future := time.Now().Add(time.Hour)
|
||||
require.NoError(t, os.Chtimes(target, future, future))
|
||||
|
||||
err = l.reload()
|
||||
require.Error(t, err, "parse error surfaced by reload()")
|
||||
|
||||
cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "previous table still available after parse failure")
|
||||
assert.InDelta(t, 0.00075, cost, 1e-9, "previous cost preserved")
|
||||
}
|
||||
|
||||
func TestLoader_ReloadNoChangeIsNoOp(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
target := filepath.Join(base, "pricing.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target)
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
ptrBefore := l.Get()
|
||||
|
||||
require.NoError(t, l.reload(), "no-change reload must not error")
|
||||
ptrAfter := l.Get()
|
||||
assert.Same(t, ptrBefore, ptrAfter, "table pointer unchanged when mtime unchanged")
|
||||
}
|
||||
|
||||
func TestLoader_ReloadDetectsChange(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
target := filepath.Join(base, "pricing.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target)
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
updated := []byte("openai:\n gpt-4o-mini:\n input_per_1k: 1.00\n output_per_1k: 2.00\n")
|
||||
require.NoError(t, os.WriteFile(target, updated, 0o600))
|
||||
future := time.Now().Add(time.Hour)
|
||||
require.NoError(t, os.Chtimes(target, future, future))
|
||||
|
||||
require.NoError(t, l.reload(), "reload must succeed on valid new content")
|
||||
|
||||
cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "updated model still present")
|
||||
assert.InDelta(t, 3.0, cost, 0.0001, "new prices are applied: 1 + 2 per 1k")
|
||||
}
|
||||
|
||||
// TestLoader_ReloadGoroutinePicksUpChanges proves the background goroutine
|
||||
// started via Reload actually swaps the pricing table when the file changes
|
||||
// on disk. Without that goroutine running, pricing edits would never reach
|
||||
// requests until a proxy restart.
|
||||
func TestLoader_ReloadGoroutinePicksUpChanges(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
target := filepath.Join(base, "pricing.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target)
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
l.SetReloadInterval(20 * time.Millisecond)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
l.Reload(ctx)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Before any rewrite, the loader holds the fixture's prices.
|
||||
costBefore, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "fixture model must resolve initially")
|
||||
assert.InDelta(t, 0.00075, costBefore, 1e-9, "fixture prices apply before rewrite")
|
||||
|
||||
updated := []byte("openai:\n gpt-4o-mini:\n input_per_1k: 1.00\n output_per_1k: 2.00\n")
|
||||
require.NoError(t, os.WriteFile(target, updated, 0o600))
|
||||
future := time.Now().Add(time.Hour)
|
||||
require.NoError(t, os.Chtimes(target, future, future))
|
||||
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for {
|
||||
cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
if ok && cost > 2.5 {
|
||||
break
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatalf("background reloader did not pick up rewrite within deadline")
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("Reload loop did not exit after cancel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoader_ReloadBackgroundLoopCancellation(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml"))
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
l.Reload(ctx)
|
||||
close(done)
|
||||
}()
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("Reload loop did not exit on context cancel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewLoader_PathValidation(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml"))
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
relPath string
|
||||
}{
|
||||
{"traversal", "../../etc/passwd"},
|
||||
{"absolute", "/etc/passwd"},
|
||||
{"slash in basename", "sub/pricing.yaml"},
|
||||
{"control chars", "pricing\x00.yaml"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := NewLoader(base, tc.relPath, "llm_observability", nil)
|
||||
require.Error(t, err, "NewLoader must reject %q", tc.relPath)
|
||||
})
|
||||
}
|
||||
|
||||
// Empty relPath is no longer a validation error: the loader treats it
|
||||
// as "no override file, defaults only" so cost metadata is still
|
||||
// emitted for the embedded models out of the box.
|
||||
t.Run("empty falls back to defaults", func(t *testing.T) {
|
||||
l, err := NewLoader(base, "", "llm_observability", nil)
|
||||
require.NoError(t, err, "empty relPath should yield a defaults-only loader")
|
||||
require.NotNil(t, l, "loader must be returned")
|
||||
require.False(t, l.WatchesFile(), "no file watching when no override is given")
|
||||
_, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
assert.True(t, ok, "embedded defaults should still resolve gpt-4o-mini")
|
||||
})
|
||||
}
|
||||
|
||||
// TestNewLoader_PathValidation_Extended covers the remaining attack shapes
|
||||
// called out in C2: dot references, embedded traversal segments, and a
|
||||
// newline in the basename. The basename regex must reject each one even
|
||||
// though filepath.Clean would otherwise collapse them.
|
||||
func TestNewLoader_PathValidation_Extended(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml"))
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
relPath string
|
||||
}{
|
||||
{"dot", "."},
|
||||
{"dotdot", ".."},
|
||||
{"relative traversal", "../pricing.yaml"},
|
||||
{"embedded slash", "pri/cing.yaml"},
|
||||
{"newline", "pricing\n.yaml"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := NewLoader(base, tc.relPath, "llm_observability", nil)
|
||||
require.Error(t, err, "NewLoader must reject %q", tc.relPath)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewLoader_ValidBasenameLoads proves the allowlist is exclusive: a
|
||||
// basename containing only safe characters under baseDir loads. Without this
|
||||
// a regression that over-tightened the regex would silently break valid
|
||||
// deployments.
|
||||
func TestNewLoader_ValidBasenameLoads(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing-v2_prod.yaml"))
|
||||
|
||||
l, err := NewLoader(base, "pricing-v2_prod.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err, "basename with _, -, . must load")
|
||||
require.NotNil(t, l.Get(), "table populated")
|
||||
}
|
||||
|
||||
// TestNewLoader_SymlinkOutsideBaseDirRejected constructs a symlink under
|
||||
// baseDir that points to a file outside it. O_NOFOLLOW must refuse to open
|
||||
// the symlink even though the symlink path itself is a valid basename under
|
||||
// baseDir.
|
||||
func TestNewLoader_SymlinkOutsideBaseDirRejected(t *testing.T) {
|
||||
outside := t.TempDir()
|
||||
target := filepath.Join(outside, "evil.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target)
|
||||
|
||||
base := t.TempDir()
|
||||
link := filepath.Join(base, "pricing.yaml")
|
||||
require.NoError(t, os.Symlink(target, link), "symlink setup")
|
||||
|
||||
_, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.Error(t, err, "O_NOFOLLOW must reject symlink even when it points outside baseDir")
|
||||
}
|
||||
|
||||
func TestNewLoader_SymlinkRejected(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
concrete := filepath.Join(base, "real.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), concrete)
|
||||
|
||||
link := filepath.Join(base, "pricing.yaml")
|
||||
require.NoError(t, os.Symlink(concrete, link), "symlink setup")
|
||||
|
||||
_, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.Error(t, err, "O_NOFOLLOW must reject symlinked targets")
|
||||
}
|
||||
|
||||
func TestTableCost_NilSafe(t *testing.T) {
|
||||
@@ -141,37 +402,31 @@ func TestTableCost_NilSafe(t *testing.T) {
|
||||
assert.False(t, t1.Has("x", "y"), "nil table has nothing")
|
||||
}
|
||||
|
||||
func TestNewTable_ValidatesRates(t *testing.T) {
|
||||
good := map[string]map[string]EntryJSON{
|
||||
"openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01, CachedInputPer1K: 0.00125}},
|
||||
}
|
||||
tbl, err := NewTable(good)
|
||||
require.NoError(t, err)
|
||||
cost, ok := tbl.Cost("openai", "gpt-4o", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "entry survives the wire conversion")
|
||||
assert.InDelta(t, 0.0125, cost, 1e-9)
|
||||
|
||||
_, ok = tbl.Cost("openai", "unknown-model", 1, 1, 0, 0)
|
||||
assert.False(t, ok, "unknown model misses")
|
||||
|
||||
for name, bad := range map[string]EntryJSON{
|
||||
"negative input": {InputPer1K: -1, OutputPer1K: 0.01},
|
||||
"NaN output": {InputPer1K: 0.01, OutputPer1K: math.NaN()},
|
||||
"Inf cache read": {InputPer1K: 0.01, OutputPer1K: 0.01, CacheReadPer1K: math.Inf(1)},
|
||||
"negative cached": {InputPer1K: 0.01, OutputPer1K: 0.01, CachedInputPer1K: -0.001},
|
||||
} {
|
||||
_, err := NewTable(map[string]map[string]EntryJSON{"openai": {"m": bad}})
|
||||
assert.Error(t, err, "case %q must be rejected so a corrupt config fails the chain build instead of mispricing", name)
|
||||
}
|
||||
func TestLoaderGet_NilSafe(t *testing.T) {
|
||||
var l *Loader
|
||||
assert.Nil(t, l.Get(), "nil loader returns nil table")
|
||||
}
|
||||
|
||||
func TestNewTable_NilAndEmpty(t *testing.T) {
|
||||
tbl, err := NewTable(nil)
|
||||
require.NoError(t, err, "nil map builds an empty (never-matching) table")
|
||||
_, ok := tbl.Cost("openai", "gpt-4o", 1, 1, 0, 0)
|
||||
assert.False(t, ok, "empty table prices nothing")
|
||||
// TestNewLoader_RejectsOversizedFile_FixesM4 proves the loader bounds reads
|
||||
// at maxPricingBytes so a hostile file cannot exhaust process memory.
|
||||
func TestNewLoader_RejectsOversizedFile_FixesM4(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
target := filepath.Join(base, "pricing.yaml")
|
||||
|
||||
entries, err := NewEntries(nil)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, entries, "nil in, empty (never-matching) map out for the per-record map")
|
||||
// Build a YAML payload larger than the cap. We pad with valid YAML
|
||||
// comments so a partial read would still fail the size check rather
|
||||
// than the parser.
|
||||
header := "openai:\n"
|
||||
bigComment := make([]byte, maxPricingBytes+1024)
|
||||
for i := range bigComment {
|
||||
bigComment[i] = ' '
|
||||
}
|
||||
bigComment[0] = '#'
|
||||
bigComment[len(bigComment)-1] = '\n'
|
||||
payload := append([]byte(header), bigComment...)
|
||||
require.NoError(t, os.WriteFile(target, payload, 0o600))
|
||||
|
||||
_, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.Error(t, err, "oversized pricing file must be rejected")
|
||||
assert.Contains(t, err.Error(), "exceeds", "rejection must reference the byte cap")
|
||||
}
|
||||
|
||||
68
proxy/internal/llm/pricing/pricing_unix.go
Normal file
68
proxy/internal/llm/pricing/pricing_unix.go
Normal file
@@ -0,0 +1,68 @@
|
||||
//go:build unix
|
||||
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// maxPricingBytes caps the size of the pricing YAML on read so a hostile or
|
||||
// runaway file cannot exhaust process memory during reload. 1 MiB is several
|
||||
// orders of magnitude larger than any reasonable pricing table.
|
||||
const maxPricingBytes int64 = 1 << 20
|
||||
|
||||
// loadPricing opens the file with O_NOFOLLOW, fstats the open descriptor,
|
||||
// and parses from that same descriptor. Never re-opens by path so a
|
||||
// mid-read rename or symlink swap cannot substitute content. Bytes are
|
||||
// capped at maxPricingBytes so the loader cannot be coerced into reading an
|
||||
// unbounded file.
|
||||
func loadPricing(path string) (*Table, time.Time, error) {
|
||||
f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW, 0)
|
||||
if err != nil {
|
||||
return nil, time.Time{}, fmt.Errorf("open %s: %w", path, err)
|
||||
}
|
||||
defer func() {
|
||||
if cerr := f.Close(); cerr != nil {
|
||||
log.Debugf("close pricing file %s: %v", path, cerr)
|
||||
}
|
||||
}()
|
||||
|
||||
info, err := f.Stat()
|
||||
if err != nil {
|
||||
return nil, time.Time{}, fmt.Errorf("fstat %s: %w", path, err)
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return nil, time.Time{}, fmt.Errorf("pricing file %s is not a regular file", path)
|
||||
}
|
||||
|
||||
data, err := io.ReadAll(io.LimitReader(f, maxPricingBytes+1))
|
||||
if err != nil {
|
||||
return nil, time.Time{}, fmt.Errorf("read %s: %w", path, err)
|
||||
}
|
||||
if int64(len(data)) > maxPricingBytes {
|
||||
return nil, time.Time{}, fmt.Errorf("pricing file %s exceeds %d bytes", path, maxPricingBytes)
|
||||
}
|
||||
|
||||
table, err := parsePricingBytes(data)
|
||||
if err != nil {
|
||||
return nil, time.Time{}, err
|
||||
}
|
||||
return table, info.ModTime(), nil
|
||||
}
|
||||
|
||||
// statMtime returns the mtime of the file at path. It uses lstat semantics
|
||||
// via os.Lstat so a symlink swap is detected even though O_NOFOLLOW will
|
||||
// later reject the open.
|
||||
func statMtime(path string) (time.Time, error) {
|
||||
info, err := os.Lstat(path)
|
||||
if err != nil {
|
||||
return time.Time{}, fmt.Errorf("lstat %s: %w", path, err)
|
||||
}
|
||||
return info.ModTime(), nil
|
||||
}
|
||||
@@ -36,13 +36,15 @@ var defaultRegistry = middleware.NewRegistry()
|
||||
|
||||
// FactoryContext is the per-process bag that concrete factories may
|
||||
// consult during construction. It carries the proxy-lifetime context,
|
||||
// the OTel meter, and the proxy logger.
|
||||
// the data directory used for static config files (pricing tables,
|
||||
// allowlists), the OTel meter, and the proxy logger.
|
||||
//
|
||||
// Configure must be called once at boot before any chain build calls
|
||||
// Resolve. Calling it twice overwrites the prior value; tests may rely
|
||||
// on this to reset state.
|
||||
type FactoryContext struct {
|
||||
Context context.Context
|
||||
DataDir string
|
||||
Meter metric.Meter
|
||||
Logger *log.Logger
|
||||
MgmtClient MgmtClient
|
||||
@@ -56,11 +58,12 @@ var (
|
||||
// Configure stores the per-process FactoryContext. Concrete factories
|
||||
// reach for it via Context(). mgmt may be nil on tests / standalone
|
||||
// builds with no management server; consumers must guard.
|
||||
func Configure(ctx context.Context, meter metric.Meter, logger *log.Logger, mgmt MgmtClient) {
|
||||
func Configure(ctx context.Context, dataDir string, meter metric.Meter, logger *log.Logger, mgmt MgmtClient) {
|
||||
ctxMu.Lock()
|
||||
defer ctxMu.Unlock()
|
||||
ctxStore = FactoryContext{
|
||||
Context: ctx,
|
||||
DataDir: dataDir,
|
||||
Meter: meter,
|
||||
Logger: logger,
|
||||
MgmtClient: mgmt,
|
||||
|
||||
@@ -12,7 +12,6 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
mgmtpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin/cost_meter"
|
||||
@@ -20,20 +19,17 @@ import (
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin/llm_response_parser"
|
||||
)
|
||||
|
||||
// Drives the real pipeline (llm_request_parser → llm_response_parser → cost_meter) on the REAL default pricing
|
||||
// table management ships (mgmtpricing.DefaultTable, catalog-derived) and asserts exact USD amounts hardcoded from
|
||||
// the vendors' published prices, including the cache split. This is the cross-stack pricing contract test: the
|
||||
// management-side Entry JSON must decode into the proxy-side table and produce these exact costs.
|
||||
// Drives the real pipeline (llm_request_parser → llm_response_parser → cost_meter) on the embedded default pricing
|
||||
// table and asserts exact USD amounts hardcoded from the vendors' published prices, including the cache split.
|
||||
func TestCostCalculation_ProviderMatrix(t *testing.T) {
|
||||
builtin.Configure(context.Background(), nil, nil, nil)
|
||||
// Empty data dir → embedded defaults, like a proxy with no pricing override.
|
||||
builtin.Configure(context.Background(), t.TempDir(), nil, nil, nil)
|
||||
|
||||
reqMW, err := llm_request_parser.Factory{}.New(nil)
|
||||
require.NoError(t, err, "build llm_request_parser")
|
||||
respMW, err := llm_response_parser.Factory{}.New(nil)
|
||||
require.NoError(t, err, "build llm_response_parser")
|
||||
costCfgJSON, err := json.Marshal(map[string]any{"pricing": map[string]any{"defaults": mgmtpricing.DefaultTable()}})
|
||||
require.NoError(t, err, "marshal management default table into cost_meter config")
|
||||
costMW, err := cost_meter.Factory{}.New(costCfgJSON)
|
||||
costMW, err := cost_meter.Factory{}.New(nil)
|
||||
require.NoError(t, err, "build cost_meter")
|
||||
t.Cleanup(func() { _ = costMW.Close() })
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ package cost_meter
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
@@ -10,27 +11,16 @@ import (
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
|
||||
)
|
||||
|
||||
// Config is the on-wire configuration for the middleware, synthesized by
|
||||
// management (buildCostMeterConfigJSON). The proxy has no embedded price
|
||||
// list: this payload is the only pricing source, and updates arrive as
|
||||
// ordinary mapping pushes that rebuild the chain (and with it this
|
||||
// middleware instance) — no per-request fetches, no reload loops.
|
||||
type Config struct {
|
||||
Pricing *PricingConfig `json:"pricing"`
|
||||
}
|
||||
// defaultPricingFilename is the basename probed inside the proxy data
|
||||
// directory when no override is configured.
|
||||
const defaultPricingFilename = "pricing.yaml"
|
||||
|
||||
// PricingConfig carries the full pricing table:
|
||||
// - Defaults: parser surface ("openai"/"anthropic"/"bedrock") ->
|
||||
// normalized model id -> rates, matched against llm.provider +
|
||||
// llm.model.
|
||||
// - Providers: provider record id -> normalized model id -> rates,
|
||||
// matched against the llm.resolved_provider_id metadata llm_router
|
||||
// stamps. Entries arrive fully materialized (management folds default
|
||||
// cache rates in at synth time), so lookup order is simply
|
||||
// per-record first, defaults second.
|
||||
type PricingConfig struct {
|
||||
Defaults map[string]map[string]pricing.EntryJSON `json:"defaults"`
|
||||
Providers map[string]map[string]pricing.EntryJSON `json:"providers"`
|
||||
// Config is the on-wire configuration for the middleware.
|
||||
type Config struct {
|
||||
// PricingPath optionally overrides the basename of the pricing
|
||||
// file probed inside the proxy data directory. When empty the
|
||||
// loader falls back to "pricing.yaml".
|
||||
PricingPath string `json:"pricing_path"`
|
||||
}
|
||||
|
||||
// Factory builds cost_meter instances from raw config bytes.
|
||||
@@ -39,45 +29,45 @@ type Factory struct{}
|
||||
// ID returns the registry identifier.
|
||||
func (Factory) ID() string { return ID }
|
||||
|
||||
// New constructs a middleware instance. Empty, null, and {} configs are
|
||||
// accepted for backward compatibility with a management server that
|
||||
// predates config-delivered pricing — the instance then skips every cost
|
||||
// computation (unknown_model) and a warning is logged once at build time.
|
||||
// Non-empty rawConfig that fails to unmarshal, or a table carrying a
|
||||
// non-finite / negative rate, is rejected so misconfigurations surface at
|
||||
// chain build time.
|
||||
// New constructs a middleware instance. Empty, null, and {} configs
|
||||
// are accepted; non-empty rawConfig that fails to unmarshal is
|
||||
// rejected so misconfigurations surface at chain build time. The
|
||||
// pricing loader is built once per instance and reused across
|
||||
// invocations.
|
||||
func (Factory) New(rawConfig []byte) (middleware.Middleware, error) {
|
||||
cfg, err := decodeConfig(rawConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if cfg.Pricing == nil {
|
||||
if logger := builtin.Context().Logger; logger != nil {
|
||||
logger.Warnf("cost_meter: no pricing table in middleware config; management predates config-delivered pricing — every request will record cost.skipped=unknown_model ($0)")
|
||||
}
|
||||
return newMiddleware(mustEmptyTable(), nil), nil
|
||||
fctx := builtin.Context()
|
||||
pricingPath := cfg.PricingPath
|
||||
if pricingPath == "" {
|
||||
pricingPath = defaultPricingFilename
|
||||
}
|
||||
|
||||
defaults, err := pricing.NewTable(cfg.Pricing.Defaults)
|
||||
loader, err := pricing.NewLoader(fctx.DataDir, pricingPath, ID, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("cost_meter pricing defaults: %w", err)
|
||||
return nil, fmt.Errorf("init pricing loader: %w", err)
|
||||
}
|
||||
perRecord, err := pricing.NewEntries(cfg.Pricing.Providers)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("cost_meter per-provider pricing: %w", err)
|
||||
}
|
||||
return newMiddleware(defaults, perRecord), nil
|
||||
|
||||
cancel := startReloader(fctx.Context, loader)
|
||||
|
||||
return newMiddleware(loader, cancel), nil
|
||||
}
|
||||
|
||||
// mustEmptyTable returns a valid empty table. NewTable on a nil map cannot
|
||||
// fail; the panic guard documents that invariant.
|
||||
func mustEmptyTable() *pricing.Table {
|
||||
t, err := pricing.NewTable(nil)
|
||||
if err != nil {
|
||||
panic(fmt.Sprintf("cost_meter: empty pricing table must build: %v", err))
|
||||
// startReloader binds the loader's mtime-poll goroutine to a context
|
||||
// derived from the proxy-lifetime context and returns its cancel func so
|
||||
// the owning middleware can stop the goroutine on teardown. Returns nil
|
||||
// when there's nothing to watch (nil context or defaults-only loader), in
|
||||
// which case the middleware's Close is a no-op.
|
||||
func startReloader(ctx context.Context, loader *pricing.Loader) context.CancelFunc {
|
||||
if ctx == nil || !loader.WatchesFile() {
|
||||
return nil
|
||||
}
|
||||
return t
|
||||
cctx, cancel := context.WithCancel(ctx)
|
||||
go loader.Reload(cctx)
|
||||
return cancel
|
||||
}
|
||||
|
||||
// decodeConfig accepts empty, null, and {} configs, returning a
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
// Package cost_meter implements the SlotOnResponse middleware that
|
||||
// converts token-usage metadata emitted by llm_response_parser into a
|
||||
// per-request USD cost estimate. Pricing arrives from management inside
|
||||
// the middleware config: a per-provider-record table (the operator's
|
||||
// stored prices, matched via llm.resolved_provider_id) consulted first,
|
||||
// then the surface-keyed defaults table.
|
||||
// per-request USD cost estimate. The middleware uses the shared pricing
|
||||
// loader so operator pricing overrides apply to the chain.
|
||||
package cost_meter
|
||||
|
||||
import (
|
||||
@@ -19,9 +17,7 @@ import (
|
||||
const ID = "cost_meter"
|
||||
|
||||
// Version is the implementation version emitted via the spec merge.
|
||||
// 1.1.0: pricing is config-delivered (defaults + per-provider-record
|
||||
// entries) instead of proxy-embedded.
|
||||
const Version = "1.1.0"
|
||||
const Version = "1.0.0"
|
||||
|
||||
// Skip reasons emitted under KeyCostSkipped. The set is closed; the
|
||||
// dashboard surfaces these verbatim.
|
||||
@@ -46,21 +42,19 @@ var metadataKeys = []string{
|
||||
}
|
||||
|
||||
// Middleware computes a per-response cost estimate from the token
|
||||
// counts emitted upstream by llm_response_parser. Both tables are
|
||||
// immutable — a pricing change arrives as a mapping push that rebuilds
|
||||
// the chain with a fresh instance.
|
||||
// counts emitted upstream by llm_response_parser.
|
||||
type Middleware struct {
|
||||
// defaults is the surface-keyed table (llm.provider x llm.model).
|
||||
defaults *pricing.Table
|
||||
// perRecord is keyed by provider record id (llm.resolved_provider_id)
|
||||
// then normalized model id; entries arrive fully materialized from
|
||||
// management. Consulted before defaults. May be nil.
|
||||
perRecord map[string]map[string]pricing.Entry
|
||||
loader *pricing.Loader
|
||||
// cancel stops this instance's pricing-reload goroutine. Non-nil only
|
||||
// when the loader watches an override file; Close calls it so a chain
|
||||
// rebuild doesn't leak a poll goroutine per retired instance.
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
// newMiddleware constructs a Middleware over the given pricing tables.
|
||||
func newMiddleware(defaults *pricing.Table, perRecord map[string]map[string]pricing.Entry) *Middleware {
|
||||
return &Middleware{defaults: defaults, perRecord: perRecord}
|
||||
// newMiddleware constructs a Middleware bound to the given pricing loader.
|
||||
// cancel may be nil (defaults-only loader with no reloader to stop).
|
||||
func newMiddleware(loader *pricing.Loader, cancel context.CancelFunc) *Middleware {
|
||||
return &Middleware{loader: loader, cancel: cancel}
|
||||
}
|
||||
|
||||
// ID returns the registry identifier.
|
||||
@@ -85,9 +79,16 @@ func (m *Middleware) MetadataKeys() []string {
|
||||
// response.
|
||||
func (m *Middleware) MutationsSupported() bool { return false }
|
||||
|
||||
// Close releases resources owned by the middleware. Stateless — the
|
||||
// pricing tables are plain maps owned by this instance.
|
||||
func (m *Middleware) Close() error { return nil }
|
||||
// Close stops this instance's pricing-reload goroutine, if any. Called by
|
||||
// the chain when a rebuild retires the instance, so the mtime-poll loop
|
||||
// doesn't outlive the chain it belonged to. Safe to call on a nil receiver
|
||||
// and on an instance with no reloader.
|
||||
func (m *Middleware) Close() error {
|
||||
if m != nil && m.cancel != nil {
|
||||
m.cancel()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Invoke reads provider, model, and token metadata, looks up pricing,
|
||||
// and emits either KeyCostUSDTotal or KeyCostSkipped. The decision is
|
||||
@@ -143,7 +144,8 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar
|
||||
return out, nil
|
||||
}
|
||||
|
||||
costs, ok := m.lookupCosts(in.Metadata, provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
|
||||
table := m.loader.Get()
|
||||
costs, ok := table.Costs(provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
|
||||
if !ok {
|
||||
out.Metadata = skip(skipUnknownModel)
|
||||
return out, nil
|
||||
@@ -162,26 +164,6 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// lookupCosts resolves the price for this request and computes the cost
|
||||
// split. Resolution order:
|
||||
//
|
||||
// 1. Per-provider-record entry: the operator's stored price for the
|
||||
// provider route that served the request, keyed by the
|
||||
// llm.resolved_provider_id metadata llm_router stamped on the allow
|
||||
// path. Absent metadata (e.g. no router in the chain) skips this tier.
|
||||
// 2. Surface defaults: the catalog-derived table keyed by llm.provider.
|
||||
//
|
||||
// The surface always selects the cache formula — a per-record entry for an
|
||||
// Anthropic route still bills its cache buckets additively.
|
||||
func (m *Middleware) lookupCosts(md []middleware.KV, surface, model string, inTokens, outTokens, cachedTokens, cacheCreationTokens int64) (pricing.Costs, bool) {
|
||||
if recordID := lookupKV(md, middleware.KeyLLMResolvedProviderID); recordID != "" {
|
||||
if entry, ok := m.perRecord[recordID][model]; ok {
|
||||
return pricing.EntryCosts(entry, surface, inTokens, outTokens, cachedTokens, cacheCreationTokens), true
|
||||
}
|
||||
}
|
||||
return m.defaults.Costs(surface, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
|
||||
}
|
||||
|
||||
// usd renders a cost as the fixed-precision string every cost.usd_* key
|
||||
// carries, so the per-bucket values and the aggregates round identically.
|
||||
//
|
||||
|
||||
@@ -3,50 +3,39 @@ package cost_meter
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/llm/pricing"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
|
||||
)
|
||||
|
||||
// fixtureConfig mirrors what management's buildCostMeterConfigJSON ships:
|
||||
// a surface-keyed defaults table. Rates match the retired YAML fixture so
|
||||
// every cost assertion below is byte-identical to the pre-feature values.
|
||||
func fixtureConfig(t *testing.T) []byte {
|
||||
t.Helper()
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Defaults: map[string]map[string]pricing.EntryJSON{
|
||||
"openai": {
|
||||
"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01},
|
||||
"gpt-4o-mini": {InputPer1K: 0.00015, OutputPer1K: 0.0006},
|
||||
},
|
||||
"anthropic": {
|
||||
"claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015},
|
||||
},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
return raw
|
||||
}
|
||||
const fixturePricing = `openai:
|
||||
gpt-4o:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.01
|
||||
gpt-4o-mini:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.0006
|
||||
anthropic:
|
||||
claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
`
|
||||
|
||||
// fixtureConfigWithCache adds the cache-rate fields.
|
||||
func fixtureConfigWithCache(t *testing.T) []byte {
|
||||
// configureBuiltin points the package-level FactoryContext at a tmp
|
||||
// directory containing the test pricing fixture. Returns the path so
|
||||
// callers can override files later if needed.
|
||||
func configureBuiltin(t *testing.T) string {
|
||||
t.Helper()
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Defaults: map[string]map[string]pricing.EntryJSON{
|
||||
"openai": {
|
||||
"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01, CachedInputPer1K: 0.00125},
|
||||
},
|
||||
"anthropic": {
|
||||
"claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015, CacheReadPer1K: 0.0003, CacheCreationPer1K: 0.00375},
|
||||
},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
return raw
|
||||
dir := t.TempDir()
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "pricing.yaml"), []byte(fixturePricing), 0o600), "write pricing fixture")
|
||||
builtin.Configure(context.Background(), dir, nil, nil, nil)
|
||||
return dir
|
||||
}
|
||||
|
||||
func metaValue(t *testing.T, kvs []middleware.KV, key string) (string, bool) {
|
||||
@@ -67,7 +56,8 @@ func buildMiddleware(t *testing.T, raw []byte) middleware.Middleware {
|
||||
}
|
||||
|
||||
func TestMiddleware_StaticSurface(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
assert.Equal(t, ID, mw.ID(), "ID must match the registered constant")
|
||||
assert.Equal(t, Version, mw.Version(), "Version must match the constant")
|
||||
@@ -89,10 +79,8 @@ func TestMiddleware_StaticSurface(t *testing.T) {
|
||||
assert.Equal(t, expected, keys, "metadata key allowlist must match the spec")
|
||||
}
|
||||
|
||||
// TestFactory_AcceptsEmptyAndJSONConfig: empty/null/{} configs are what an
|
||||
// old management (pre config-delivered pricing) sends — they must build a
|
||||
// working (all-skip) instance, never fail the chain.
|
||||
func TestFactory_AcceptsEmptyAndJSONConfig(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
cases := [][]byte{nil, {}, []byte("null"), []byte("{}"), []byte(" ")}
|
||||
for _, raw := range cases {
|
||||
mw, err := Factory{}.New(raw)
|
||||
@@ -102,57 +90,15 @@ func TestFactory_AcceptsEmptyAndJSONConfig(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestFactory_RejectsMalformedConfig(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw, err := Factory{}.New([]byte("{not json"))
|
||||
require.Error(t, err, "malformed config must surface at construction")
|
||||
assert.Nil(t, mw, "no instance is returned on error")
|
||||
}
|
||||
|
||||
// TestFactory_RejectsInvalidRates: a non-finite or negative rate anywhere
|
||||
// in the table fails the chain build (defense-in-depth behind management's
|
||||
// API validation) rather than silently mispricing.
|
||||
func TestFactory_RejectsInvalidRates(t *testing.T) {
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Defaults: map[string]map[string]pricing.EntryJSON{
|
||||
"openai": {"gpt-4o": {InputPer1K: -0.0025, OutputPer1K: 0.01}},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
mw, err := Factory{}.New(raw)
|
||||
require.Error(t, err, "negative rate must fail the build")
|
||||
assert.Nil(t, mw)
|
||||
|
||||
raw, err = json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Providers: map[string]map[string]pricing.EntryJSON{
|
||||
"prov-1": {"m": {InputPer1K: 0.01, OutputPer1K: 0.01, CacheReadPer1K: -1}},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
_, err = Factory{}.New(raw)
|
||||
require.Error(t, err, "per-record tables validate too")
|
||||
}
|
||||
|
||||
// TestFactory_NilPricingSkipsEverything is the version-skew contract: a
|
||||
// new proxy under an old management ({} config) must build, allow, and
|
||||
// skip with unknown_model — degraded but never broken.
|
||||
func TestFactory_NilPricingSkipsEverything(t *testing.T) {
|
||||
mw := buildMiddleware(t, []byte("{}"))
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "1000"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "1000"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "cost_meter always allows")
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
|
||||
require.True(t, ok, "no pricing table means every request skips")
|
||||
assert.Equal(t, skipUnknownModel, value)
|
||||
}
|
||||
|
||||
func TestFactory_ConfigDefaultsPriceRequests(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
func TestFactory_DefaultPricingPathLoadsFixture(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -170,78 +116,34 @@ func TestFactory_ConfigDefaultsPriceRequests(t *testing.T) {
|
||||
assert.Equal(t, "0.000750000", value, "0.00015 + 0.0006 per 1k tokens, 9-decimal format")
|
||||
}
|
||||
|
||||
// TestInvoke_PerRecordEntryBeatsDefaults: when llm_router resolved a
|
||||
// provider record whose operator pinned a price for the model, that price
|
||||
// wins over the surface default.
|
||||
func TestInvoke_PerRecordEntryBeatsDefaults(t *testing.T) {
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Defaults: map[string]map[string]pricing.EntryJSON{
|
||||
"openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}},
|
||||
},
|
||||
Providers: map[string]map[string]pricing.EntryJSON{
|
||||
"prov-azure": {"gpt-4o": {InputPer1K: 0.005, OutputPer1K: 0.02}},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
mw := buildMiddleware(t, raw)
|
||||
func TestFactory_PricingPathOverride(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "custom.yaml"), []byte(fixturePricing), 0o600), "write custom pricing")
|
||||
builtin.Configure(context.Background(), dir, nil, nil, nil)
|
||||
|
||||
raw, err := json.Marshal(Config{PricingPath: "custom.yaml"})
|
||||
require.NoError(t, err)
|
||||
|
||||
mw := buildMiddleware(t, raw)
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||
{Key: middleware.KeyLLMResolvedProviderID, Value: "prov-azure"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "1000"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "2000"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "1000"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "0.025000000", value, "operator's per-record price (0.005+0.02) wins over the default (0.0025+0.01)")
|
||||
require.True(t, ok, "cost.usd_total must be emitted with custom pricing path")
|
||||
assert.Equal(t, "0.015000000", value, "2*0.0025 + 1*0.01 = 0.015 with 9-decimal format")
|
||||
}
|
||||
|
||||
// TestInvoke_PerRecordMissFallsBackToDefaults: a resolved record with no
|
||||
// entry for this model (or no entries at all) falls through to the
|
||||
// surface defaults — gateway providers rely on exactly this.
|
||||
func TestInvoke_PerRecordMissFallsBackToDefaults(t *testing.T) {
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Defaults: map[string]map[string]pricing.EntryJSON{
|
||||
"openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}},
|
||||
},
|
||||
Providers: map[string]map[string]pricing.EntryJSON{
|
||||
"prov-1": {"some-other-model": {InputPer1K: 1, OutputPer1K: 1}},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
mw := buildMiddleware(t, raw)
|
||||
func TestInvoke_ComputesCostForKnownModel(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
for name, recordID := range map[string]string{
|
||||
"record with other models": "prov-1",
|
||||
"record with no entries": "prov-gateway",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||
{Key: middleware.KeyLLMResolvedProviderID, Value: recordID},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "1000"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "1000"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok, "per-record miss must fall back to the surface default, not skip")
|
||||
assert.Equal(t, "0.012500000", value, "default rates apply")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestInvoke_NoResolvedProviderIDUsesDefaults: metadata without a
|
||||
// resolved provider id (router denied, or a chain without llm_router)
|
||||
// prices from the defaults table directly.
|
||||
func TestInvoke_NoResolvedProviderIDUsesDefaults(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "anthropic"},
|
||||
@@ -251,15 +153,17 @@ func TestInvoke_NoResolvedProviderIDUsesDefaults(t *testing.T) {
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok)
|
||||
require.True(t, ok, "cost.usd_total must be emitted")
|
||||
assert.Equal(t, "0.018000000", value, "0.003 + 0.015 = 0.018 with 9-decimal format")
|
||||
_, skipped := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
|
||||
assert.False(t, skipped, "cost.skipped must not be set when cost is computed")
|
||||
}
|
||||
|
||||
func TestInvoke_MissingProvider(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -275,7 +179,8 @@ func TestInvoke_MissingProvider(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_MissingModel(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -291,7 +196,8 @@ func TestInvoke_MissingModel(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_MissingTokens(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
@@ -334,7 +240,8 @@ func TestInvoke_MissingTokens(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_UnparseableTokens(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
@@ -364,7 +271,8 @@ func TestInvoke_UnparseableTokens(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_ZeroTokens(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -383,7 +291,8 @@ func TestInvoke_ZeroTokens(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_UnknownModel(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -400,7 +309,8 @@ func TestInvoke_UnknownModel(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_NilInput(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
@@ -409,12 +319,36 @@ func TestInvoke_NilInput(t *testing.T) {
|
||||
assert.Empty(t, out.Metadata, "no metadata must be emitted on nil input")
|
||||
}
|
||||
|
||||
const fixturePricingWithCache = `openai:
|
||||
gpt-4o:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.01
|
||||
cached_input_per_1k: 0.00125
|
||||
anthropic:
|
||||
claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
`
|
||||
|
||||
// configureBuiltinWithCacheRates points the package-level
|
||||
// FactoryContext at a tmp directory containing pricing entries that
|
||||
// include the cache rate fields.
|
||||
func configureBuiltinWithCacheRates(t *testing.T) {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "pricing.yaml"), []byte(fixturePricingWithCache), 0o600), "write cache-aware pricing fixture")
|
||||
builtin.Configure(context.Background(), dir, nil, nil, nil)
|
||||
}
|
||||
|
||||
// TestInvoke_OpenAICachedSubsetDiscount proves the OpenAI shape end
|
||||
// to end through the middleware: cached_input_tokens is treated as a
|
||||
// SUBSET of input_tokens and discounted at the configured rate, not
|
||||
// added on top.
|
||||
func TestInvoke_OpenAICachedSubsetDiscount(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfigWithCache(t))
|
||||
configureBuiltinWithCacheRates(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -456,7 +390,8 @@ func TestInvoke_OpenAICachedSubsetDiscount(t *testing.T) {
|
||||
// shape: cache_read and cache_creation are additive to input_tokens
|
||||
// and each carries its own rate.
|
||||
func TestInvoke_AnthropicCacheBucketsAdditive(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfigWithCache(t))
|
||||
configureBuiltinWithCacheRates(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -494,37 +429,8 @@ func TestInvoke_AnthropicCacheBucketsAdditive(t *testing.T) {
|
||||
"output bucket bills 200 tokens at 0.015/1k")
|
||||
}
|
||||
|
||||
// TestInvoke_PerRecordEntryUsesSurfaceFormula: a per-record entry for an
|
||||
// anthropic-surface request must bill its cache buckets additively — the
|
||||
// formula follows llm.provider, not which table the entry came from.
|
||||
func TestInvoke_PerRecordEntryUsesSurfaceFormula(t *testing.T) {
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Providers: map[string]map[string]pricing.EntryJSON{
|
||||
"prov-ant": {"claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015, CacheReadPer1K: 0.0003, CacheCreationPer1K: 0.00375}},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
mw := buildMiddleware(t, raw)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "anthropic"},
|
||||
{Key: middleware.KeyLLMModel, Value: "claude-sonnet-4-5"},
|
||||
{Key: middleware.KeyLLMResolvedProviderID, Value: "prov-ant"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "256"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "200"},
|
||||
{Key: middleware.KeyLLMCachedInputTokens, Value: "768"},
|
||||
{Key: middleware.KeyLLMCacheCreationTokens, Value: "512"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "0.005918400", value, "identical math to the defaults-table entry with the same rates")
|
||||
}
|
||||
|
||||
// assertBucket asserts one per-bucket cost key carries the expected
|
||||
// 9-decimal value.
|
||||
// 6-decimal value.
|
||||
func assertBucket(t *testing.T, md []middleware.KV, key, want, msg string) {
|
||||
t.Helper()
|
||||
got, ok := metaValue(t, md, key)
|
||||
@@ -533,10 +439,13 @@ func assertBucket(t *testing.T, md []middleware.KV, key, want, msg string) {
|
||||
}
|
||||
|
||||
// TestInvoke_CachedTokensAbsentFallsBackToBaseFormula covers the
|
||||
// no-cache-metadata path: with no cached keys emitted, the meter must
|
||||
// produce exactly the input+output cost.
|
||||
// "operator hasn't opted in" path: with no cached metadata keys
|
||||
// emitted, the meter must produce exactly the same cost as before
|
||||
// the feature landed. Critical so operators with the new binary but
|
||||
// no YAML changes see no behavioural drift on OpenAI requests.
|
||||
func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfigWithCache(t))
|
||||
configureBuiltinWithCacheRates(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -551,7 +460,7 @@ func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) {
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok)
|
||||
// 1000 input * 0.0025 + 500 output * 0.01 = 0.0025 + 0.005 = 0.0075
|
||||
assert.Equal(t, "0.007500000", value, "no cached metadata = plain input+output cost")
|
||||
assert.Equal(t, "0.007500000", value, "no cached metadata = same cost as before the feature landed")
|
||||
}
|
||||
|
||||
// TestInvoke_UnparseableCachedTokensSkippedSilently proves the
|
||||
@@ -560,7 +469,8 @@ func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) {
|
||||
// regular formula. Cache buckets are a refinement, never a reason to
|
||||
// abort cost computation.
|
||||
func TestInvoke_UnparseableCachedTokensSkippedSilently(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfigWithCache(t))
|
||||
configureBuiltinWithCacheRates(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -577,10 +487,22 @@ func TestInvoke_UnparseableCachedTokensSkippedSilently(t *testing.T) {
|
||||
assert.Equal(t, "0.007500000", value, "same as the no-cached-metadata path")
|
||||
}
|
||||
|
||||
// TestMiddleware_CloseNilSafe confirms Close is a no-op (no panic) even
|
||||
// for a nil receiver.
|
||||
// TestMiddleware_CloseCancelsReloader proves Close stops the per-instance
|
||||
// pricing-reload goroutine: a chain rebuild retires the old instance and
|
||||
// calls Close, which must invoke the cancel func startReloader handed it so
|
||||
// the mtime-poll loop doesn't outlive the chain.
|
||||
func TestMiddleware_CloseCancelsReloader(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
m := newMiddleware(nil, cancel)
|
||||
|
||||
require.NoError(t, m.Close(), "Close must not error")
|
||||
require.Error(t, ctx.Err(), "Close must cancel the reloader context so the poll goroutine exits")
|
||||
}
|
||||
|
||||
// TestMiddleware_CloseNilSafe confirms Close is a no-op (no panic) for an
|
||||
// instance with no reloader and for a nil receiver.
|
||||
func TestMiddleware_CloseNilSafe(t *testing.T) {
|
||||
require.NoError(t, newMiddleware(nil, nil).Close(), "Close must be a no-op")
|
||||
require.NoError(t, newMiddleware(nil, nil).Close(), "no-reloader Close must be a no-op")
|
||||
var m *Middleware
|
||||
require.NoError(t, m.Close(), "nil-receiver Close must be safe")
|
||||
}
|
||||
|
||||
@@ -6,6 +6,23 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestNormalizeBedrockModel(t *testing.T) {
|
||||
cases := map[string]string{
|
||||
"eu.anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5",
|
||||
"us.anthropic.claude-opus-4-8-20250101-v1:0": "anthropic.claude-opus-4-8",
|
||||
"apac.anthropic.claude-haiku-4-5-v1:0": "anthropic.claude-haiku-4-5",
|
||||
"anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5",
|
||||
"meta.llama3-3-70b-instruct-v1:0": "meta.llama3-3-70b-instruct",
|
||||
"amazon.nova-pro-v1:0": "amazon.nova-pro",
|
||||
"amazon.nova-2-lite-v1:0": "amazon.nova-2-lite",
|
||||
// Inference-profile ARN — model id lives in the last path segment.
|
||||
"arn:aws:bedrock:eu-central-1:123456789012:inference-profile/eu.anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5",
|
||||
}
|
||||
for in, want := range cases {
|
||||
require.Equal(t, want, normalizeBedrockModel(in), "normalize %q", in)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseBedrockPath(t *testing.T) {
|
||||
tests := []struct {
|
||||
path string
|
||||
|
||||
@@ -8,6 +8,7 @@ package llm_request_parser
|
||||
import (
|
||||
"context"
|
||||
"net/url"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
@@ -252,7 +253,9 @@ func parseVertexPath(reqPath string) (vertexRequest, bool) {
|
||||
if c := strings.LastIndex(rest, ":"); c >= 0 {
|
||||
model, action = rest[:c], rest[c+1:]
|
||||
}
|
||||
model = llm.NormalizeVertexModel(model)
|
||||
if at := strings.Index(model, "@"); at >= 0 {
|
||||
model = model[:at]
|
||||
}
|
||||
if model == "" {
|
||||
return vertexRequest{}, false
|
||||
}
|
||||
@@ -340,6 +343,14 @@ func trimBedrockNamespace(reqPath string) string {
|
||||
return reqPath
|
||||
}
|
||||
|
||||
// bedrockRegionPrefixes are the cross-region inference-profile prefixes that
|
||||
// front a Bedrock model id (e.g. "eu.anthropic.claude-...").
|
||||
var bedrockRegionPrefixes = []string{"us.", "eu.", "apac.", "global."}
|
||||
|
||||
// bedrockVersionSuffix matches the trailing "-vN[:N]" or "-YYYYMMDD-vN[:N]"
|
||||
// version/throughput suffix of a Bedrock model id.
|
||||
var bedrockVersionSuffix = regexp.MustCompile(`-(\d{8}-)?v\d+(:\d+)?$`)
|
||||
|
||||
// parseBedrockPath extracts the model and streaming/converse flags from an AWS
|
||||
// Bedrock runtime model endpoint:
|
||||
//
|
||||
@@ -364,7 +375,7 @@ func parseBedrockPath(reqPath string) (bedrockRequest, bool) {
|
||||
if decoded, err := url.PathUnescape(rawModel); err == nil {
|
||||
rawModel = decoded
|
||||
}
|
||||
model := llm.NormalizeBedrockModel(rawModel)
|
||||
model := normalizeBedrockModel(rawModel)
|
||||
if model == "" {
|
||||
return bedrockRequest{}, false
|
||||
}
|
||||
@@ -378,6 +389,30 @@ func parseBedrockPath(reqPath string) (bedrockRequest, bool) {
|
||||
}
|
||||
}
|
||||
|
||||
// normalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile
|
||||
// prefix, and the version/throughput suffix from a Bedrock model id so it
|
||||
// matches the catalog/pricing key, e.g.
|
||||
// "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -> "anthropic.claude-sonnet-4-5"
|
||||
// and "arn:aws:bedrock:eu-central-1:123:inference-profile/eu.anthropic.claude-sonnet-4-5-20250929-v1:0"
|
||||
// -> "anthropic.claude-sonnet-4-5".
|
||||
func normalizeBedrockModel(modelID string) string {
|
||||
m := modelID
|
||||
// A full ARN (inference-profile / provisioned-throughput / foundation-model)
|
||||
// carries the model id in its last path segment.
|
||||
if strings.HasPrefix(m, "arn:") {
|
||||
if i := strings.LastIndex(m, "/"); i >= 0 {
|
||||
m = m[i+1:]
|
||||
}
|
||||
}
|
||||
for _, p := range bedrockRegionPrefixes {
|
||||
if strings.HasPrefix(m, p) {
|
||||
m = m[len(p):]
|
||||
break
|
||||
}
|
||||
}
|
||||
return bedrockVersionSuffix.ReplaceAllString(m, "")
|
||||
}
|
||||
|
||||
// invokeBedrock emits the model/provider/session/prompt for an AWS Bedrock
|
||||
// request. Bedrock is metered under the dedicated "bedrock" parser, which reads
|
||||
// both the InvokeModel and Converse response shapes.
|
||||
|
||||
@@ -18,10 +18,10 @@ import (
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
"google.golang.org/grpc/test/bufconn"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
mgmtgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/bodytap"
|
||||
@@ -134,18 +134,14 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) {
|
||||
RedactPii: true,
|
||||
}))
|
||||
require.NoError(t, st.SaveAgentNetworkProvider(ctx, &agentNetworkTypes.Provider{
|
||||
ID: providerID,
|
||||
AccountID: testAccountID,
|
||||
ProviderID: "openai_api",
|
||||
Name: "openai-fullchain-test",
|
||||
UpstreamURL: upstream.URL, // router rewrites to this
|
||||
APIKey: "sk-test",
|
||||
Enabled: true,
|
||||
// Operator-pinned prices deliberately differ from the catalog's
|
||||
// gpt-5.4 rates (0.0025/0.015) so the cost assertion below proves
|
||||
// the per-provider-record price — not the default table — billed
|
||||
// this request: stored price → synth → wire → cost_meter.
|
||||
Models: []agentNetworkTypes.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.004, OutputPer1k: 0.02}},
|
||||
ID: providerID,
|
||||
AccountID: testAccountID,
|
||||
ProviderID: "openai_api",
|
||||
Name: "openai-fullchain-test",
|
||||
UpstreamURL: upstream.URL, // router rewrites to this
|
||||
APIKey: "sk-test",
|
||||
Enabled: true,
|
||||
Models: []agentNetworkTypes.ProviderModel{{ID: "gpt-5.4"}},
|
||||
SessionPrivateKey: "priv",
|
||||
SessionPublicKey: "pub",
|
||||
}))
|
||||
@@ -180,7 +176,7 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) {
|
||||
|
||||
// ---- 5. Wire the middleware framework — same registry the proxy uses
|
||||
// in production, configured with our bufconn-backed management client.
|
||||
mwbuiltin.Configure(ctx, nil, testLogger, mgmtClient)
|
||||
mwbuiltin.Configure(ctx, t.TempDir(), nil, testLogger, mgmtClient)
|
||||
registry := mwbuiltin.DefaultRegistry()
|
||||
mwMetrics, err := middleware.NewMetrics(nil)
|
||||
require.NoError(t, err)
|
||||
@@ -287,23 +283,13 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) {
|
||||
if r.DimensionKind == agentNetworkTypes.DimensionGroup &&
|
||||
r.DimensionID == adminGroupID &&
|
||||
r.WindowSeconds == 60 &&
|
||||
r.TokensInput+r.TokensOutput > 0 &&
|
||||
r.CostUSD > 0 {
|
||||
r.TokensInput+r.TokensOutput > 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}, 5*time.Second, 50*time.Millisecond,
|
||||
"Admins group consumption row must increment via the response leg WITH a non-zero cost — a zero cost means the operator's stored price never reached cost_meter (synth → wire → per-record lookup broken)")
|
||||
|
||||
// 8a-cost. Exact cost from the OPERATOR's stored price, not the catalog
|
||||
// default: 12 prompt tokens × 0.004/1k + 40 completion tokens × 0.02/1k
|
||||
// = 0.000048 + 0.0008 = 0.000848. With catalog rates it would be 0.00063
|
||||
// — this assertion distinguishes the two, closing the loop on the whole
|
||||
// dynamic-pricing feature (dashboard save → synth → gRPC wire →
|
||||
// llm.resolved_provider_id lookup → billing).
|
||||
assert.Equal(t, "0.000848000", cd.GetMetadata()["cost.usd_total"],
|
||||
"cost must be computed from the provider record's operator-pinned price")
|
||||
"Admins group consumption row must increment via the response leg — if this fails the proxy's respInput dropped UserGroups again or the parser/recorder wiring is broken")
|
||||
|
||||
// 8b. Both the captured prompt and the captured completion are
|
||||
// redacted — proves the synth threads redact_pii=true into BOTH parser
|
||||
|
||||
@@ -246,6 +246,10 @@ type Server struct {
|
||||
// in processMappings before the receive loop reconnects to resync.
|
||||
// Zero uses defaultMappingBatchWatchdog.
|
||||
MappingBatchWatchdog time.Duration
|
||||
// MiddlewareDataDir is the base directory the middleware system uses to
|
||||
// resolve file-backed configuration (e.g. the cost_meter pricing table).
|
||||
// Empty means any middleware that requires a file fails at configure time.
|
||||
MiddlewareDataDir string
|
||||
// MiddlewareCaptureBudgetBytes overrides the proxy-wide in-flight capture
|
||||
// budget passed to middleware.NewManager. Zero or negative values fall
|
||||
// back to defaultMiddlewareCaptureBudgetBytes (256 MiB).
|
||||
@@ -2089,7 +2093,7 @@ func (s *Server) initMiddlewareManager(ctx context.Context) error {
|
||||
return fmt.Errorf("middleware manager requires metrics bundle")
|
||||
}
|
||||
otelMeter := s.meter.Meter()
|
||||
mwbuiltin.Configure(ctx, otelMeter, s.Logger, s.mgmtClient)
|
||||
mwbuiltin.Configure(ctx, s.MiddlewareDataDir, otelMeter, s.Logger, s.mgmtClient)
|
||||
|
||||
mwMetrics, err := middleware.NewMetrics(otelMeter)
|
||||
if err != nil {
|
||||
|
||||
@@ -1,58 +0,0 @@
|
||||
// Package llm holds LLM model-identifier helpers shared by the proxy and
|
||||
// the management server. The proxy normalizes model ids parsed off inbound
|
||||
// requests; management normalizes the operator's registered model ids at
|
||||
// synthesis time so both sides of the pricing / routing contract compare
|
||||
// equal.
|
||||
package llm
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// bedrockRegionPrefixes are the cross-region inference-profile prefixes that
|
||||
// front a Bedrock model id (e.g. "eu.anthropic.claude-...").
|
||||
var bedrockRegionPrefixes = []string{"us.", "eu.", "apac.", "global."}
|
||||
|
||||
// bedrockVersionSuffix matches the trailing "-vN[:N]" or "-YYYYMMDD-vN[:N]"
|
||||
// version/throughput suffix of a Bedrock model id.
|
||||
var bedrockVersionSuffix = regexp.MustCompile(`-(\d{8}-)?v\d+(:\d+)?$`)
|
||||
|
||||
// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile
|
||||
// prefix, and the version/throughput suffix from a Bedrock model id so it
|
||||
// matches the catalog/pricing key, e.g.
|
||||
// "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -> "anthropic.claude-sonnet-4-5"
|
||||
// and the inference-profile ARN's last segment likewise. It is the single
|
||||
// source of truth shared by the proxy's request parser (which normalizes the
|
||||
// request model from the URL path), the proxy's router (which normalizes the
|
||||
// operator's registered Bedrock model ids so both sides compare equal), and
|
||||
// the management synthesizer (which keys per-provider pricing entries by the
|
||||
// normalized id the parser will emit at billing time).
|
||||
func NormalizeBedrockModel(modelID string) string {
|
||||
m := modelID
|
||||
// A full ARN (inference-profile / provisioned-throughput / foundation-model)
|
||||
// carries the model id in its last path segment.
|
||||
if strings.HasPrefix(m, "arn:") {
|
||||
if i := strings.LastIndex(m, "/"); i >= 0 {
|
||||
m = m[i+1:]
|
||||
}
|
||||
}
|
||||
for _, p := range bedrockRegionPrefixes {
|
||||
if strings.HasPrefix(m, p) {
|
||||
m = m[len(p):]
|
||||
break
|
||||
}
|
||||
}
|
||||
return bedrockVersionSuffix.ReplaceAllString(m, "")
|
||||
}
|
||||
|
||||
// NormalizeVertexModel strips the "@version" suffix from a Vertex AI model id
|
||||
// (e.g. "claude-sonnet-4-5@20250929" -> "claude-sonnet-4-5") so it matches
|
||||
// the catalog/pricing key. Vertex publisher models are priced under their
|
||||
// vendor surface with the bare, unversioned id.
|
||||
func NormalizeVertexModel(modelID string) string {
|
||||
if at := strings.Index(modelID, "@"); at >= 0 {
|
||||
return modelID[:at]
|
||||
}
|
||||
return modelID
|
||||
}
|
||||
@@ -5271,21 +5271,6 @@ components:
|
||||
format: double
|
||||
description: Cost per 1k output tokens, in USD.
|
||||
example: 0.0006
|
||||
cached_input_per_1k:
|
||||
type: number
|
||||
format: double
|
||||
description: OpenAI-shape cache rate — cost per 1k cached prompt tokens (a subset of input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means no discount (cached tokens bill at input_per_1k).
|
||||
example: 0.000075
|
||||
cache_read_per_1k:
|
||||
type: number
|
||||
format: double
|
||||
description: Anthropic-shape cache rate — cost per 1k cache-read tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache reads bill at input_per_1k.
|
||||
example: 0.0003
|
||||
cache_creation_per_1k:
|
||||
type: number
|
||||
format: double
|
||||
description: Anthropic-shape cache rate — cost per 1k cache-creation tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache writes bill at input_per_1k.
|
||||
example: 0.00375
|
||||
required:
|
||||
- id
|
||||
- input_per_1k
|
||||
@@ -5311,21 +5296,6 @@ components:
|
||||
format: double
|
||||
description: Output token price per 1k tokens, in USD.
|
||||
example: 0.015
|
||||
cached_input_per_1k:
|
||||
type: number
|
||||
format: double
|
||||
description: OpenAI-shape cache rate — default cost per 1k cached prompt tokens (a subset of input tokens), in USD. Absent when the model has no cached-input discount.
|
||||
example: 0.000075
|
||||
cache_read_per_1k:
|
||||
type: number
|
||||
format: double
|
||||
description: Anthropic-shape cache rate — default cost per 1k cache-read tokens (additive to input tokens), in USD. Absent when the model has no cache-read rate.
|
||||
example: 0.0003
|
||||
cache_creation_per_1k:
|
||||
type: number
|
||||
format: double
|
||||
description: Anthropic-shape cache rate — default cost per 1k cache-creation tokens (additive to input tokens), in USD. Absent when the model has no cache-creation rate.
|
||||
example: 0.00375
|
||||
context_window:
|
||||
type: integer
|
||||
description: Maximum context window in tokens.
|
||||
@@ -5384,13 +5354,6 @@ components:
|
||||
$ref: '#/components/schemas/AgentNetworkCatalogExtraHeader'
|
||||
identity_injection:
|
||||
$ref: '#/components/schemas/AgentNetworkCatalogIdentityInjection'
|
||||
pricing_surfaces:
|
||||
type: array
|
||||
description: |
|
||||
Cost-meter pricing surfaces this provider's traffic is metered under ("openai", "anthropic", "bedrock"). Tells the dashboard which cache-rate fields apply to this provider's models: "openai" → cached_input_per_1k (cached prompt tokens are a subset of input); "anthropic"/"bedrock" → cache_read_per_1k + cache_creation_per_1k (additive buckets). Absent/empty for gateway and custom entries, whose upstream shape NetBird cannot know ahead of time — surface all cache fields for those.
|
||||
items:
|
||||
type: string
|
||||
example: ["openai"]
|
||||
models:
|
||||
type: array
|
||||
description: Catalog models available for this provider.
|
||||
|
||||
@@ -2017,15 +2017,6 @@ type AgentNetworkCatalogJSONMetadataInjection struct {
|
||||
|
||||
// AgentNetworkCatalogModel defines model for AgentNetworkCatalogModel.
|
||||
type AgentNetworkCatalogModel struct {
|
||||
// CacheCreationPer1k Anthropic-shape cache rate — default cost per 1k cache-creation tokens (additive to input tokens), in USD. Absent when the model has no cache-creation rate.
|
||||
CacheCreationPer1k *float64 `json:"cache_creation_per_1k,omitempty"`
|
||||
|
||||
// CacheReadPer1k Anthropic-shape cache rate — default cost per 1k cache-read tokens (additive to input tokens), in USD. Absent when the model has no cache-read rate.
|
||||
CacheReadPer1k *float64 `json:"cache_read_per_1k,omitempty"`
|
||||
|
||||
// CachedInputPer1k OpenAI-shape cache rate — default cost per 1k cached prompt tokens (a subset of input tokens), in USD. Absent when the model has no cached-input discount.
|
||||
CachedInputPer1k *float64 `json:"cached_input_per_1k,omitempty"`
|
||||
|
||||
// ContextWindow Maximum context window in tokens.
|
||||
ContextWindow int `json:"context_window"`
|
||||
|
||||
@@ -2079,9 +2070,6 @@ type AgentNetworkCatalogProvider struct {
|
||||
|
||||
// Name Display name for the provider.
|
||||
Name string `json:"name"`
|
||||
|
||||
// PricingSurfaces Cost-meter pricing surfaces this provider's traffic is metered under ("openai", "anthropic", "bedrock"). Tells the dashboard which cache-rate fields apply to this provider's models: "openai" → cached_input_per_1k (cached prompt tokens are a subset of input); "anthropic"/"bedrock" → cache_read_per_1k + cache_creation_per_1k (additive buckets). Absent/empty for gateway and custom entries, whose upstream shape NetBird cannot know ahead of time — surface all cache fields for those.
|
||||
PricingSurfaces *[]string `json:"pricing_surfaces,omitempty"`
|
||||
}
|
||||
|
||||
// AgentNetworkCatalogProviderKind Presentation grouping for the provider Select on the dashboard.
|
||||
@@ -2305,15 +2293,6 @@ type AgentNetworkProvider struct {
|
||||
|
||||
// AgentNetworkProviderModel A model exposed by the provider, with the operator's per-1k input/output prices in USD.
|
||||
type AgentNetworkProviderModel struct {
|
||||
// CacheCreationPer1k Anthropic-shape cache rate — cost per 1k cache-creation tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache writes bill at input_per_1k.
|
||||
CacheCreationPer1k *float64 `json:"cache_creation_per_1k,omitempty"`
|
||||
|
||||
// CacheReadPer1k Anthropic-shape cache rate — cost per 1k cache-read tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache reads bill at input_per_1k.
|
||||
CacheReadPer1k *float64 `json:"cache_read_per_1k,omitempty"`
|
||||
|
||||
// CachedInputPer1k OpenAI-shape cache rate — cost per 1k cached prompt tokens (a subset of input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means no discount (cached tokens bill at input_per_1k).
|
||||
CachedInputPer1k *float64 `json:"cached_input_per_1k,omitempty"`
|
||||
|
||||
// Id Model identifier (e.g. "gpt-4o-mini").
|
||||
Id string `json:"id"`
|
||||
|
||||
|
||||
Reference in New Issue
Block a user