[management,proxy,client] 0.74.0 version (#6563)

* [management,proxy] Agent network: per-account LLM gateway (policy, metering, multi-provider) (#6555)

* [agent-network] Shared proto, OpenAPI schema, and generated types

* [agent-network] Management: store, manager, synthesizer, policy engine, provider catalog, HTTP/gRPC API

Adds the account-scoped agent-network module: provider/policy/budget CRUD and
store, the reverse-proxy service synthesizer, policy selection + limit
enforcement, the provider catalog (incl. Vertex AI and AWS Bedrock entries),
and the management HTTP + proxy gRPC surfaces.

* [management] Fix agent-network proxy-peer fan-out on affected-peer recompute

The affected-peers resolver loaded only persisted reverse-proxy services, but
agent-network services are synthesized on demand and never persisted. As a
result the embedded proxy peer was never folded into the affected set when a
client's group changed, so the proxy received no network-map update for a newly
authorised client and rejected its handshake until a full resync (restart).

loadProxyServices now merges the synthesized agent-network services (injected
via a registration hook to avoid an import cycle), so proxy peers learn newly
authorised clients immediately.

* [proxy] Reverse-proxy middleware framework, chain, and request plumbing

The per-target middleware chain (slots, dispatcher, mutation gate, metadata
merger), body capture, access-log terminal sink, and the proxy wiring that
builds + runs chains for synthesized agent-network services.

* [proxy] LLM parsers, pricing, and builtin middlewares (OpenAI, Anthropic, Vertex AI, AWS Bedrock)

Request/response parsers and SSE/event-stream metering, the embedded pricing
table, and the builtin middleware set: request parser, router, policy
limit-check/record, cost meter, guardrail, identity inject, response parser.
Includes the path-routed providers — Google Vertex AI (keyfile:: service-account
OAuth minting) and AWS Bedrock (bearer auth, invoke/converse/streaming, optional
/bedrock prefix) — plus the Models allowlist and unmeterable-publisher deny.

* [proxy] IPv6 in-place apply and TCP accept-loop hardening on netstack listeners

* [agent-network] End-to-end test suite, module docs, and deployment preset

* [agent-network] Fix codespell typos and exclude false positives

- labelgen word pool: vermillion -> vermilion, racoon -> raccoon.
- codespell ignore list: add flate (Go compress/flate package), recordin
  (a test-local identifier), and unparseable (a valid alternative spelling used
  consistently across identifiers + a metadata-value constant).

* [management] Set LastSeen on injected proxy peer in realstack test (MySQL strict-mode)

The injected embedded proxy peer had a PeerStatus with a zero LastSeen, which
serializes to '0000-00-00' and is rejected by MySQL in strict mode (SQLite
tolerates it). Set LastSeen to a valid time so SaveAccount succeeds on both
engines.

* [agent-network] Remove e2e shell-script suite from this branch

The end-to-end shell scripts under scripts/e2e/ are maintained in a separate
testing suite and are not part of this change set.

* [agent-network] Polish module docs: remove internal review scaffolding, fix links, verify diagrams

Strip PR-review framing, commit references, absolute paths, and stale internal
references from the agent-network module docs; fix broken relative links; verify
all diagrams against the current architecture. Remove the internal AI-reviewer
prompt file.

* [management] Refine session expiration handling to support 3-state encoding for SSO deadlines

* [agent-network] Relocate agentnetwork package to internals/modules

Move management/server/agentnetwork (and its catalog/, labelgen/, types/
subpackages) to management/internals/modules/agentnetwork, alongside the
reverse-proxy module, and rewrite all importers. Pure relocation: package names,
the synthesizer + affectedpeers registration hook, and store access (shared
store.Store) are unchanged, so no import cycle is introduced (affectedpeers
still depends only on the agentnetwork/types leaf).

* [agent-network] Co-locate HTTP handlers in the module (RegisterEndpoints)

Move the agent-network HTTP handlers from server/http/handlers/agentnetwork into
the module at internals/modules/agentnetwork/handlers (package handlers) and
rename the entrypoint AddEndpoints -> RegisterEndpoints, matching the
reverse-proxy module convention. Wiring in http/handler.go updated accordingly.

* Update getting started to point to rc when agent network enabled

* Add a reference to a commercial license

* Fix docs localhost link

* Fix docs localhost link

* Add private services domain note

* [management] Add agent-network telemetry metrics (#6561)

Surface agent-network adoption and usage in the self-hosted metrics
worker: distinct accounts, providers, policies, budget rules, accounts
with log collection enabled, and aggregated input/output tokens plus
cost.

Tokens and cost are summed from agent_network_request_usage (the
always-written per-request ledger) so the figures are accurate
regardless of the log-collection toggle and carry no double-counting.
All values come from a handful of indexed aggregate queries run only on
the worker's periodic tick.

Adds store.AgentNetworkMetrics with GetAgentNetworkMetrics on the Store
interface, the SqlStore implementation, and a zero-valued FileStore stub.

* Update NetBird server and proxy image versions to 0.74.0-rc.2

* [management,proxy] Reduce agent-network cognitive complexity (#6566)

Address the SonarCloud quality-gate findings in new agent-network code
by extracting focused helpers. No behavior change.

- synthesizer.go: split buildIdentityInjectConfigJSON into per-shape
  rule builders; extract mergeGuardrail from mergeGuardrails to cut
  nesting depth.
- llm_identity_inject: extract injectionEmitsAnything validation
  predicate from New.
- llm_response_parser/streaming.go: extract applyOpenAIStreamUsage and
  applyAnthropicStreamUsage (via a named anthropicStreamUsage type) and
  simplify the OpenAI scanner loop.
- reverseproxy.go: decompose ServeHTTP into serveRouteError,
  buildTargetContext, serveDirect, serveWithChain, captureRequestForChain,
  serveDeny, newResponseWriter, observeResponse, and forwardUpstream,
  preserving the defer ordering so response observation still reads the
  captured writer before it is released.

* [management] Move agent-network access-log ingest into the agentnetwork module (#6568)

The agent-network access-log ingest path (metaKey wire contract, flatten,
usage derivation, and the dual-write of the usage ledger + settings-gated
full row) lived in the reverseproxy accesslogs manager, even though the
agentnetwork module already owns the rest of that domain — types, read
(ListAccessLogs / GetUsageOverview), the budget-counter writes, and
retention cleanup.

Move it next to the rest: a stateless agentnetwork.IngestAccessLog(ctx,
store, entry) that the reverseproxy SaveAccessLog delegates to when the
entry is agent-network. Removes the agentNetworkTypes import from the
reverseproxy manager. No behavior change; the write/read table separation
is unchanged.

Adds real-store coverage for the disable->enable log-collection toggle
(usage ledger always written, full row gated) plus the metadata parse and
group-dedup helpers, which previously had no dedicated tests.

* Add session view support in the access log

* [management,proxy] Container-based agent-network e2e harness (#6577)

* [e2e] Add container-based agent-network e2e harness (Pillar 1)

Introduce a self-contained, OIDC-free e2e harness that stands up NetBird
in containers, so suites no longer depend on the hand-maintained Tilt
stack or a real IdP.

- harness brings up the combined server (management + signal + relay +
  STUN + embedded IdP) in a single container built from
  combined/Dockerfile.multistage, and mints an admin PAT through the
  unauthenticated /api/setup bootstrap (NB_SETUP_PAT_ENABLED). API access
  goes through the existing shared/management/client/rest typed client.
- the image is built via the docker CLI (BuildKit) so the Dockerfile's
  cache mounts are honored; testcontainers then runs the tagged image.
- everything is behind the `e2e` build tag so normal builds and unit
  tests never pull in testcontainers.

Adds BuildKit cache mounts to combined/Dockerfile.multistage so source
changes recompile incrementally rather than from scratch.

Pillar 1 proven by TestCombinedBootstrap: server builds, boots, mints a
PAT, and the PAT authenticates a real management API call.

* [e2e] Add management-side agent-network scenarios (Pillar 2)

Port the API-driven agent-network scenarios from the bash suites to Go,
sharing one combined server per package run (TestMain) with each test
owning its resource cleanup. Drives the /api/agent-network/* endpoints
through the shared REST client's NewRequest primitive with the generated
api types.

Scenarios:
- provider lifecycle (create/get/list/delete + 404 after delete)
- provider validation (missing api_key, unknown catalog id → 4xx)
- settings collection-toggle round-trip with cluster/subdomain immutability
- policy window floor (reject <60s enabled limit, accept at 60s)
- consumption read endpoint returns an array

All deterministic and dependency-free (dummy provider keys; no upstream
calls), so they run headless in CI.

* [e2e] Add live chat-through-proxy scenario (Pillar 3)

Stand up the full agent-network data path in containers and drive a real
chat-completion through the gateway:

- harness: a shared docker network (combined server reachable by alias),
  a proxy container built from the published reverse-proxy image
  (NB_PROXY_PRIVATE, NB_PROXY_ALLOW_INSECURE, NB_RELAY_TRANSPORT=ws to match
  the combined server's WS-multiplexed relay) with a generated self-signed
  wildcard cert, and a netbird client container that joins via a setup key.
- the combined image, proxy image, and client image default to the
  published rc.2 releases (overridable via NB_E2E_*_IMAGE; a bare local tag
  is built from source instead). Geolocation download is disabled so the
  server starts without external fetches.
- one shared domain is used for the management exposed address, the proxy
  domain, and the agent-network cluster; the proxy token is minted via the
  server CLI (global) to match the manual install.

TestChatCompletionThroughProxy provisions provider+policy+group+setup key,
runs proxy+client, drives an OpenAI chat-completion through the tunnel, and
asserts a 200 plus the ingested access-log row. Requires OPENAI_TOKEN
(skips otherwise). The provider must be created with enabled=true explicitly
— the create default is false despite the API doc.

* [e2e] Run the live chat scenario across a provider matrix

Replace the single-provider chat test with a data-driven matrix that runs
the same scenario through every provider whose credentials are present in
the environment (keys/URLs sourced from ~/.llm-keys locally, Actions
secrets in CI):

- OpenAI (chat), Anthropic (messages), Vercel, OpenRouter, Cloudflare
  (OpenAI-compatible gateways), and Bedrock (path-routed, bearer, via the
  messages shape) — covering both wire shapes and the gateway routing.
- all providers are created enabled with a unique model string so the
  proxy's connect-time snapshot carries them all and model->provider
  routing is unambiguous (provider toggles after connect don't reconcile
  to a connected proxy).
- the client supports both wire shapes (/v1/chat/completions and
  /v1/messages); Cloudflare gets the openai provider segment appended to
  its gateway URL.

Each provider must return 200 through the tunnel and produce an ingested
access-log row. Vertex is intentionally excluded from the uniform matrix:
it needs a bespoke rawPredict request shape rather than the shared
chat/messages path, so it warrants a dedicated scenario.

* [ci] Add manual workflow for the agent-network e2e suite

The e2e suite (build tag `e2e`) stands up the combined server + proxy +
client in Docker and drives live chat-completions, so it is slow and needs
provider credentials. Gate it out of normal CI (it already is, via the
build tag) and run it on demand via workflow_dispatch. Provider scenarios
skip when their secret is unset, so it degrades gracefully.

* [e2e] Add Vertex to the provider matrix; run e2e on ubuntu-latest

Vertex (Anthropic-on-Vertex) doesn't share the chat/messages wire shapes:
the model travels in a rawPredict path and the proxy mints the service
account's OAuth token. Add a Vertex client method that posts
/v1/projects/<project>/locations/<region>/publishers/anthropic/models/<model>:rawPredict
with the Vertex anthropic_version body, and wire it into the matrix as a
path-routed provider (created without a models array). It is keyed off
GOOGLE_VERTEX_SA_BASE64 + GOOGLE_VERTEX_PROJECT (region defaults to
"global", model to a pinned claude snapshot, both overridable).

Also bump the e2e workflow runner to ubuntu-latest and add the Vertex
secrets.

* Add docker/docker and docker/go-connections as direct dependencies in go.mod

* [ci] Trigger agent-network e2e workflow on push to main and pull requests

* [e2e] Fix proxy cert permission denied on Linux CI runners

The proxy bind-mounts a temp dir of self-signed certs. MkdirTemp creates
it 0700 and the key was 0600, which Docker Desktop on macOS ignores but a
non-root proxy container on Linux runners cannot traverse/read, so the
cert watcher failed with "open /certs/tls.crt: permission denied" and the
container exited. Widen the cert dir to 0755 and write the throwaway key
0644 so the proxy uid can read the bind-mounted material.

* [e2e] Build images from source by default instead of pulling rc.2

The agent-network code under test lives in this branch, so the e2e should
exercise it rather than a frozen published release. Flip the harness
default: combined/proxy/client are now built from their in-repo
Dockerfiles (combined/Dockerfile.multistage, proxy/Dockerfile.multistage,
e2e/harness/Dockerfile.client) under local tags. Pulling a published image
stays available by setting NB_E2E_*_IMAGE to a registry reference.

Builds now go through buildx --load so the Dockerfile cache mounts are
honored and the result is loaded for testcontainers. The CI workflow adds
a container-driver builder and a local layer cache (NB_E2E_BUILDX_CACHE)
persisted via actions/cache, which caches the base/apt/dep-download layers
across runs. The Go compile still re-runs each time, as BuildKit mount
caches cannot be exported to the GitHub cache.

* [e2e] Cover real providers in lifecycle + assert real consumption metering

- TestProviderLifecycle now runs per available real provider (create → get →
  list → delete → 404) instead of a single dummy provider, exercising each
  catalog's create and field round-trip. Create is offline, so it stays fast
  and burns no provider quota; falls back to a synthetic OpenAI provider when
  no keys are set.
- TestProvidersMatrix attaches a token limit (high caps, 60s window) to its
  policy, which switches on usage metering, and asserts consumption rows are
  recorded with positive token counts after the live traffic. Consumption is
  account-scoped (keyed by source group / user and window, not per provider),
  so the assertion is aggregate.
- TestProviderValidation gains invalid-upstream and blank-name cases. Create
  validation is uniform across catalogs (no per-provider required-field rules),
  so per-provider rejection cases would be redundant.

* [e2e] Assert session id propagates per provider

Each matrix request now sends a unique session id as the universal
x-session-id header and asserts it round-trips into that provider's
access-log row. This guards the session-grouping contract end to end for
every provider (header extraction runs in llm_request_parser ahead of the
parser-specific body extraction, so it is provider-agnostic).

* [e2e] Drop accidentally committed sync-phases dashboard

netbird-sync-phases.json was swept into the Pillar 1 commit by a broad
git add; it belongs to the unrelated sync-phases metrics work, not this
e2e harness. Remove it from the branch so the PR diff is scoped to the
e2e changes.

* [e2e] Revert accidentally committed sync-phase ingest spec

The netbird_sync_phase measurement spec in metrics ingest was swept into
the Pillar 1 commit; it belongs to the unrelated sync-phases metrics work,
not this e2e harness. Its emission side never landed here, so the spec was
orphaned anyway. Restore ingest/main.go to its origin/main state.

* Fix golint issues

* Fix sonar

* Add access log session test

* Fix access log tests

---------

Co-authored-by: braginini <bangvalo@gmail.com>
Co-authored-by: Zoltan Papp <zoltan.pmail@gmail.com>
This commit is contained in:
Maycon Santos
2026-07-01 12:45:14 +02:00
committed by GitHub
co-authored by braginini Zoltan Papp
parent 3be90f06b2
commit 92a66cdd20
208 changed files with 39957 additions and 688 deletions
@@ -0,0 +1,43 @@
package llm_response_parser
import (
"bytes"
"encoding/json"
"fmt"
"github.com/netbirdio/netbird/proxy/internal/middleware"
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
)
// Factory constructs configured Middleware instances for the registry.
type Factory struct{}
// ID returns the registry identifier.
func (Factory) ID() string { return ID }
// New decodes RawConfig (empty / null / "{}" all accepted) and returns
// a configured Middleware. Construction never fails on a well-formed
// empty config; only structurally invalid JSON is rejected.
func (Factory) New(rawConfig []byte) (middleware.Middleware, error) {
cfg, err := decodeConfig(rawConfig)
if err != nil {
return nil, fmt.Errorf("decode config: %w", err)
}
return New(cfg), nil
}
func decodeConfig(raw []byte) (config, error) {
trimmed := bytes.TrimSpace(raw)
if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) {
return config{}, nil
}
var cfg config
if err := json.Unmarshal(trimmed, &cfg); err != nil {
return config{}, err
}
return cfg, nil
}
func init() {
builtin.Register(Factory{})
}
@@ -0,0 +1,133 @@
package llm_response_parser
import (
"bytes"
"compress/flate"
"compress/gzip"
"compress/zlib"
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/proxy/internal/middleware"
)
// gzipBytes returns data gzip-compressed — the wire shape Anthropic
// returns when the client (Claude Code) negotiated Accept-Encoding: gzip.
func gzipBytes(t *testing.T, data []byte) []byte {
t.Helper()
var buf bytes.Buffer
w := gzip.NewWriter(&buf)
_, err := w.Write(data)
require.NoError(t, err, "gzip write must succeed")
require.NoError(t, w.Close(), "gzip close must succeed")
return buf.Bytes()
}
// TestInvoke_AnthropicStreaming_Gzip is the regression guard for the live
// bug: Claude Code negotiates gzip, Anthropic gzips the SSE stream, the
// proxy captures the compressed bytes, and the parser must decompress
// before accumulating — otherwise token usage is silently dropped and
// cost_meter skips with missing_tokens.
func TestInvoke_AnthropicStreaming_Gzip(t *testing.T) {
m := newTestMiddleware(t)
body := gzipBytes(t, loadFixture(t, "anthropic_stream.txt"))
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{
{Key: "Content-Type", Value: "text/event-stream; charset=utf-8"},
{Key: "Content-Encoding", Value: "gzip"},
},
RespBody: body,
Metadata: []middleware.KV{
{Key: middleware.KeyLLMProvider, Value: "anthropic"},
{Key: middleware.KeyLLMModel, Value: "claude-opus-4-8"},
},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "Invoke must not error on a gzip-encoded streaming body")
in123, ok := metaValue(out.Metadata, middleware.KeyLLMInputTokens)
require.True(t, ok, "input tokens must be emitted from a gzip SSE stream")
assert.Equal(t, "123", in123, "input tokens must survive gzip decompression")
outTok, _ := metaValue(out.Metadata, middleware.KeyLLMOutputTokens)
assert.Equal(t, "45", outTok, "output tokens must survive gzip decompression")
totTok, _ := metaValue(out.Metadata, middleware.KeyLLMTotalTokens)
assert.Equal(t, "168", totTok, "total tokens must survive gzip decompression")
}
// TestInvoke_AnthropicBuffered_Gzip covers the non-streaming JSON path
// under gzip — the same decode must happen before ParseResponse.
func TestInvoke_AnthropicBuffered_Gzip(t *testing.T) {
m := newTestMiddleware(t)
body := gzipBytes(t, loadFixture(t, "anthropic_messages.json"))
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{
{Key: "Content-Type", Value: "application/json"},
{Key: "Content-Encoding", Value: "gzip"},
},
RespBody: body,
Metadata: []middleware.KV{
{Key: middleware.KeyLLMProvider, Value: "anthropic"},
{Key: middleware.KeyLLMModel, Value: "claude-opus-4-8"},
},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "Invoke must not error on a gzip-encoded buffered body")
_, ok := metaValue(out.Metadata, middleware.KeyLLMInputTokens)
require.True(t, ok, "input tokens must be emitted from a gzip JSON body")
}
// TestDecodeResponseBody covers the encoding matrix directly.
func TestDecodeResponseBody(t *testing.T) {
plain := []byte(`{"hello":"world"}`)
t.Run("identity passthrough", func(t *testing.T) {
assert.Equal(t, plain, decodeResponseBody(plain, ""))
assert.Equal(t, plain, decodeResponseBody(plain, "identity"))
})
t.Run("gzip", func(t *testing.T) {
assert.Equal(t, plain, decodeResponseBody(gzipBytes(t, plain), "gzip"))
})
t.Run("gzip with multi-coding header takes outermost", func(t *testing.T) {
assert.Equal(t, plain, decodeResponseBody(gzipBytes(t, plain), "identity, gzip"))
})
t.Run("deflate zlib-wrapped", func(t *testing.T) {
var buf bytes.Buffer
zw := zlib.NewWriter(&buf)
_, _ = zw.Write(plain)
_ = zw.Close()
assert.Equal(t, plain, decodeResponseBody(buf.Bytes(), "deflate"))
})
t.Run("deflate raw flate fallback", func(t *testing.T) {
var buf bytes.Buffer
fw, _ := flate.NewWriter(&buf, flate.DefaultCompression)
_, _ = fw.Write(plain)
_ = fw.Close()
assert.Equal(t, plain, decodeResponseBody(buf.Bytes(), "deflate"))
})
t.Run("gzip header but not actually gzip falls back to raw", func(t *testing.T) {
assert.Equal(t, plain, decodeResponseBody(plain, "gzip"))
})
t.Run("unknown encoding (br) returns raw", func(t *testing.T) {
assert.Equal(t, plain, decodeResponseBody(plain, "br"))
})
}
@@ -0,0 +1,339 @@
// Package llm_response_parser implements the SlotOnResponse middleware
// that decodes OpenAI- and Anthropic-shaped LLM responses (buffered or
// streaming) and emits token usage and completion metadata. Provider
// and model are read from the request-side metadata bag emitted by
// llm_request_parser; without that context the middleware is a no-op.
package llm_response_parser
import (
"bytes"
"compress/flate"
"compress/gzip"
"compress/zlib"
"context"
"io"
"strconv"
"strings"
"unicode/utf8"
"github.com/netbirdio/netbird/proxy/internal/llm"
"github.com/netbirdio/netbird/proxy/internal/middleware"
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin/llm_guardrail"
)
// ID is the registry identifier for this middleware.
const ID = "llm_response_parser"
const version = "1.0.0"
// maxCompletionBytes is the rune-safe cap applied to the extracted
// completion text before emitting it as metadata.
const maxCompletionBytes = 3500
// maxDecodedBytes bounds the inflated size of a compressed response body
// so a small gzip/deflate payload can't expand into a memory blow-up. The
// captured input is already capped (per-direction body cap), so this only
// bounds the decompression ratio; the parser is best-effort and tolerates a
// truncated decode.
const maxDecodedBytes = 16 << 20
var (
acceptedContentTypes = []string{"application/json", "text/event-stream"}
metadataKeys = []string{
middleware.KeyLLMInputTokens,
middleware.KeyLLMOutputTokens,
middleware.KeyLLMTotalTokens,
middleware.KeyLLMCachedInputTokens,
middleware.KeyLLMCacheCreationTokens,
middleware.KeyLLMResponseCompletion,
}
)
// config is the wire-side configuration for this middleware. RedactPii, when
// true, runs PII redaction on the extracted completion text BEFORE it is
// emitted as llm.response_completion — keeping the access-log row free of
// emails / SSNs / phone numbers the model itself generated. CaptureCompletion
// gates emission of the completion key entirely: a nil pointer preserves
// legacy emission (so callers without the toggle aren't broken), an explicit
// false suppresses the key so the access-log row carries token / cost facts
// only. Both are sourced by the synthesiser from the account's redact_pii
// and enable_prompt_collection toggles respectively.
type config struct {
RedactPii bool `json:"redact_pii,omitempty"`
CaptureCompletion *bool `json:"capture_completion,omitempty"`
}
// Middleware implements middleware.Middleware.
type Middleware struct {
parsers []llm.Parser
redactPii bool
captureCompletion bool
}
// New constructs a configured Middleware instance.
func New(cfg config) *Middleware {
capture := true
if cfg.CaptureCompletion != nil {
capture = *cfg.CaptureCompletion
}
return &Middleware{parsers: llm.Parsers(), redactPii: cfg.RedactPii, captureCompletion: capture}
}
// ID returns the registry identifier.
func (m *Middleware) ID() string { return ID }
// Version returns the implementation version.
func (m *Middleware) Version() string { return version }
// Slot reports that the middleware runs after the upstream call.
func (m *Middleware) Slot() middleware.Slot { return middleware.SlotOnResponse }
// AcceptedContentTypes lists the response content types the middleware
// inspects.
func (m *Middleware) AcceptedContentTypes() []string {
return append([]string(nil), acceptedContentTypes...)
}
// MetadataKeys returns the closed allowlist of keys this middleware
// may emit.
func (m *Middleware) MetadataKeys() []string {
return append([]string(nil), metadataKeys...)
}
// MutationsSupported reports that this middleware never mutates the
// response.
func (m *Middleware) MutationsSupported() bool { return false }
// Close releases any resources held by the middleware. The parser-set
// is stateless so this is a no-op.
func (m *Middleware) Close() error { return nil }
// Invoke decodes the response body and emits token-usage and completion
// metadata. The decision is always DecisionAllow; parse errors degrade
// silently to omitted metadata rather than chain failures.
func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middleware.Output, error) {
out := &middleware.Output{Decision: middleware.DecisionAllow}
if in == nil {
return out, nil
}
provider := lookupKV(in.Metadata, middleware.KeyLLMProvider)
if provider == "" {
return out, nil
}
parser := m.parserByName(provider)
if parser == nil {
return out, nil
}
// Upstreams compress the response when the client negotiated it
// (Claude Code sends Accept-Encoding: gzip). The transport leaves it
// compressed because the request carried an explicit Accept-Encoding,
// so the captured copy is gzip/deflate bytes — decompress it before
// parsing or token usage is silently lost. The forwarded client
// stream is untouched; this only affects our parse copy.
body := decodeResponseBody(in.RespBody, headerLookup(in.RespHeaders, "Content-Encoding"))
contentType := headerLookup(in.RespHeaders, "Content-Type")
switch {
case isEventStream(contentType), isAWSEventStream(contentType):
out.Metadata = m.invokeStreaming(parser, body)
case isJSON(contentType):
out.Metadata = m.invokeBuffered(parser, in, contentType, body)
}
return out, nil
}
// invokeBuffered decodes a non-streaming JSON response body. Status
// codes >= 400 short-circuit because providers don't include usage on
// error responses.
func (m *Middleware) invokeBuffered(parser llm.Parser, in *middleware.Input, contentType string, body []byte) []middleware.KV {
if in.Status >= 400 {
return nil
}
var md []middleware.KV
usage, err := parser.ParseResponse(in.Status, contentType, body)
if err == nil {
md = appendUsage(md, usage)
}
if completion := truncateCompletion(parser.ExtractCompletion(in.Status, contentType, body)); completion != "" && m.captureCompletion {
if m.redactPii {
completion = llm_guardrail.RedactPII(completion)
}
md = append(md, middleware.KV{Key: middleware.KeyLLMResponseCompletion, Value: completion})
}
return md
}
// invokeStreaming walks the buffered SSE prefix and accumulates token
// deltas plus completion text. Truncated bodies are processed
// best-effort; partial usage is preferred over no metadata.
func (m *Middleware) invokeStreaming(parser llm.Parser, body []byte) []middleware.KV {
if len(body) == 0 {
return nil
}
usage, completion := accumulateStream(parser.ProviderName(), body)
var md []middleware.KV
if usage.InputTokens > 0 || usage.OutputTokens > 0 || usage.TotalTokens > 0 {
md = appendUsage(md, usage)
}
if c := truncateCompletion(completion); c != "" && m.captureCompletion {
if m.redactPii {
c = llm_guardrail.RedactPII(c)
}
md = append(md, middleware.KV{Key: middleware.KeyLLMResponseCompletion, Value: c})
}
return md
}
// parserByName returns the parser matching the provider label emitted
// by llm_request_parser, or nil when none claims it.
func (m *Middleware) parserByName(name string) llm.Parser {
for _, p := range m.parsers {
if p.ProviderName() == name {
return p
}
}
return nil
}
// appendUsage emits the three baseline token-count metadata keys plus
// optional cached / cache-creation bucket counts when nonzero. Total
// is computed when the provider omitted one but reported per-direction
// counts; cache buckets are excluded from the legacy total because
// llm.input_tokens already absorbs the OpenAI cached subset and the
// sum-of-everything is a separate downstream concern.
func appendUsage(md []middleware.KV, usage llm.Usage) []middleware.KV {
total := usage.TotalTokens
if total == 0 && (usage.InputTokens > 0 || usage.OutputTokens > 0) {
total = usage.InputTokens + usage.OutputTokens
}
md = append(md,
middleware.KV{Key: middleware.KeyLLMInputTokens, Value: strconv.FormatInt(usage.InputTokens, 10)},
middleware.KV{Key: middleware.KeyLLMOutputTokens, Value: strconv.FormatInt(usage.OutputTokens, 10)},
middleware.KV{Key: middleware.KeyLLMTotalTokens, Value: strconv.FormatInt(total, 10)},
)
if usage.CachedInputTokens > 0 {
md = append(md, middleware.KV{
Key: middleware.KeyLLMCachedInputTokens,
Value: strconv.FormatInt(usage.CachedInputTokens, 10),
})
}
if usage.CacheCreationTokens > 0 {
md = append(md, middleware.KV{
Key: middleware.KeyLLMCacheCreationTokens,
Value: strconv.FormatInt(usage.CacheCreationTokens, 10),
})
}
return md
}
// truncateCompletion clamps an extracted completion to maxCompletionBytes.
// The cut is rune-safe so we never split a multi-byte UTF-8 sequence.
func truncateCompletion(s string) string {
if len(s) <= maxCompletionBytes {
return s
}
cut := maxCompletionBytes
for cut > 0 && !utf8.RuneStart(s[cut]) {
cut--
}
return s[:cut]
}
func lookupKV(kvs []middleware.KV, key string) string {
for _, kv := range kvs {
if kv.Key == key {
return kv.Value
}
}
return ""
}
func headerLookup(h []middleware.KV, name string) string {
lower := strings.ToLower(name)
for _, kv := range h {
if strings.ToLower(kv.Key) == lower {
return kv.Value
}
}
return ""
}
func isEventStream(contentType string) bool {
return strings.Contains(strings.ToLower(contentType), "text/event-stream")
}
// isAWSEventStream reports whether contentType is the AWS binary event-stream
// framing used by Bedrock's streaming endpoints.
func isAWSEventStream(contentType string) bool {
return strings.Contains(strings.ToLower(contentType), "application/vnd.amazon.eventstream")
}
func isJSON(contentType string) bool {
lower := strings.ToLower(contentType)
return strings.Contains(lower, "application/json") || strings.Contains(lower, "+json")
}
// decodeResponseBody returns body decompressed per its Content-Encoding,
// or the original bytes when the encoding is identity, unrecognised
// (e.g. br — no stdlib decoder), or the body isn't actually compressed.
// Decoding is best-effort: a truncated stream (capture hit the byte cap)
// yields the decompressed prefix rather than an error, which is enough to
// recover the leading message_start usage on Anthropic SSE.
func decodeResponseBody(body []byte, contentEncoding string) []byte {
enc := strings.ToLower(strings.TrimSpace(contentEncoding))
// Content-Encoding may list multiple codings; the last applied is
// the outermost on the wire.
if idx := strings.LastIndex(enc, ","); idx >= 0 {
enc = strings.TrimSpace(enc[idx+1:])
}
switch enc {
case "", "identity":
return body
case "gzip", "x-gzip":
zr, err := gzip.NewReader(bytes.NewReader(body))
if err != nil {
return body
}
defer zr.Close()
if out := readCapped(zr); len(out) > 0 {
return out
}
return body
case "deflate":
// "deflate" on the wire is usually zlib-wrapped; fall back to raw
// flate when there's no zlib header.
if zr, err := zlib.NewReader(bytes.NewReader(body)); err == nil {
defer zr.Close()
if out := readCapped(zr); len(out) > 0 {
return out
}
return body
}
fr := flate.NewReader(bytes.NewReader(body))
defer fr.Close()
if out := readCapped(fr); len(out) > 0 {
return out
}
return body
default:
return body
}
}
// readCapped reads at most maxDecodedBytes from r, discarding any excess.
// Best-effort: a read error returns whatever was decoded so far, which is
// enough for the parser to recover leading usage events.
func readCapped(r io.Reader) []byte {
out, _ := io.ReadAll(io.LimitReader(r, maxDecodedBytes))
return out
}
@@ -0,0 +1,433 @@
package llm_response_parser
import (
"context"
"os"
"path/filepath"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/proxy/internal/middleware"
)
func loadFixture(t *testing.T, name string) []byte {
t.Helper()
root, err := os.Getwd()
require.NoError(t, err, "must resolve cwd to locate fixture")
dir := root
for i := 0; i < 8; i++ {
candidate := filepath.Join(dir, "proxy", "internal", "llm", "fixtures", name)
if data, err := os.ReadFile(candidate); err == nil {
return data
}
parent := filepath.Dir(dir)
if parent == dir {
break
}
dir = parent
}
t.Fatalf("fixture %q not found relative to %q", name, root)
return nil
}
func metaValue(kvs []middleware.KV, key string) (string, bool) {
for _, kv := range kvs {
if kv.Key == key {
return kv.Value, true
}
}
return "", false
}
func newTestMiddleware(t *testing.T) *Middleware {
t.Helper()
mw, err := Factory{}.New(nil)
require.NoError(t, err, "factory must accept empty config")
concrete, ok := mw.(*Middleware)
require.True(t, ok, "factory must return *Middleware")
return concrete
}
func TestMiddleware_StaticSurface(t *testing.T) {
m := newTestMiddleware(t)
assert.Equal(t, ID, m.ID(), "ID must match registry constant")
assert.Equal(t, "1.0.0", m.Version(), "Version must be 1.0.0")
assert.Equal(t, middleware.SlotOnResponse, m.Slot(), "Slot must be SlotOnResponse")
assert.False(t, m.MutationsSupported(), "response parser does not mutate")
assert.ElementsMatch(t,
[]string{"application/json", "text/event-stream"},
m.AcceptedContentTypes(),
"AcceptedContentTypes must list JSON and SSE",
)
assert.ElementsMatch(t,
[]string{
middleware.KeyLLMInputTokens,
middleware.KeyLLMOutputTokens,
middleware.KeyLLMTotalTokens,
middleware.KeyLLMCachedInputTokens,
middleware.KeyLLMCacheCreationTokens,
middleware.KeyLLMResponseCompletion,
},
m.MetadataKeys(),
"MetadataKeys must be the documented response-side keys, including the optional cache buckets emitted only when nonzero",
)
require.NoError(t, m.Close(), "Close must be a no-op")
}
func TestFactory_AcceptsEmptyAndNullConfig(t *testing.T) {
for name, raw := range map[string][]byte{
"nil": nil,
"empty": {},
"null": []byte("null"),
"obj": []byte("{}"),
"ws": []byte(" "),
} {
t.Run(name, func(t *testing.T) {
mw, err := Factory{}.New(raw)
require.NoError(t, err, "factory must accept %s config", name)
require.NotNil(t, mw, "factory must return middleware for %s", name)
})
}
}
func TestFactory_RejectsMalformedJSON(t *testing.T) {
_, err := Factory{}.New([]byte("not-json"))
require.Error(t, err, "malformed config must surface a decode error")
}
func TestInvoke_OpenAIBuffered(t *testing.T) {
m := newTestMiddleware(t)
body := loadFixture(t, "openai_chat_completion.json")
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: body,
Metadata: []middleware.KV{
{Key: middleware.KeyLLMProvider, Value: "openai"},
{Key: middleware.KeyLLMModel, Value: "gpt-4o-mini"},
},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "Invoke must not error on a valid buffered response")
require.Equal(t, middleware.DecisionAllow, out.Decision, "decision must be Allow")
in123, ok := metaValue(out.Metadata, middleware.KeyLLMInputTokens)
require.True(t, ok, "input tokens must be emitted")
assert.Equal(t, "123", in123, "input tokens must match fixture prompt_tokens")
outTok, ok := metaValue(out.Metadata, middleware.KeyLLMOutputTokens)
require.True(t, ok, "output tokens must be emitted")
assert.Equal(t, "45", outTok, "output tokens must match fixture completion_tokens")
totTok, ok := metaValue(out.Metadata, middleware.KeyLLMTotalTokens)
require.True(t, ok, "total tokens must be emitted")
assert.Equal(t, "168", totTok, "total tokens must match fixture")
completion, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
require.True(t, ok, "completion must be emitted")
assert.Equal(t, "Hello, world!", completion, "completion text must match fixture")
}
func TestInvoke_AnthropicBuffered(t *testing.T) {
m := newTestMiddleware(t)
body := loadFixture(t, "anthropic_messages.json")
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: body,
Metadata: []middleware.KV{
{Key: middleware.KeyLLMProvider, Value: "anthropic"},
{Key: middleware.KeyLLMModel, Value: "claude-sonnet-4-5"},
},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "Invoke must not error on a valid buffered response")
require.Equal(t, middleware.DecisionAllow, out.Decision, "decision must be Allow")
in123, _ := metaValue(out.Metadata, middleware.KeyLLMInputTokens)
assert.Equal(t, "123", in123, "input tokens must match anthropic fixture")
outTok, _ := metaValue(out.Metadata, middleware.KeyLLMOutputTokens)
assert.Equal(t, "45", outTok, "output tokens must match anthropic fixture")
totTok, _ := metaValue(out.Metadata, middleware.KeyLLMTotalTokens)
assert.Equal(t, "168", totTok, "total tokens must be input+output for anthropic")
completion, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
require.True(t, ok, "completion must be emitted for anthropic")
assert.Equal(t, "Hello, world!", completion, "completion text must match fixture")
}
// TestInvoke_OpenAICachedTokensSurfaceOnMetadata covers the
// end-to-end path from the JSON usage block to the
// llm.cached_input_tokens metadata key the cost meter consumes.
// llm.cache_creation_tokens is NOT emitted for OpenAI because
// OpenAI has no cache_creation analogue.
func TestInvoke_OpenAICachedTokensSurfaceOnMetadata(t *testing.T) {
m := newTestMiddleware(t)
body := []byte(`{"usage":{"prompt_tokens":1024,"completion_tokens":200,"total_tokens":1224,"prompt_tokens_details":{"cached_tokens":768}}}`)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: body,
Metadata: []middleware.KV{
{Key: middleware.KeyLLMProvider, Value: "openai"},
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err)
cached, ok := metaValue(out.Metadata, middleware.KeyLLMCachedInputTokens)
require.True(t, ok, "cached_input_tokens must land on the bag when the OpenAI response carries cached_tokens")
assert.Equal(t, "768", cached)
_, hasCreation := metaValue(out.Metadata, middleware.KeyLLMCacheCreationTokens)
assert.False(t, hasCreation, "cache_creation_tokens must NOT be emitted for OpenAI — no analogue in the OpenAI shape")
}
// TestInvoke_AnthropicCacheBucketsSurfaceOnMetadata covers the
// Anthropic shape: both cache_read and cache_creation values flow
// onto the metadata bag so the cost meter can apply per-bucket
// rates.
func TestInvoke_AnthropicCacheBucketsSurfaceOnMetadata(t *testing.T) {
m := newTestMiddleware(t)
body := []byte(`{"usage":{"input_tokens":256,"output_tokens":200,"cache_read_input_tokens":768,"cache_creation_input_tokens":512}}`)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: body,
Metadata: []middleware.KV{
{Key: middleware.KeyLLMProvider, Value: "anthropic"},
{Key: middleware.KeyLLMModel, Value: "claude-sonnet-4-5"},
},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err)
cached, ok := metaValue(out.Metadata, middleware.KeyLLMCachedInputTokens)
require.True(t, ok, "cache_read_input_tokens lands under cached_input_tokens — same key carries OpenAI cached subset and Anthropic cache reads, meter switches formula on provider")
assert.Equal(t, "768", cached)
creation, ok := metaValue(out.Metadata, middleware.KeyLLMCacheCreationTokens)
require.True(t, ok, "cache_creation_input_tokens lands under cache_creation_tokens for Anthropic")
assert.Equal(t, "512", creation)
}
func TestInvoke_NoProviderMetadata_NoOp(t *testing.T) {
m := newTestMiddleware(t)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: loadFixture(t, "openai_chat_completion.json"),
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "missing provider metadata is not an error")
assert.Equal(t, middleware.DecisionAllow, out.Decision, "decision must be Allow")
assert.Empty(t, out.Metadata, "no metadata when provider context is missing")
}
func TestInvoke_UnknownProvider_NoOp(t *testing.T) {
m := newTestMiddleware(t)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: loadFixture(t, "openai_chat_completion.json"),
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "cohere"}},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "unknown provider must not surface an error")
assert.Empty(t, out.Metadata, "unknown providers emit no metadata")
}
func TestInvoke_ErrorStatus_NoUsageEmitted(t *testing.T) {
m := newTestMiddleware(t)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 500,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: []byte(`{"error":{"message":"upstream blew up"}}`),
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "openai"}},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "error responses must not surface as middleware error")
_, ok := metaValue(out.Metadata, middleware.KeyLLMInputTokens)
assert.False(t, ok, "no usage metadata on >=400 responses")
}
func TestInvoke_NonInspectedContentType_NoOp(t *testing.T) {
m := newTestMiddleware(t)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "text/plain"}},
RespBody: []byte("not json"),
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "openai"}},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "Invoke must tolerate non-inspected content types")
assert.Empty(t, out.Metadata, "no metadata for non-JSON, non-SSE bodies")
}
func TestInvoke_NilInput(t *testing.T) {
m := newTestMiddleware(t)
out, err := m.Invoke(context.Background(), nil)
require.NoError(t, err, "nil input must not error")
require.Equal(t, middleware.DecisionAllow, out.Decision, "decision must be Allow even on nil input")
assert.Empty(t, out.Metadata, "no metadata for nil input")
}
func TestInvoke_CompletionTruncatedAt3500Bytes(t *testing.T) {
m := newTestMiddleware(t)
long := strings.Repeat("x", 5000)
body := []byte(`{"id":"x","choices":[{"message":{"role":"assistant","content":"` + long + `"}}],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: body,
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "openai"}},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "long-completion body must parse cleanly")
completion, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
require.True(t, ok, "completion must be emitted for long body")
assert.LessOrEqual(t, len(completion), maxCompletionBytes, "completion must be truncated to <=3500 bytes")
assert.Equal(t, maxCompletionBytes, len(completion), "completion must be truncated exactly at the cap when input is ASCII and longer")
}
// TestInvoke_RedactPii_RedactsCompletionBeforeEmit covers the GC contract on
// the response leg: when the synthesiser sets redact_pii=true, the value
// emitted as llm.response_completion must already be redacted, so the
// access-log row never carries raw emails / SSNs / phones the model generated.
// Without this, the response side leaked dozens of raw PII tokens per request.
func TestInvoke_RedactPii_RedactsCompletionBeforeEmit(t *testing.T) {
mw, err := Factory{}.New([]byte(`{"redact_pii":true}`))
require.NoError(t, err)
piiCompletion := "Sample record: Alice Johnson, alice.johnson@example.com, SSN 123-45-6789, phone (202) 555-0147. Bob: 202/555/0108."
body := []byte(`{"id":"x","choices":[{"message":{"role":"assistant","content":"` + piiCompletion + `"}}],"usage":{"prompt_tokens":10,"completion_tokens":50,"total_tokens":60}}`)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: body,
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "openai"}},
}
out, err := mw.Invoke(context.Background(), in)
require.NoError(t, err)
require.NotNil(t, out)
completion, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
require.True(t, ok, "completion key must be emitted")
assert.Contains(t, completion, "[REDACTED:email]", "email must be redacted before emit")
assert.Contains(t, completion, "[REDACTED:ssn]", "ssn must be redacted before emit")
assert.Contains(t, completion, "[REDACTED:phone]", "phone must be redacted before emit")
assert.NotContains(t, completion, "alice.johnson@example.com", "raw email must not survive")
assert.NotContains(t, completion, "123-45-6789", "raw SSN must not survive")
assert.NotContains(t, completion, "(202) 555-0147", "parens-phone must not survive")
assert.NotContains(t, completion, "202/555/0108", "slash-phone must not survive")
}
// TestInvoke_CaptureCompletionOff_DoesNotEmitCompletion mirrors the request
// parser test: when capture_completion=false (operator has enable_prompt_
// collection off), llm.response_completion MUST NOT appear in the access log.
// The token / cost / usage facts the response parser also emits stay so
// operators still get billing data on log-only mode.
func TestInvoke_CaptureCompletionOff_DoesNotEmitCompletion(t *testing.T) {
mw, err := Factory{}.New([]byte(`{"capture_completion":false}`))
require.NoError(t, err)
body := []byte(`{"id":"x","choices":[{"message":{"role":"assistant","content":"alice@example.com 123-45-6789"}}],"usage":{"prompt_tokens":10,"completion_tokens":20,"total_tokens":30}}`)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: body,
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "openai"}},
}
out, err := mw.Invoke(context.Background(), in)
require.NoError(t, err)
_, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
assert.False(t, ok, "llm.response_completion must NOT be emitted when capture_completion is false")
// Token facts must still flow.
_, ok = metaValue(out.Metadata, middleware.KeyLLMInputTokens)
assert.True(t, ok, "input tokens fact must still be emitted")
_, ok = metaValue(out.Metadata, middleware.KeyLLMOutputTokens)
assert.True(t, ok, "output tokens fact must still be emitted")
}
// TestInvoke_CaptureCompletionUnset_PreservesLegacyEmission documents the
// default behavior: empty config keeps emitting completion, so callers
// without the toggle aren't broken.
func TestInvoke_CaptureCompletionUnset_PreservesLegacyEmission(t *testing.T) {
mw, err := Factory{}.New([]byte(`{}`))
require.NoError(t, err)
body := []byte(`{"id":"x","choices":[{"message":{"role":"assistant","content":"hello"}}],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: body,
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "openai"}},
}
out, err := mw.Invoke(context.Background(), in)
require.NoError(t, err)
_, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
assert.True(t, ok, "absent capture_completion must preserve emission (backwards-compatible default)")
}
// TestInvoke_RedactPii_OffShipsRawCompletion covers the inverse: with
// redact_pii=false (default) the model output is shipped verbatim.
func TestInvoke_RedactPii_OffShipsRawCompletion(t *testing.T) {
mw, err := Factory{}.New(nil)
require.NoError(t, err)
body := []byte(`{"id":"x","choices":[{"message":{"role":"assistant","content":"alice@example.com 123-45-6789"}}],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
RespBody: body,
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "openai"}},
}
out, err := mw.Invoke(context.Background(), in)
require.NoError(t, err)
completion, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
require.True(t, ok)
assert.Contains(t, completion, "alice@example.com", "redact off → raw email passes through")
assert.Contains(t, completion, "123-45-6789", "redact off → raw SSN passes through")
assert.NotContains(t, completion, "[REDACTED:", "redact off → no markers")
}
func TestInvoke_CompletionTruncationRuneSafe(t *testing.T) {
rune4 := "\xf0\x9f\x98\x80" // 4-byte emoji
body := strings.Repeat("a", maxCompletionBytes-1) + rune4
require.Greater(t, len(body), maxCompletionBytes, "test setup must exceed the cap")
got := truncateCompletion(body)
assert.True(t, len(got) < maxCompletionBytes, "truncated bytes must drop the partial rune entirely")
assert.NotContains(t, got, "\x80", "truncated text must not end on a continuation byte")
}
@@ -0,0 +1,69 @@
package llm_response_parser
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/proxy/internal/middleware"
)
// TestInvoke_OpenAIResponsesStreaming is the regression guard for the live
// bug where Codex hits /v1/responses (the OpenAI Responses API), whose SSE
// shape differs from chat.completions: completion text rides
// response.output_text.delta and usage rides response.completed under
// response.usage. The old parser only knew the chat.completions shape, so
// resp_meta came back empty (no tokens, no cost).
func TestInvoke_OpenAIResponsesStreaming(t *testing.T) {
m := newTestMiddleware(t)
body := loadFixture(t, "openai_responses_stream.txt")
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "text/event-stream; charset=utf-8"}},
RespBody: body,
Metadata: []middleware.KV{
{Key: middleware.KeyLLMProvider, Value: "openai"},
{Key: middleware.KeyLLMModel, Value: "gpt-5.5"},
},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "Invoke must not error on a Responses-API streaming body")
inTok, ok := metaValue(out.Metadata, middleware.KeyLLMInputTokens)
require.True(t, ok, "input tokens must be emitted from a Responses-API stream")
assert.Equal(t, "123", inTok, "input_tokens must come from response.completed usage")
outTok, _ := metaValue(out.Metadata, middleware.KeyLLMOutputTokens)
assert.Equal(t, "45", outTok, "output_tokens must come from response.completed usage")
totTok, _ := metaValue(out.Metadata, middleware.KeyLLMTotalTokens)
assert.Equal(t, "168", totTok, "total_tokens must come from response.completed usage")
cached, ok := metaValue(out.Metadata, middleware.KeyLLMCachedInputTokens)
require.True(t, ok, "cached input tokens must surface from input_tokens_details")
assert.Equal(t, "40", cached, "cached_tokens subset must surface for cost discounting")
completion, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
require.True(t, ok, "completion must be emitted for Responses-API streams")
assert.Equal(t, "Hello, world!", completion, "output_text.delta events must concatenate")
}
// TestAccumulateOpenAIStream_ResponsesNoUsage confirms that a Responses-API
// stream with text but no terminal usage frame still yields the completion
// and leaves tokens at zero rather than erroring.
func TestAccumulateOpenAIStream_ResponsesNoUsage(t *testing.T) {
body := []byte(`event: response.output_text.delta
data: {"type":"response.output_text.delta","delta":"partial"}
`)
usage, completion := accumulateOpenAIStream(body)
assert.Equal(t, int64(0), usage.InputTokens, "no usage frame leaves input tokens at zero")
assert.Equal(t, int64(0), usage.OutputTokens, "no usage frame leaves output tokens at zero")
assert.Equal(t, "partial", completion, "output_text deltas accumulate even without a usage frame")
}
@@ -0,0 +1,270 @@
package llm_response_parser
import (
"bytes"
"encoding/json"
"errors"
"io"
"strings"
"github.com/netbirdio/netbird/proxy/internal/llm"
)
// openAIDoneSentinel is the OpenAI end-of-stream marker. The scanner
// stops once this data frame is observed.
const openAIDoneSentinel = "[DONE]"
// accumulateStream walks the SSE byte slice, dispatches per provider,
// and returns the running token-usage and concatenated completion text.
// Errors from the scanner short-circuit accumulation but never panic
// — partial results are returned for truncated bodies.
func accumulateStream(provider string, body []byte) (llm.Usage, string) {
switch provider {
case "openai":
return accumulateOpenAIStream(body)
case "anthropic":
return accumulateAnthropicStream(body)
case llm.ProviderNameBedrock:
return accumulateBedrockStream(body)
default:
return llm.Usage{}, ""
}
}
// openAIStreamUsage is the usage block shared by both OpenAI streaming
// envelopes. Pointer fields tell "absent" from zero; the chat.completions
// (prompt_/completion_) and Responses-API (input_/output_) names are both
// accepted so a single decode covers either endpoint.
type openAIStreamUsage struct {
PromptTokens *int64 `json:"prompt_tokens"`
CompletionTokens *int64 `json:"completion_tokens"`
InputTokens *int64 `json:"input_tokens"`
OutputTokens *int64 `json:"output_tokens"`
TotalTokens *int64 `json:"total_tokens"`
PromptTokensDetails *struct {
CachedTokens *int64 `json:"cached_tokens"`
} `json:"prompt_tokens_details"`
InputTokensDetails *struct {
CachedTokens *int64 `json:"cached_tokens"`
} `json:"input_tokens_details"`
}
// openAIStreamChunk matches both OpenAI streaming envelopes. The
// chat.completions chunk carries text in choices[].delta.content and a
// trailing top-level usage block. The Responses API (/v1/responses) emits
// typed events instead: completion text rides response.output_text.delta
// (top-level "delta" string) and the final usage rides response.completed
// under response.usage. Only fields used for accumulation are declared.
type openAIStreamChunk struct {
Choices []struct {
Delta struct {
Content string `json:"content"`
} `json:"delta"`
} `json:"choices"`
Usage *openAIStreamUsage `json:"usage"`
Type string `json:"type"`
Delta json.RawMessage `json:"delta"`
Response *struct {
Usage *openAIStreamUsage `json:"usage"`
} `json:"response"`
}
// accumulateOpenAIStream sums per-chunk content deltas and lifts the usage
// block off the final frame, handling both the chat.completions and the
// Responses-API event shapes. Clients without stream_options.include_usage
// (chat.completions) and any provider that omits the final usage simply
// leave tokens at zero; the caller chooses what to emit.
func accumulateOpenAIStream(body []byte) (llm.Usage, string) {
var (
usage llm.Usage
completion strings.Builder
)
scanner := llm.NewScanner(bytes.NewReader(body))
for {
ev, err := scanner.Next()
if err != nil {
break
}
if ev.Data == openAIDoneSentinel {
break
}
if ev.Data == "" {
continue
}
var chunk openAIStreamChunk
if err := json.Unmarshal([]byte(ev.Data), &chunk); err != nil {
continue
}
for _, c := range chunk.Choices {
completion.WriteString(c.Delta.Content)
}
if chunk.Type == "response.output_text.delta" {
if s, ok := decodeJSONString(chunk.Delta); ok {
completion.WriteString(s)
}
}
u := chunk.Usage
if u == nil && chunk.Response != nil {
u = chunk.Response.Usage
}
applyOpenAIStreamUsage(u, &usage)
}
return usage, completion.String()
}
// applyOpenAIStreamUsage lifts the token counts off a final-frame usage
// block into the running usage, normalising the chat.completions
// (prompt_/completion_) and Responses-API (input_/output_) names and
// backfilling total tokens when the provider omits them.
func applyOpenAIStreamUsage(u *openAIStreamUsage, usage *llm.Usage) {
if u == nil {
return
}
usage.InputTokens = pickInt64(u.InputTokens, u.PromptTokens)
usage.OutputTokens = pickInt64(u.OutputTokens, u.CompletionTokens)
usage.TotalTokens = derefInt64(u.TotalTokens)
if u.InputTokensDetails != nil {
if v := derefInt64(u.InputTokensDetails.CachedTokens); v > 0 {
usage.CachedInputTokens = v
}
}
if usage.CachedInputTokens == 0 && u.PromptTokensDetails != nil {
usage.CachedInputTokens = derefInt64(u.PromptTokensDetails.CachedTokens)
}
if usage.TotalTokens == 0 && (usage.InputTokens > 0 || usage.OutputTokens > 0) {
usage.TotalTokens = usage.InputTokens + usage.OutputTokens
}
}
// decodeJSONString unmarshals a JSON-encoded string value, returning
// ok=false when the raw message is empty or not a string.
func decodeJSONString(raw json.RawMessage) (string, bool) {
if len(raw) == 0 {
return "", false
}
var s string
if err := json.Unmarshal(raw, &s); err != nil {
return "", false
}
return s, true
}
// anthropicStreamEvent captures the union of Messages-API stream event
// payloads we care about. Each named event on the wire fills only its
// shape's fields; unknown keys are ignored.
type anthropicStreamUsage struct {
InputTokens *int64 `json:"input_tokens"`
OutputTokens *int64 `json:"output_tokens"`
CacheReadInputTokens *int64 `json:"cache_read_input_tokens"`
CacheCreationInputTokens *int64 `json:"cache_creation_input_tokens"`
}
type anthropicStreamEvent struct {
Type string `json:"type"`
Message *struct {
Usage *anthropicStreamUsage `json:"usage"`
} `json:"message"`
Delta *struct {
Type string `json:"type"`
Text string `json:"text"`
} `json:"delta"`
Usage *anthropicStreamUsage `json:"usage"`
}
// accumulateAnthropicStream tracks input_tokens from message_start,
// output_tokens from message_delta, and concatenates text_delta payloads
// from content_block_delta events. Final usage prefers message_delta
// values which carry the post-completion totals.
func accumulateAnthropicStream(body []byte) (llm.Usage, string) {
var (
usage llm.Usage
completion strings.Builder
)
scanner := llm.NewScanner(bytes.NewReader(body))
for {
ev, err := scanner.Next()
if err != nil {
if errors.Is(err, io.EOF) {
break
}
break
}
if ev.Data == "" {
continue
}
var payload anthropicStreamEvent
if err := json.Unmarshal([]byte(ev.Data), &payload); err != nil {
continue
}
eventType := ev.Type
if eventType == "" {
eventType = payload.Type
}
applyAnthropicStreamEvent(eventType, payload, &usage, &completion)
}
if usage.InputTokens > 0 || usage.OutputTokens > 0 {
usage.TotalTokens = usage.InputTokens + usage.OutputTokens + usage.CachedInputTokens + usage.CacheCreationTokens
}
return usage, completion.String()
}
// applyAnthropicStreamEvent folds one parsed Anthropic Messages stream event
// into the running usage/completion. Shared by the SSE accumulator and the
// Bedrock InvokeModel event-stream, whose chunks wrap the same event JSON.
func applyAnthropicStreamEvent(eventType string, payload anthropicStreamEvent, usage *llm.Usage, completion *strings.Builder) {
switch eventType {
case "message_start":
if payload.Message != nil {
applyAnthropicStreamUsage(payload.Message.Usage, usage)
}
case "content_block_delta":
if payload.Delta != nil && payload.Delta.Type == "text_delta" {
completion.WriteString(payload.Delta.Text)
}
case "message_delta":
applyAnthropicStreamUsage(payload.Usage, usage)
case "message_stop":
// No-op; Anthropic does not emit usage here.
}
}
// applyAnthropicStreamUsage folds a non-nil Anthropic usage block into the
// running totals. Each field overwrites only when present and positive, so
// message_delta's post-completion counts supersede the message_start seed
// without zeroing dimensions a later event omits.
func applyAnthropicStreamUsage(u *anthropicStreamUsage, usage *llm.Usage) {
if u == nil {
return
}
if v := derefInt64(u.InputTokens); v > 0 {
usage.InputTokens = v
}
if v := derefInt64(u.OutputTokens); v > 0 {
usage.OutputTokens = v
}
if v := derefInt64(u.CacheReadInputTokens); v > 0 {
usage.CachedInputTokens = v
}
if v := derefInt64(u.CacheCreationInputTokens); v > 0 {
usage.CacheCreationTokens = v
}
}
func pickInt64(preferred, fallback *int64) int64 {
if preferred != nil {
return *preferred
}
return derefInt64(fallback)
}
func derefInt64(v *int64) int64 {
if v == nil {
return 0
}
return *v
}
@@ -0,0 +1,110 @@
package llm_response_parser
import (
"bytes"
"encoding/json"
"strings"
"github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream"
"github.com/netbirdio/netbird/proxy/internal/llm"
)
// bedrockEventTypeHeader names each AWS event-stream frame's event type.
const bedrockEventTypeHeader = ":event-type"
// accumulateBedrockStream decodes the AWS binary event-stream returned by
// Bedrock's streaming endpoints and folds it into running usage/completion.
// Two framings are handled:
// - InvokeModel (invoke-with-response-stream): each "chunk" frame's payload is
// {"bytes":"<base64>"} wrapping a vendor-native (Anthropic) stream event.
// - Converse (converse-stream): native frames (contentBlockDelta, metadata, …)
// whose payload JSON carries text deltas and a final usage block.
//
// A truncated stream (cut at the capture cap) decodes best-effort: frames up to
// the cut are applied and the partial usage is returned.
func accumulateBedrockStream(body []byte) (llm.Usage, string) {
var (
usage llm.Usage
completion strings.Builder
)
dec := eventstream.NewDecoder()
r := bytes.NewReader(body)
for {
msg, err := dec.Decode(r, nil)
if err != nil {
break // EOF or a partial trailing frame — return what we have.
}
eventType := ""
if v := msg.Headers.Get(bedrockEventTypeHeader); v != nil {
eventType = v.String()
}
if eventType == "chunk" {
applyBedrockInvokeChunk(msg.Payload, &usage, &completion)
continue
}
applyConverseStreamEvent(eventType, msg.Payload, &usage, &completion)
}
if usage.TotalTokens == 0 && (usage.InputTokens > 0 || usage.OutputTokens > 0) {
usage.TotalTokens = usage.InputTokens + usage.OutputTokens + usage.CachedInputTokens + usage.CacheCreationTokens
}
return usage, completion.String()
}
// applyBedrockInvokeChunk decodes an InvokeModel stream "chunk" frame
// ({"bytes":"<base64 anthropic event>"}) and folds the wrapped Anthropic event
// into usage/completion via the shared accumulator.
func applyBedrockInvokeChunk(payload []byte, usage *llm.Usage, completion *strings.Builder) {
var wrap struct {
Bytes []byte `json:"bytes"` // base64 string — encoding/json decodes it
}
if err := json.Unmarshal(payload, &wrap); err != nil || len(wrap.Bytes) == 0 {
return
}
var ev anthropicStreamEvent
if err := json.Unmarshal(wrap.Bytes, &ev); err != nil {
return
}
applyAnthropicStreamEvent(ev.Type, ev, usage, completion)
}
// converseStreamEvent captures the Converse stream frames carrying completion
// text (contentBlockDelta) and the final token usage (metadata).
type converseStreamEvent struct {
Delta *struct {
Text string `json:"text"`
} `json:"delta"`
Usage *struct {
InputTokens int64 `json:"inputTokens"`
OutputTokens int64 `json:"outputTokens"`
TotalTokens int64 `json:"totalTokens"`
} `json:"usage"`
}
// applyConverseStreamEvent folds one native Converse stream frame into the
// running usage/completion: contentBlockDelta carries assistant text, and the
// trailing metadata frame carries the final usage block.
func applyConverseStreamEvent(eventType string, payload []byte, usage *llm.Usage, completion *strings.Builder) {
var ev converseStreamEvent
if err := json.Unmarshal(payload, &ev); err != nil {
return
}
switch eventType {
case "contentBlockDelta":
if ev.Delta != nil {
completion.WriteString(ev.Delta.Text)
}
case "metadata":
if ev.Usage != nil {
if ev.Usage.InputTokens > 0 {
usage.InputTokens = ev.Usage.InputTokens
}
if ev.Usage.OutputTokens > 0 {
usage.OutputTokens = ev.Usage.OutputTokens
}
if ev.Usage.TotalTokens > 0 {
usage.TotalTokens = ev.Usage.TotalTokens
}
}
}
}
@@ -0,0 +1,74 @@
package llm_response_parser
import (
"bytes"
"encoding/base64"
"encoding/json"
"testing"
"github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream"
"github.com/stretchr/testify/require"
)
// bedrockFrame encodes a single AWS event-stream frame with the given
// :event-type header and JSON payload, mirroring what Bedrock sends.
func bedrockFrame(t *testing.T, eventType string, payload []byte) []byte {
t.Helper()
var buf bytes.Buffer
enc := eventstream.NewEncoder()
err := enc.Encode(&buf, eventstream.Message{
Headers: eventstream.Headers{{Name: ":event-type", Value: eventstream.StringValue(eventType)}},
Payload: payload,
})
require.NoError(t, err, "encode event-stream frame")
return buf.Bytes()
}
func mustJSON(t *testing.T, v any) []byte {
t.Helper()
b, err := json.Marshal(v)
require.NoError(t, err)
return b
}
func TestAccumulateBedrockStream_Invoke(t *testing.T) {
// invoke-with-response-stream: each "chunk" frame wraps a base64-encoded
// Anthropic stream event under {"bytes": ...}.
events := [][]byte{
mustJSON(t, map[string]any{"type": "message_start", "message": map[string]any{"usage": map[string]any{"input_tokens": 13}}}),
mustJSON(t, map[string]any{"type": "content_block_delta", "delta": map[string]any{"type": "text_delta", "text": "po"}}),
mustJSON(t, map[string]any{"type": "content_block_delta", "delta": map[string]any{"type": "text_delta", "text": "ng"}}),
mustJSON(t, map[string]any{"type": "message_delta", "usage": map[string]any{"output_tokens": 5}}),
}
var body bytes.Buffer
for _, ev := range events {
wrap := mustJSON(t, map[string]any{"bytes": base64.StdEncoding.EncodeToString(ev)})
body.Write(bedrockFrame(t, "chunk", wrap))
}
usage, completion := accumulateBedrockStream(body.Bytes())
require.Equal(t, int64(13), usage.InputTokens, "input tokens from message_start")
require.Equal(t, int64(5), usage.OutputTokens, "output tokens from message_delta")
require.Equal(t, int64(18), usage.TotalTokens, "total is additive")
require.Equal(t, "pong", completion, "text deltas concatenated")
}
func TestAccumulateBedrockStream_Converse(t *testing.T) {
var body bytes.Buffer
body.Write(bedrockFrame(t, "contentBlockDelta", mustJSON(t, map[string]any{"delta": map[string]any{"text": "po"}})))
body.Write(bedrockFrame(t, "contentBlockDelta", mustJSON(t, map[string]any{"delta": map[string]any{"text": "ng"}})))
body.Write(bedrockFrame(t, "metadata", mustJSON(t, map[string]any{"usage": map[string]any{"inputTokens": 11, "outputTokens": 3, "totalTokens": 14}})))
usage, completion := accumulateBedrockStream(body.Bytes())
require.Equal(t, int64(11), usage.InputTokens, "input tokens from metadata frame")
require.Equal(t, int64(3), usage.OutputTokens, "output tokens from metadata frame")
require.Equal(t, int64(14), usage.TotalTokens, "total from metadata frame")
require.Equal(t, "pong", completion, "converse text deltas concatenated")
}
func TestAccumulateBedrockStream_Truncated(t *testing.T) {
// A body cut mid-frame must not panic; partial usage is returned.
full := bedrockFrame(t, "metadata", mustJSON(t, map[string]any{"usage": map[string]any{"inputTokens": 11, "outputTokens": 3}}))
usage, _ := accumulateBedrockStream(full[:len(full)-4])
require.Zero(t, usage.OutputTokens, "truncated trailing frame is dropped, not panicked on")
}
@@ -0,0 +1,169 @@
package llm_response_parser
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/proxy/internal/middleware"
)
func TestInvoke_OpenAIStreamingWithUsage(t *testing.T) {
m := newTestMiddleware(t)
body := loadFixture(t, "openai_stream.txt")
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "text/event-stream"}},
RespBody: body,
Metadata: []middleware.KV{
{Key: middleware.KeyLLMProvider, Value: "openai"},
{Key: middleware.KeyLLMModel, Value: "gpt-4o-mini"},
},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "Invoke must not error on streaming OpenAI body")
in123, _ := metaValue(out.Metadata, middleware.KeyLLMInputTokens)
assert.Equal(t, "123", in123, "input tokens must come from final-chunk usage block")
outTok, _ := metaValue(out.Metadata, middleware.KeyLLMOutputTokens)
assert.Equal(t, "45", outTok, "output tokens must come from final-chunk usage block")
totTok, _ := metaValue(out.Metadata, middleware.KeyLLMTotalTokens)
assert.Equal(t, "168", totTok, "total tokens must come from final-chunk usage block")
completion, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
require.True(t, ok, "completion must be emitted for streaming responses")
assert.Equal(t, "Hello, world!", completion, "deltas must concatenate into the buffered fixture's text")
}
func TestInvoke_OpenAIStreamingWithoutUsage(t *testing.T) {
body := []byte(`data: {"choices":[{"delta":{"content":"Hi"}}]}
data: {"choices":[{"delta":{"content":" there"}}]}
data: [DONE]
`)
usage, completion := accumulateOpenAIStream(body)
assert.Equal(t, int64(0), usage.InputTokens, "input tokens must stay zero without a usage frame")
assert.Equal(t, int64(0), usage.OutputTokens, "output tokens must stay zero without a usage frame")
assert.Equal(t, int64(0), usage.TotalTokens, "total tokens must stay zero without a usage frame")
assert.Equal(t, "Hi there", completion, "deltas must still accumulate when usage is absent")
}
func TestInvoke_OpenAIStreamingNoUsage_OmitsUsageMetadata(t *testing.T) {
m := newTestMiddleware(t)
body := []byte(`data: {"choices":[{"delta":{"content":"Hello"}}]}
data: [DONE]
`)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "text/event-stream"}},
RespBody: body,
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "openai"}},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "Invoke must not error on usage-less streams")
_, hasIn := metaValue(out.Metadata, middleware.KeyLLMInputTokens)
_, hasOut := metaValue(out.Metadata, middleware.KeyLLMOutputTokens)
_, hasTot := metaValue(out.Metadata, middleware.KeyLLMTotalTokens)
assert.False(t, hasIn, "input tokens omitted when no usage frame")
assert.False(t, hasOut, "output tokens omitted when no usage frame")
assert.False(t, hasTot, "total tokens omitted when no usage frame")
completion, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
require.True(t, ok, "completion must still be emitted from deltas")
assert.Equal(t, "Hello", completion, "completion must come from delta accumulation")
}
func TestInvoke_AnthropicStreaming(t *testing.T) {
m := newTestMiddleware(t)
body := loadFixture(t, "anthropic_stream.txt")
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "text/event-stream"}},
RespBody: body,
Metadata: []middleware.KV{
{Key: middleware.KeyLLMProvider, Value: "anthropic"},
{Key: middleware.KeyLLMModel, Value: "claude-sonnet-4-5"},
},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "Invoke must not error on streaming Anthropic body")
in123, _ := metaValue(out.Metadata, middleware.KeyLLMInputTokens)
assert.Equal(t, "123", in123, "input tokens must come from message_start usage")
outTok, _ := metaValue(out.Metadata, middleware.KeyLLMOutputTokens)
assert.Equal(t, "45", outTok, "output tokens must come from message_delta usage")
totTok, _ := metaValue(out.Metadata, middleware.KeyLLMTotalTokens)
assert.Equal(t, "168", totTok, "total tokens must be input+output for anthropic streaming")
completion, ok := metaValue(out.Metadata, middleware.KeyLLMResponseCompletion)
require.True(t, ok, "completion must be emitted from text_delta accumulation")
assert.Equal(t, "Hello, world!", completion, "anthropic streaming text must accumulate across content_block_delta events")
}
func TestInvoke_StreamingTruncatedBody_BestEffort(t *testing.T) {
m := newTestMiddleware(t)
full := loadFixture(t, "anthropic_stream.txt")
cut := len(full) / 2
truncated := full[:cut]
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "text/event-stream"}},
RespBody: truncated,
RespBodyTruncated: true,
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "anthropic"}},
}
require.NotPanics(t, func() {
_, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "truncated streaming body must not surface as error")
}, "Invoke must never panic on a truncated SSE body")
}
func TestInvoke_StreamingEmptyBody(t *testing.T) {
m := newTestMiddleware(t)
in := &middleware.Input{
Slot: middleware.SlotOnResponse,
Status: 200,
RespHeaders: []middleware.KV{{Key: "Content-Type", Value: "text/event-stream"}},
RespBody: nil,
Metadata: []middleware.KV{{Key: middleware.KeyLLMProvider, Value: "openai"}},
}
out, err := m.Invoke(context.Background(), in)
require.NoError(t, err, "empty SSE body must not surface as error")
assert.Empty(t, out.Metadata, "no metadata for empty SSE body")
}
func TestAccumulateAnthropicStream_PartialUsage(t *testing.T) {
body := []byte(`event: message_start
data: {"type":"message_start","message":{"usage":{"input_tokens":10}}}
event: content_block_delta
data: {"type":"content_block_delta","delta":{"type":"text_delta","text":"hi"}}
`)
usage, completion := accumulateAnthropicStream(body)
assert.Equal(t, int64(10), usage.InputTokens, "partial input_tokens must survive truncated stream")
assert.Equal(t, int64(0), usage.OutputTokens, "output_tokens stays zero without message_delta")
assert.Equal(t, "hi", completion, "completion must come from observed text_delta events")
}