From e93eac6f1611ef358bd342f9f7d46675f20b5e3d Mon Sep 17 00:00:00 2001 From: Ramapitecus Date: Sat, 12 Sep 2026 18:21:22 +0500 Subject: [PATCH 1/9] fix: preserve billing dimensions and deliver replayable usage events --- internal/api/handlers/management/api_tools.go | 30 +++ internal/api/handlers/management/usage.go | 33 +++ internal/api/server.go | 5 + internal/api/server_management.go | 2 + internal/api/server_routes.go | 7 + internal/api/usage_coverage.go | 34 +++ internal/client/codex/live/sideband.go | 22 +- internal/client/codex/live/usage.go | 100 +++++++++ internal/client/codex/live/usage_test.go | 20 ++ internal/client/codex/live/websocket.go | 4 +- internal/redisqueue/journal.go | 202 ++++++++++++++++++ internal/redisqueue/journal_test.go | 78 +++++++ internal/redisqueue/plugin.go | 55 ++++- internal/redisqueue/queue.go | 4 +- internal/redisqueue/usage_integrity_test.go | 18 ++ .../executor/codex_websockets_execute.go | 1 + .../executor/codex_websockets_stream.go | 2 + .../codex_websockets_usage_integrity_test.go | 96 +++++++++ .../helps/openai_compat_tool_results.go | 8 + .../executor/helps/plugin_executor_usage.go | 23 ++ .../runtime/executor/helps/usage_helpers.go | 71 ++++-- .../executor/helps/usage_integrity_test.go | 32 +++ .../executor/helps/usage_measurements.go | 65 ++++++ ...penai_compat_executor_tool_results_test.go | 18 +- .../runtime/executor/xai_executor_media.go | 22 +- .../runtime/executor/xai_executor_test.go | 4 +- .../openai/openai_responses_websocket.go | 4 + sdk/cliproxy/usage/manager.go | 40 ++++ sdk/cliproxy/usage/scope.go | 35 +++ sdk/cliproxy/usage/scope_test.go | 38 ++++ 30 files changed, 1038 insertions(+), 35 deletions(-) create mode 100644 internal/api/usage_coverage.go create mode 100644 internal/client/codex/live/usage.go create mode 100644 internal/client/codex/live/usage_test.go create mode 100644 internal/redisqueue/journal.go create mode 100644 internal/redisqueue/journal_test.go create mode 100644 internal/redisqueue/usage_integrity_test.go create mode 100644 internal/runtime/executor/codex_websockets_usage_integrity_test.go create mode 100644 internal/runtime/executor/helps/usage_integrity_test.go create mode 100644 internal/runtime/executor/helps/usage_measurements.go create mode 100644 sdk/cliproxy/usage/scope.go create mode 100644 sdk/cliproxy/usage/scope_test.go diff --git a/internal/api/handlers/management/api_tools.go b/internal/api/handlers/management/api_tools.go index a619afd1076..9572f42f2e0 100644 --- a/internal/api/handlers/management/api_tools.go +++ b/internal/api/handlers/management/api_tools.go @@ -4,6 +4,8 @@ import ( "context" "encoding/json" "fmt" + "github.com/router-for-me/CLIProxyAPI/v7/internal/runtime/executor/helps" + "github.com/tidwall/gjson" "io" "net/http" "net/url" @@ -189,8 +191,24 @@ func (h *Handler) APICall(c *gin.Context) { } httpClient.Transport = h.apiCallTransport(auth, requestProxyURL) + // Management model probes bypass normal executors and need their own record. + var reporter *helps.UsageReporter + usageCtx := context.WithValue(c.Request.Context(), "gin", c) + model := gjson.Get(body.Data, "model").String() + if method == http.MethodPost && model != "" { + provider := "openai-compatible" + if auth != nil { + provider = auth.Provider + } else if req.URL.Hostname() == "api.x.ai" { + provider = "xai" + } + reporter = helps.NewUsageReporter(usageCtx, provider, model, auth) + reporter.SetOperation("management_call", method+" "+req.URL.Path) + defer reporter.EnsurePublished(usageCtx) + } resp, errDo := httpClient.Do(req) if errDo != nil { + reporter.PublishFailure(usageCtx, errDo) log.WithError(errDo).Debug("management APICall request failed") c.JSON(http.StatusBadGateway, gin.H{"error": "request failed"}) return @@ -203,10 +221,22 @@ func (h *Handler) APICall(c *gin.Context) { respBody, errReadAll := io.ReadAll(resp.Body) if errReadAll != nil { + reporter.PublishFailure(usageCtx, errReadAll) c.JSON(http.StatusBadGateway, gin.H{"error": "failed to read response"}) return } + if reporter != nil { + detail := helps.ParseOpenAIUsage(respBody) + if strings.HasSuffix(req.URL.Path, "/messages") { + detail = helps.ParseClaudeUsage(respBody) + } + if resp.StatusCode >= 400 { + reporter.PublishFailureWithDetail(usageCtx, detail, fmt.Errorf("upstream HTTP %d", resp.StatusCode)) + } else { + reporter.Publish(usageCtx, detail) + } + } c.JSON(http.StatusOK, apiCallResponse{ StatusCode: resp.StatusCode, Header: resp.Header, diff --git a/internal/api/handlers/management/usage.go b/internal/api/handlers/management/usage.go index c1602c0423e..54164ae4251 100644 --- a/internal/api/handlers/management/usage.go +++ b/internal/api/handlers/management/usage.go @@ -53,3 +53,36 @@ func parseUsageQueueCount(value string) (int, error) { } return count, nil } + +// GetUsageJournal returns replayable events without deleting them. +func (h *Handler) GetUsageJournal(c *gin.Context) { + count, errCount := parseUsageQueueCount(c.Query("count")) + if errCount != nil || count > 1000 { + c.JSON(http.StatusBadRequest, gin.H{"error": "count must be between 1 and 1000"}) + return + } + items, errRead := redisqueue.ReadUsageJournal(count) + if errRead != nil { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": errRead.Error()}) + return + } + records := make([]usageQueueRecord, 0, len(items)) + for _, item := range items { + records = append(records, usageQueueRecord(item)) + } + c.JSON(http.StatusOK, records) +} +func (h *Handler) AckUsageJournal(c *gin.Context) { + var body struct { + IDs []string `json:"event_ids"` + } + if errBind := c.ShouldBindJSON(&body); errBind != nil || len(body.IDs) > 1000 { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid event_ids"}) + return + } + if errAck := redisqueue.AckUsageJournal(body.IDs); errAck != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": errAck.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{"acknowledged": len(body.IDs)}) +} diff --git a/internal/api/server.go b/internal/api/server.go index 8108d5bf9e7..d000edc760a 100644 --- a/internal/api/server.go +++ b/internal/api/server.go @@ -12,6 +12,7 @@ import ( "net" "net/http" "os" + "path/filepath" "strings" "sync" "sync/atomic" @@ -158,6 +159,7 @@ func NewServer(cfg *config.Config, authManager *auth.Manager, accessManager *sdk } engine.Use(corsMiddleware()) + engine.Use(usageCoverageMiddleware()) wd, err := os.Getwd() if err != nil { wd = configFilePath @@ -234,6 +236,9 @@ func NewServer(cfg *config.Config, authManager *auth.Manager, accessManager *sdk // or when a local management password is provided (e.g. TUI mode). hasManagementSecret := cfg.RemoteManagement.SecretKey != "" || envManagementSecret || s.localPassword != "" s.managementRoutesEnabled.Store(hasManagementSecret) + if strings.TrimSpace(configFilePath) != "" { + redisqueue.ConfigureUsageJournal(filepath.Join(filepath.Dir(configFilePath), "usage-journal")) + } redisqueue.SetEnabled(hasManagementSecret || (cfg != nil && cfg.Home.Enabled)) if hasManagementSecret { s.registerManagementRoutes() diff --git a/internal/api/server_management.go b/internal/api/server_management.go index 247bd3e1967..67c53acb547 100644 --- a/internal/api/server_management.go +++ b/internal/api/server_management.go @@ -82,6 +82,8 @@ func (s *Server) registerManagementRoutes() { mgmt.DELETE("/api-keys", s.mgmt.DeleteAPIKeys) mgmt.GET("/api-key-usage", s.mgmt.GetAPIKeyUsage) mgmt.GET("/usage-queue", s.mgmt.GetUsageQueue) + mgmt.GET("/usage-journal", s.mgmt.GetUsageJournal) + mgmt.POST("/usage-journal/ack", s.mgmt.AckUsageJournal) mgmt.GET("/gemini-api-key", s.mgmt.GetGeminiKeys) mgmt.PUT("/gemini-api-key", s.mgmt.PutGeminiKeys) diff --git a/internal/api/server_routes.go b/internal/api/server_routes.go index 6483fb6161e..a28b0ab1063 100644 --- a/internal/api/server_routes.go +++ b/internal/api/server_routes.go @@ -502,6 +502,13 @@ func (s *Server) codexAlphaSearch(c *gin.Context) { c.JSON(clienterror.HTTPStatusFromErrorOr(err, http.StatusBadGateway), gin.H{"error": "Failed to read Codex search response"}) return } + reporter := helps.NewUsageReporter(ctx, "codex", selectionModel, selected) + detail := helps.ParseOpenAIUsage(upstreamBody) + if resp.StatusCode >= 400 { + reporter.PublishFailureWithDetail(ctx, detail, fmt.Errorf("search returned HTTP %d", resp.StatusCode)) + } else { + reporter.Publish(ctx, detail) + } helps.AppendAPIResponseChunk(ctx, s.cfg, upstreamBody) if selection != nil && resp.StatusCode == http.StatusUnauthorized { s.handlers.AuthManager.ReportHomeUnauthorized(ctx, selected, "codex", selectionModel, upstreamBody) diff --git a/internal/api/usage_coverage.go b/internal/api/usage_coverage.go new file mode 100644 index 00000000000..b51d2b7961d --- /dev/null +++ b/internal/api/usage_coverage.go @@ -0,0 +1,34 @@ +package api + +import ( + "context" + "net/http" + "strings" + "time" + + "github.com/gin-gonic/gin" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" +) + +// coverageMiddleware records operations whose handlers have no token reporter. +// Control-plane calls and opaque WebRTC media are explicitly unmeasured; token +// estimates (for example count_tokens) are never recorded as consumed tokens. +func usageCoverageMiddleware() gin.HandlerFunc { + return func(c *gin.Context) { + path := c.Request.URL.Path + covered := strings.HasPrefix(path, "/v1/") || strings.HasPrefix(path, "/v1beta/") || strings.HasPrefix(path, "/backend-api/codex/") || strings.HasPrefix(path, "/openai/v1/") + if !covered || strings.HasSuffix(path, "/models") { + c.Next() + return + } + ctx, scope := usage.WithAccountingScope(c.Request.Context()) + ctx = usage.WithNewGeneration(ctx) + c.Request = c.Request.WithContext(ctx) + started := time.Now() + c.Next() + if scope.Published() { + return + } + usage.PublishRecord(context.WithValue(ctx, "gin", c), usage.Record{Provider: "unknown", Model: "unknown", ExecutorType: "EndpointCoverage", Endpoint: c.Request.Method + " " + path, Kind: "unmeasured", Generate: usage.GenerateFlag(false), RequestedAt: started, Latency: time.Since(started), Failed: c.Writer.Status() >= http.StatusBadRequest, Fail: usage.Failure{StatusCode: c.Writer.Status()}}) + } +} diff --git a/internal/client/codex/live/sideband.go b/internal/client/codex/live/sideband.go index 0f9b27509c0..da472f8a297 100644 --- a/internal/client/codex/live/sideband.go +++ b/internal/client/codex/live/sideband.go @@ -509,7 +509,9 @@ func (h *Handler) HandleSideband(c *gin.Context) { } consumeSession = true - if errRelay := relayWebsockets(downstream, upstream); errRelay != nil && !isNormalWebsocketClose(errRelay) { + observer := newLiveUsageObserver(ctx, selected, session.model) + defer observer.close() + if errRelay := relayWebsockets(downstream, upstream, observer.observe); errRelay != nil && !isNormalWebsocketClose(errRelay) { helps.RecordAPIWebsocketError(ctx, runtimeConfig, "relay", errRelay) log.WithError(errRelay).Debug("codex live sideband relay closed") } @@ -632,10 +634,10 @@ func websocketCloseFunc(name string, conn *websocket.Conn) func() error { } } -func relayWebsockets(downstream, upstream *websocket.Conn) error { +func relayWebsockets(downstream, upstream *websocket.Conn, observers ...func([]byte)) error { results := make(chan error, 2) go func() { results <- copyWebsocket(upstream, downstream) }() - go func() { results <- copyWebsocket(downstream, upstream) }() + go func() { results <- copyWebsocket(downstream, upstream, observers...) }() firstErr := <-results closeCode, closeReason := websocketCloseDetails(firstErr) @@ -648,7 +650,7 @@ func relayWebsockets(downstream, upstream *websocket.Conn) error { return firstErr } -func copyWebsocket(destination, source *websocket.Conn) error { +func copyWebsocket(destination, source *websocket.Conn, observers ...func([]byte)) error { for { messageType, reader, errReader := source.NextReader() if errReader != nil { @@ -658,7 +660,17 @@ func copyWebsocket(destination, source *websocket.Conn) error { if errWriter != nil { return errWriter } - _, errCopy := io.Copy(writer, reader) + capture := &usageFrameCapture{} + var target io.Writer = writer + if len(observers) > 0 && messageType == websocket.TextMessage { + target = io.MultiWriter(writer, capture) + } + _, errCopy := io.Copy(target, reader) + if errCopy == nil && !capture.overflow && len(capture.data) > 0 { + for _, observe := range observers { + observe(capture.data) + } + } errClose := writer.Close() if errCopy != nil { return errCopy diff --git a/internal/client/codex/live/usage.go b/internal/client/codex/live/usage.go new file mode 100644 index 00000000000..3aed1baae21 --- /dev/null +++ b/internal/client/codex/live/usage.go @@ -0,0 +1,100 @@ +package live + +import ( + "context" + "errors" + + "github.com/router-for-me/CLIProxyAPI/v7/internal/runtime/executor/helps" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" + "github.com/tidwall/gjson" +) + +// Observe only upstream data events. Client frames and repeated response.done +// notifications must never manufacture or double-count provider usage. +type liveUsageObserver struct { + ctx context.Context + auth *auth.Auth + model string + seen map[string]bool + pending map[string]*helps.UsageReporter +} + +func newLiveUsageObserver(ctx context.Context, selected *auth.Auth, model string) *liveUsageObserver { + return &liveUsageObserver{ctx: ctx, auth: selected, model: model, seen: map[string]bool{}, pending: map[string]*helps.UsageReporter{}} +} +func liveUsageDetail(payload []byte) (usage.Detail, bool) { + root := gjson.ParseBytes(payload) + switch root.Get("type").String() { + case "response.done", "response.completed", "response.incomplete": + return helps.ParseCodexUsage(payload) + case "conversation.item.input_audio_transcription.completed": + d := helps.ParseOpenAIUsage(payload) + return d, d.UsageObserved + } + return usage.Detail{}, false +} +func (o *liveUsageObserver) observe(payload []byte) { + root := gjson.ParseBytes(payload) + if model := root.Get("session.model").String(); model != "" { + o.model = model + } + kind := root.Get("type").String() + id := root.Get("response.id").String() + if id == "" { + id = root.Get("item_id").String() + } + if kind == "response.created" && id != "" && !o.seen[id] { + r := helps.NewUsageReporter(usage.WithNewGeneration(o.ctx), "codex", o.model, o.auth) + r.SetStream(true) + r.SetTransport("websocket") + o.pending[id] = r + return + } + terminal := kind == "response.done" || kind == "response.completed" || kind == "response.incomplete" || kind == "conversation.item.input_audio_transcription.completed" + if !terminal || id != "" && o.seen[id] { + return + } + detail, _ := liveUsageDetail(payload) + model := root.Get("response.model").String() + if model == "" { + model = o.model + } + r := o.pending[id] + if r == nil { + r = helps.NewUsageReporter(usage.WithNewGeneration(o.ctx), "codex", model, o.auth) + r.SetStream(true) + r.SetTransport("websocket") + } + ctx := usage.WithNewGeneration(o.ctx) + status := root.Get("response.status").String() + if status == "failed" || status == "cancelled" || status == "incomplete" { + r.PublishFailureWithDetail(ctx, detail, errors.New("realtime response "+status)) + } else { + r.Publish(ctx, detail) + } + if id != "" { + o.seen[id] = true + delete(o.pending, id) + } +} +func (o *liveUsageObserver) close() { + for _, r := range o.pending { + r.PublishFailure(o.ctx, errors.New("realtime connection closed before terminal usage")) + } +} + +// Bounded capture never limits forwarding of large media frames. +type usageFrameCapture struct { + data []byte + overflow bool +} + +func (b *usageFrameCapture) Write(p []byte) (int, error) { + if len(b.data)+len(p) > 1<<20 { + b.overflow = true + } else if !b.overflow { + b.data = append(b.data, p...) + } + return len(p), nil +} diff --git a/internal/client/codex/live/usage_test.go b/internal/client/codex/live/usage_test.go new file mode 100644 index 00000000000..ed3e11e1cc2 --- /dev/null +++ b/internal/client/codex/live/usage_test.go @@ -0,0 +1,20 @@ +package live + +import "testing" + +func TestLiveUsageTerminalDimensions(t *testing.T) { + d, ok := liveUsageDetail([]byte(`{"type":"response.done","response":{"usage":{"input_tokens":100,"output_tokens":20,"input_token_details":{"audio_tokens":60,"text_tokens":40},"output_token_details":{"audio_tokens":20}}}}`)) + if !ok || !d.UsageObserved || d.InputTokens != 100 || d.OutputTokens != 20 || d.RawUsage == "" { + t.Fatalf("lost live dimensions: %+v", d) + } + if _, ok := liveUsageDetail([]byte(`{"type":"response.audio.delta","usage":{"input_tokens":100}}`)); ok { + t.Fatal("nonterminal counted") + } +} +func TestLiveUsageCaptureDoesNotLimitForwarding(t *testing.T) { + b := &usageFrameCapture{} + n, err := b.Write(make([]byte, 2<<20)) + if n != 2<<20 || err != nil || !b.overflow || len(b.data) != 0 { + t.Fatal("capture must be bounded without limiting stream") + } +} diff --git a/internal/client/codex/live/websocket.go b/internal/client/codex/live/websocket.go index 949f3970743..c24cdfbd186 100644 --- a/internal/client/codex/live/websocket.go +++ b/internal/client/codex/live/websocket.go @@ -232,7 +232,9 @@ func (h *Handler) HandleDirectWebsocket(c *gin.Context) { } } - if errRelay := relayWebsockets(downstream, upstream); errRelay != nil && !isNormalWebsocketClose(errRelay) { + observer := newLiveUsageObserver(ctx, selected, selectionModel) + defer observer.close() + if errRelay := relayWebsockets(downstream, upstream, observer.observe); errRelay != nil && !isNormalWebsocketClose(errRelay) { helps.RecordAPIWebsocketError(ctx, h.currentConfig(), "relay", errRelay) log.WithError(errRelay).Debug("codex realtime direct websocket relay closed") } diff --git a/internal/redisqueue/journal.go b/internal/redisqueue/journal.go new file mode 100644 index 00000000000..d476521c314 --- /dev/null +++ b/internal/redisqueue/journal.go @@ -0,0 +1,202 @@ +package redisqueue + +import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "os" + "path/filepath" + "runtime" + "sort" + "strings" + "sync" + "unicode" + + log "github.com/sirupsen/logrus" +) + +// The journal is independent of the legacy destructive queue. A single local +// accounting consumer acknowledges event IDs only after committing its inbox. +// Unacknowledged events do not expire and survive core restarts. +type usageJournal struct { + mu sync.Mutex + dir string + lastErr error +} + +var journal usageJournal + +func ConfigureUsageJournal(dir string) { + journal.mu.Lock() + defer journal.mu.Unlock() + journal.dir = dir + journal.lastErr = nil +} +func validEventID(id string) bool { + if len(id) == 0 || len(id) > 128 { + return false + } + for _, c := range id { + if !unicode.IsLetter(c) && !unicode.IsDigit(c) && c != '-' && c != '_' { + return false + } + } + return true +} +func (j *usageJournal) append(payload []byte) error { + j.mu.Lock() + defer j.mu.Unlock() + if j.dir == "" { + return nil + } + var event struct { + EventID string `json:"event_id"` + } + if err := json.Unmarshal(payload, &event); err != nil { + return err + } + // The legacy wire still carries the key for old consumers. Durable storage + // only needs a fingerprint for attribution, never the credential itself. + var fields map[string]json.RawMessage + if errDecode := json.Unmarshal(payload, &fields); errDecode != nil { + return errDecode + } + var key string + if errKey := json.Unmarshal(fields["api_key"], &key); errKey == nil && key != "" { + sum := sha256.Sum256([]byte(key)) + fingerprint := hex.EncodeToString(sum[:]) + fields["api_key_hash"], _ = json.Marshal(fingerprint) + fields["api_key_display"], _ = json.Marshal("sha256:" + fingerprint[:12]) + } + var source, authType string + _ = json.Unmarshal(fields["source"], &source) + _ = json.Unmarshal(fields["auth_type"], &authType) + if source != "" && (source == key || authType == "api_key" || authType == "apikey") { + sum := sha256.Sum256([]byte(source)) + fields["source"], _ = json.Marshal("sha256:" + hex.EncodeToString(sum[:])) + } + delete(fields, "api_key") + delete(fields, "response_headers") + var errEncode error + payload, errEncode = json.Marshal(fields) + if errEncode != nil { + return errEncode + } + if !validEventID(event.EventID) { + return errors.New("usage event requires a valid event_id") + } + if err := os.MkdirAll(j.dir, 0700); err != nil { + return err + } + target := filepath.Join(j.dir, event.EventID+".json") + if _, err := os.Stat(target); err == nil { + return nil + } + file, err := os.CreateTemp(j.dir, ".pending-") + if err != nil { + return err + } + name := file.Name() + defer func() { + if errRemove := os.Remove(name); errRemove != nil && !os.IsNotExist(errRemove) { + log.WithError(errRemove).Warn("remove pending usage event") + } + }() + _, errWrite := file.Write(payload) + if errWrite == nil { + errWrite = file.Sync() + } + errClose := file.Close() + if errWrite != nil { + return errWrite + } + if errClose != nil { + return errClose + } + if errRename := os.Rename(name, target); errRename != nil { + return errRename + } + return syncJournalDirectory(j.dir) +} +func (j *usageJournal) read(count int) ([][]byte, error) { + j.mu.Lock() + defer j.mu.Unlock() + if j.dir == "" { + return nil, errors.New("durable usage journal is not configured") + } + if j.lastErr != nil { + return nil, j.lastErr + } + entries, err := os.ReadDir(j.dir) + if os.IsNotExist(err) { + return [][]byte{}, nil + } + if err != nil { + return nil, err + } + sort.Slice(entries, func(i, k int) bool { return entries[i].Name() < entries[k].Name() }) + result := make([][]byte, 0) + for _, entry := range entries { + if entry.IsDir() || !strings.HasSuffix(entry.Name(), ".json") { + continue + } + payload, errRead := os.ReadFile(filepath.Join(j.dir, entry.Name())) + if errRead != nil { + return nil, errRead + } + result = append(result, payload) + if len(result) >= count { + break + } + } + return result, nil +} +func (j *usageJournal) ack(ids []string) error { + j.mu.Lock() + defer j.mu.Unlock() + if j.dir == "" { + return errors.New("durable usage journal is not configured") + } + for _, id := range ids { + if !validEventID(id) { + return errors.New("invalid usage event_id") + } + } + for _, id := range ids { + if err := os.Remove(filepath.Join(j.dir, id+".json")); err != nil && !os.IsNotExist(err) { + return err + } + } + if len(ids) == 0 { + return nil + } + return syncJournalDirectory(j.dir) +} +func syncJournalDirectory(dir string) error { + // Windows does not support fsync on directory handles; file contents were + // already flushed before the atomic rename. + if runtime.GOOS == "windows" { + return nil + } + file, errOpen := os.Open(dir) + if errOpen != nil { + return errOpen + } + errSync := file.Sync() + errClose := file.Close() + if errSync != nil { + return errSync + } + return errClose +} +func ReadUsageJournal(count int) ([][]byte, error) { return journal.read(count) } +func AckUsageJournal(ids []string) error { return journal.ack(ids) } +func journalUsage(payload []byte) { + if err := journal.append(payload); err != nil { + journal.mu.Lock() + journal.lastErr = err + journal.mu.Unlock() + log.WithError(err).Error("durable usage journal write failed; accounting requires attention") + } +} diff --git a/internal/redisqueue/journal_test.go b/internal/redisqueue/journal_test.go new file mode 100644 index 00000000000..f78543da18d --- /dev/null +++ b/internal/redisqueue/journal_test.go @@ -0,0 +1,78 @@ +package redisqueue + +import ( + "encoding/json" + "testing" +) + +func TestUsageJournalReplayRequiresAcknowledgement(t *testing.T) { + j := &usageJournal{dir: t.TempDir()} + if err := j.append([]byte(`{"event_id":"a","request_id":"same"}`)); err != nil { + t.Fatal(err) + } + if err := j.append([]byte(`{"event_id":"b","request_id":"same"}`)); err != nil { + t.Fatal(err) + } + for i := 0; i < 2; i++ { + items, err := j.read(100) + if err != nil || len(items) != 2 { + t.Fatalf("read %d: %d %v", i, len(items), err) + } + } + restarted := &usageJournal{dir: j.dir} + items, err := restarted.read(100) + if err != nil || len(items) != 2 { + t.Fatalf("restart: %d %v", len(items), err) + } + var event map[string]any + json.Unmarshal(items[0], &event) + if err := restarted.ack([]string{event["event_id"].(string)}); err != nil { + t.Fatal(err) + } + items, err = restarted.read(100) + if err != nil || len(items) != 1 { + t.Fatalf("ack: %d %v", len(items), err) + } + if err := restarted.ack([]string{"../../other"}); err == nil { + t.Fatal("unsafe id accepted") + } +} + +func TestUsageJournalDoesNotPersistPlaintextAPIKeys(t *testing.T) { + j := &usageJournal{dir: t.TempDir()} + if err := j.append([]byte(`{"event_id":"secret-test","request_id":"r","api_key":"not-a-real-key"}`)); err != nil { + t.Fatal(err) + } + items, err := j.read(10) + if err != nil { + t.Fatal(err) + } + var value map[string]any + if err := json.Unmarshal(items[0], &value); err != nil { + t.Fatal(err) + } + if _, exists := value["api_key"]; exists { + t.Fatal("durable journal contains API key") + } + if value["api_key_hash"] == nil { + t.Fatal("key attribution was lost") + } +} + +func TestUsageJournalFingerprintsCredentialSource(t *testing.T) { + j := &usageJournal{dir: t.TempDir()} + if err := j.append([]byte(`{"event_id":"source-test","source":"upstream-secret","auth_type":"api_key"}`)); err != nil { + t.Fatal(err) + } + items, err := j.read(10) + if err != nil { + t.Fatal(err) + } + var value map[string]any + if err := json.Unmarshal(items[0], &value); err != nil { + t.Fatal(err) + } + if value["source"] == "upstream-secret" { + t.Fatal("upstream credential persisted in source") + } +} diff --git a/internal/redisqueue/plugin.go b/internal/redisqueue/plugin.go index 0b3bf108b8a..7c886663ea7 100644 --- a/internal/redisqueue/plugin.go +++ b/internal/redisqueue/plugin.go @@ -7,6 +7,7 @@ import ( "strings" "time" + "github.com/google/uuid" internallogging "github.com/router-for-me/CLIProxyAPI/v7/internal/logging" coresession "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/session" coreusage "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" @@ -18,6 +19,8 @@ func init() { type usageQueuePlugin struct{} +func (*usageQueuePlugin) Synchronous() bool { return true } + func (p *usageQueuePlugin) HandleUsage(ctx context.Context, record coreusage.Record) { if p == nil { return @@ -121,7 +124,42 @@ func (p *usageQueuePlugin) HandleUsage(ctx context.Context, record coreusage.Rec ResponseHeaders: record.ResponseHeaders, } + eventID := record.EventID + if eventID == "" { + eventID = uuid.NewString() + } + if requestID == "" { + requestID = eventID + } + endpoint := record.Endpoint + if endpoint == "" { + endpoint = resolveEndpoint(ctx) + } + kind := record.Kind + if kind == "" { + kind = "attempt" + } + if kind == "attempt" && !coreusage.GenerateEnabled(record.Generate) { + kind = "prewarm" + } + transport := record.Transport + if transport == "" { + if stream { + transport = "sse" + } else { + transport = "http" + } + } + rawUsage := json.RawMessage(usageDetail.RawUsage) + if len(rawUsage) > 0 && !json.Valid(rawUsage) { + rawUsage = nil + } payload, err := json.Marshal(queuedUsageDetail{ + BillingID: usageDetail.BillingID, CostScope: usageDetail.CostScope, GenerationID: record.GenerationID, EventID: eventID, AttemptID: record.AttemptID, Kind: kind, Transport: transport, BaseURL: record.BaseURL, + UsageObserved: usageDetail.UsageObserved, RawUsage: rawUsage, + CacheCreation5mTokens: usageDetail.CacheCreation5mTokens, CacheCreation1hTokens: usageDetail.CacheCreation1hTokens, + CostUSD: usageDetail.CostUSD, + requestDetail: detail, AccountingVersion: coreusage.TokenAccountingSchemaVersion, TokenBreakdown: usageDetail.TokenBreakdown, @@ -129,7 +167,7 @@ func (p *usageQueuePlugin) HandleUsage(ctx context.Context, record coreusage.Rec ExecutorType: executorType, Model: modelName, Alias: aliasName, - Endpoint: resolveEndpoint(ctx), + Endpoint: endpoint, AuthType: authType, APIKey: apiKey, RequestID: requestID, @@ -142,10 +180,25 @@ func (p *usageQueuePlugin) HandleUsage(ctx context.Context, record coreusage.Rec if err != nil { return } + journalUsage(payload) Enqueue(payload) } type queuedUsageDetail struct { + Transport string `json:"transport"` + BillingID string `json:"billing_id,omitempty"` + CostScope string `json:"cost_scope,omitempty"` + GenerationID string `json:"generation_id,omitempty"` + EventID string `json:"event_id"` + AttemptID string `json:"attempt_id,omitempty"` + Kind string `json:"kind,omitempty"` + BaseURL string `json:"base_url,omitempty"` + UsageObserved bool `json:"usage_observed"` + RawUsage json.RawMessage `json:"raw_usage,omitempty"` + CacheCreation5mTokens int64 `json:"cache_creation_5m_tokens"` + CacheCreation1hTokens int64 `json:"cache_creation_1h_tokens"` + CostUSD *string `json:"cost_usd,omitempty"` + requestDetail AccountingVersion int `json:"accounting_version"` TokenBreakdown coreusage.TokenBreakdown `json:"token_breakdown"` diff --git a/internal/redisqueue/queue.go b/internal/redisqueue/queue.go index 85bd4a8fc33..8c5c6269e52 100644 --- a/internal/redisqueue/queue.go +++ b/internal/redisqueue/queue.go @@ -146,17 +146,19 @@ func (q *queue) publishToSubscribers(payload []byte) bool { return false } + delivered := false for id, subscriber := range q.subscribers { cloned := append([]byte(nil), payload...) select { case subscriber <- cloned: + delivered = true default: delete(q.subscribers, id) close(subscriber) } } - return true + return delivered } func (q *queue) subscribe(buffer int, initialPayload []byte) (<-chan []byte, func()) { diff --git a/internal/redisqueue/usage_integrity_test.go b/internal/redisqueue/usage_integrity_test.go new file mode 100644 index 00000000000..138222eb995 --- /dev/null +++ b/internal/redisqueue/usage_integrity_test.go @@ -0,0 +1,18 @@ +package redisqueue + +import "testing" + +func TestUsageIntegrityOverflowFallsBack(t *testing.T) { + SetEnabled(false) + SetEnabled(true) + defer SetEnabled(false) + _, cancel := SubscribeUsage() + defer cancel() + for i := 0; i < usageSubscriberBuffer-1; i++ { + Enqueue([]byte(`{"request_id":"a"}`)) + } + Enqueue([]byte(`{"request_id":"overflow"}`)) + if got := PopOldest(10); len(got) != 1 { + t.Fatalf("lost overflow event: got %d", len(got)) + } +} diff --git a/internal/runtime/executor/codex_websockets_execute.go b/internal/runtime/executor/codex_websockets_execute.go index e2c8026dbeb..d3d85d6866a 100644 --- a/internal/runtime/executor/codex_websockets_execute.go +++ b/internal/runtime/executor/codex_websockets_execute.go @@ -356,6 +356,7 @@ func (e *CodexWebsocketsExecutor) Execute(ctx context.Context, auth *cliproxyaut } else { reporter.EnsurePublished(ctx) } + publishCodexImageToolUsage(ctx, reporter, body, payload) var param any clientPayload := applyCodexIdentityExposeResponsePayload(payload, identityState) out := sdktranslator.TranslateNonStream(ctx, to, responseFormat, req.Model, originalPayload, clientBody, clientPayload, ¶m) diff --git a/internal/runtime/executor/codex_websockets_stream.go b/internal/runtime/executor/codex_websockets_stream.go index 0247c4f3f17..72ac8789f77 100644 --- a/internal/runtime/executor/codex_websockets_stream.go +++ b/internal/runtime/executor/codex_websockets_stream.go @@ -446,6 +446,7 @@ func (e *CodexWebsocketsExecutor) ExecuteStream(ctx context.Context, auth *clipr } else { reporter.EnsurePublished(ctx) } + publishCodexImageToolUsage(ctx, reporter, body, completedPayload) } var currentChunks [][]byte @@ -673,6 +674,7 @@ func (e *CodexWebsocketsExecutor) ExecuteStream(ctx context.Context, auth *clipr } else { reporter.EnsurePublished(ctx) } + publishCodexImageToolUsage(ctx, reporter, body, completedPayload) } clientPayload := applyCodexIdentityExposeResponsePayload(payload, identityState) diff --git a/internal/runtime/executor/codex_websockets_usage_integrity_test.go b/internal/runtime/executor/codex_websockets_usage_integrity_test.go new file mode 100644 index 00000000000..0954c483f8c --- /dev/null +++ b/internal/runtime/executor/codex_websockets_usage_integrity_test.go @@ -0,0 +1,96 @@ +package executor + +import ( + "context" + "net/http" + "net/http/httptest" + "sync" + "testing" + + "github.com/gorilla/websocket" + "github.com/router-for-me/CLIProxyAPI/v7/internal/config" + cliproxyauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + cliproxyexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" + sdktranslator "github.com/router-for-me/CLIProxyAPI/v7/sdk/translator" +) + +type websocketUsageCapture struct { + mu sync.Mutex + authID string + records []usage.Record +} + +func (*websocketUsageCapture) Synchronous() bool { return true } +func (p *websocketUsageCapture) HandleUsage(_ context.Context, r usage.Record) { + p.mu.Lock() + defer p.mu.Unlock() + if r.AuthID == p.authID { + p.records = append(p.records, r) + } +} +func TestCodexWebsocketImageUsageIntegrity(t *testing.T) { + for _, mode := range []string{"execute", "stream", "downstream-websocket"} { + t.Run(mode, func(t *testing.T) { + capture := &websocketUsageCapture{authID: t.Name()} + usage.RegisterNamedPlugin(t.Name(), capture) + defer usage.RegisterNamedPlugin(t.Name(), &websocketUsageCapture{}) + upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, errUpgrade := upgrader.Upgrade(w, r, nil) + if errUpgrade != nil { + t.Error(errUpgrade) + return + } + defer func() { + if errClose := conn.Close(); errClose != nil { + t.Log(errClose) + } + }() + if _, _, errRead := conn.ReadMessage(); errRead != nil { + t.Error(errRead) + return + } + if errWrite := conn.WriteMessage(websocket.TextMessage, []byte(`{"type":"response.completed","response":{"id":"response-one","output":[],"usage":{"input_tokens":100,"output_tokens":10,"total_tokens":110},"tool_usage":{"image_gen":{"input_tokens":40,"output_tokens":60,"total_tokens":100}}}}`)); errWrite != nil { + t.Error(errWrite) + } + })) + defer server.Close() + executor := NewCodexWebsocketsExecutor(&config.Config{}) + auth := &cliproxyauth.Auth{ID: t.Name(), Provider: "codex", Attributes: map[string]string{"api_key": "test", "base_url": server.URL}} + req := cliproxyexecutor.Request{Model: "gpt-5.4", Payload: []byte(`{"model":"gpt-5.4","input":"draw","tools":[{"type":"image_generation","model":"gpt-image-1"}]}`)} + opts := cliproxyexecutor.Options{SourceFormat: sdktranslator.FromString("openai-response")} + ctx := usage.WithNewGeneration(context.Background()) + if mode == "execute" { + if _, errExecute := executor.Execute(ctx, auth, req, opts); errExecute != nil { + t.Fatal(errExecute) + } + } else { + if mode == "downstream-websocket" { + ctx = cliproxyexecutor.WithDownstreamWebsocket(ctx) + } + result, errStream := executor.ExecuteStream(ctx, auth, req, opts) + if errStream != nil { + t.Fatal(errStream) + } + for chunk := range result.Chunks { + if chunk.Err != nil { + t.Fatal(chunk.Err) + } + } + } + capture.mu.Lock() + defer capture.mu.Unlock() + if len(capture.records) != 2 { + t.Fatalf("got %d usage events, want main and image", len(capture.records)) + } + a, b := capture.records[0], capture.records[1] + if a.Detail.TotalTokens != 110 || b.Detail.TotalTokens != 100 || b.Model != "gpt-image-1" || b.Kind != "tool" { + t.Fatalf("incorrect records: %+v / %+v", a, b) + } + if a.EventID == b.EventID || a.AttemptID != b.AttemptID || a.GenerationID != b.GenerationID { + t.Fatal("event identity must differ while attempt/generation stay linked") + } + }) + } +} diff --git a/internal/runtime/executor/helps/openai_compat_tool_results.go b/internal/runtime/executor/helps/openai_compat_tool_results.go index d62591e24ab..88eff435cbc 100644 --- a/internal/runtime/executor/helps/openai_compat_tool_results.go +++ b/internal/runtime/executor/helps/openai_compat_tool_results.go @@ -37,6 +37,14 @@ func NormalizeOpenAIToolResultsTextOnly(payload []byte) []byte { out := payload messageIndex := 0 messages.ForEach(func(_, message gjson.Result) bool { + // The Claude translator relays tool images in a following user message + // because Chat Completions tool messages only accept text content. + if message.Get("role").String() == "user" && message.Get("content.0.text").String() == "Images returned by the preceding tool call(s):" { + path := fmt.Sprintf("messages.%d.content", messageIndex) + if updated, errSet := sjson.SetBytes(out, path, flattenOpenAIToolResultContent(message.Get("content"))); errSet == nil { + out = updated + } + } if message.Get("role").String() == "tool" { content := message.Get("content") if content.Exists() && content.Type != gjson.String { diff --git a/internal/runtime/executor/helps/plugin_executor_usage.go b/internal/runtime/executor/helps/plugin_executor_usage.go index a0564ef03ac..74e90eceb95 100644 --- a/internal/runtime/executor/helps/plugin_executor_usage.go +++ b/internal/runtime/executor/helps/plugin_executor_usage.go @@ -2,6 +2,7 @@ package helps import ( "bytes" + "encoding/json" "strings" "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" @@ -143,6 +144,28 @@ func ObserveMergedStreamUsage(buffer *StreamUsageBuffer, update usage.Detail) { // MergeStreamUsageDetail merges existing stream usage with a newer update. func MergeStreamUsageDetail(existing, update usage.Detail) usage.Detail { merged := update + var rawExisting, rawUpdate map[string]json.RawMessage + if json.Unmarshal([]byte(existing.RawUsage), &rawExisting) == nil && json.Unmarshal([]byte(update.RawUsage), &rawUpdate) == nil { + for key, value := range rawUpdate { + rawExisting[key] = value + } + if encoded, errEncode := json.Marshal(rawExisting); errEncode == nil { + merged.RawUsage = string(encoded) + } + } else if merged.RawUsage == "" { + merged.RawUsage = existing.RawUsage + } + merged.UsageObserved = existing.UsageObserved || update.UsageObserved + if merged.CacheCreation5mTokens == 0 { + merged.CacheCreation5mTokens = existing.CacheCreation5mTokens + } + if merged.CacheCreation1hTokens == 0 { + merged.CacheCreation1hTokens = existing.CacheCreation1hTokens + } + if merged.CostUSD == nil { + merged.CostUSD = existing.CostUSD + } + if merged.InputTokens == 0 && existing.InputTokens > 0 { merged.InputTokens = existing.InputTokens } diff --git a/internal/runtime/executor/helps/usage_helpers.go b/internal/runtime/executor/helps/usage_helpers.go index c2a43db29d7..84cc646a72b 100644 --- a/internal/runtime/executor/helps/usage_helpers.go +++ b/internal/runtime/executor/helps/usage_helpers.go @@ -13,6 +13,7 @@ import ( "time" "github.com/gin-gonic/gin" + "github.com/google/uuid" "github.com/router-for-me/CLIProxyAPI/v7/internal/clienterror" internallogging "github.com/router-for-me/CLIProxyAPI/v7/internal/logging" "github.com/router-for-me/CLIProxyAPI/v7/internal/thinking" @@ -24,6 +25,12 @@ import ( ) type UsageReporter struct { + generationID string + transport string + endpoint string + kind string + attemptID string + provider string baseURL string executorType string @@ -63,6 +70,9 @@ func NewExecutorUsageReporter(ctx context.Context, executor usageExecutor, model } reporter := NewUsageReporter(ctx, provider, model, auth) reporter.executorType = ExecutorTypeName(executor) + if strings.Contains(strings.ToLower(reporter.executorType), "websocket") { + reporter.transport = "websocket" + } return reporter } @@ -94,6 +104,9 @@ func NewUsageReporter(ctx context.Context, provider, model string, auth *cliprox } } reporter := &UsageReporter{ + attemptID: uuid.NewString(), + generationID: usage.GenerationFromContext(ctx), + kind: "attempt", provider: provider, baseURL: baseURL, model: model, @@ -366,7 +379,10 @@ func (r *UsageReporter) buildAdditionalModelRecord(model string, detail usage.De if !hasNonZeroTokenUsage(detail) { return usage.Record{}, false } - return r.buildRecordForModel(model, detail, false, usage.Failure{}), true + record := r.buildRecordForModel(model, detail, false, usage.Failure{}) + record.Kind = "tool" + record.Alias = model + return record, true } func (r *UsageReporter) PublishFailure(ctx context.Context, errs ...error) { @@ -401,7 +417,7 @@ func normalizeUsageDetailTotal(detail usage.Detail, provider, executorType strin } func hasNonZeroTokenUsage(detail usage.Detail) bool { - return detail.InputTokens != 0 || + return detail.CostUSD != nil || detail.InputTokens != 0 || detail.OutputTokens != 0 || detail.ReasoningTokens != 0 || detail.CachedTokens != 0 || @@ -445,6 +461,12 @@ func (r *UsageReporter) buildRecordForModel(model string, detail usage.Detail, f return usage.Record{Model: model, Detail: detail, Failed: failed, Fail: fail, Generate: usage.GenerateFlag(true)} } return usage.Record{ + EventID: uuid.NewString(), + GenerationID: r.generationID, + AttemptID: r.attemptID, + Kind: r.kind, + Endpoint: r.endpoint, + Transport: r.transport, Provider: r.provider, BaseURL: r.baseURL, ExecutorType: r.executorType, @@ -756,7 +778,7 @@ func ParseCodexUsage(data []byte) (usage.Detail, bool) { } detail := parseOpenAIStyleUsageNode(usageNode) detail.ResponseServiceTier = responseServiceTier - return detail, true + return withResponseBilling(detail, gjson.GetBytes(data, "response")), true } func ParseCodexImageToolUsage(data []byte) (usage.Detail, bool) { @@ -775,14 +797,14 @@ func ParseOpenAIUsage(data []byte) usage.Detail { } detail := parseOpenAIStyleUsageNode(usageNode) detail.ResponseServiceTier = responseServiceTier - return detail + return withResponseBilling(detail, gjson.ParseBytes(data)) } func hasOpenAIStyleUsageTokenFields(usageNode gjson.Result) bool { if !usageNode.Exists() || !usageNode.IsObject() { return false } - return usageNode.Get("total_tokens").Exists() || hasOpenAIStyleUsageBucketFields(usageNode) + return usageNode.Get("cost_in_usd_ticks").Exists() || usageNode.Get("total_tokens").Exists() || hasOpenAIStyleUsageBucketFields(usageNode) } func hasOpenAIStyleUsageBucketFields(usageNode gjson.Result) bool { @@ -818,6 +840,12 @@ func parseOpenAIStyleUsageNode(usageNode gjson.Result) usage.Detail { if !cached.Exists() { cached = usageNode.Get("input_tokens_details.cached_tokens") } + if !cached.Exists() { + cached = usageNode.Get("input_token_details.cached_tokens") + } + if !cached.Exists() { + cached = usageNode.Get("prompt_cache_hit_tokens") + } if cached.Exists() { detail.CachedTokens = cached.Int() detail.CacheReadTokens = cached.Int() @@ -875,7 +903,7 @@ func parseOpenAIStyleUsageNode(usageNode gjson.Result) usage.Detail { if detail.TotalTokens == 0 { detail.TotalTokens = detail.TokenBreakdown.TotalTokens } - return detail + return withUsageMeasurements(detail, usageNode) } func ParseOpenAIStreamUsage(line []byte) (usage.Detail, bool) { @@ -966,7 +994,7 @@ func parseClaudeUsageNode(usageNode gjson.Result) usage.Detail { detail.ReasoningTokens, detail.TotalTokens, ) - return detail + return withUsageMeasurements(detail, usageNode) } func parseGeminiFamilyUsageDetail(node gjson.Result) usage.Detail { @@ -983,7 +1011,7 @@ func parseGeminiFamilyUsageDetail(node gjson.Result) usage.Detail { } if !okInput { detail.TokenBreakdown = invalidUsageTokenBreakdown(detail.TotalTokens) - return detail + return withUsageMeasurements(detail, node) } if detail.TotalTokens == 0 { var okTotal bool @@ -991,7 +1019,7 @@ func parseGeminiFamilyUsageDetail(node gjson.Result) usage.Detail { if !okTotal { detail.TotalTokens = 0 detail.TokenBreakdown = invalidUsageTokenBreakdown(0) - return detail + return withUsageMeasurements(detail, node) } } detail.TokenBreakdown = usage.NewSeparateReasoningTokenBreakdown( @@ -1002,7 +1030,7 @@ func parseGeminiFamilyUsageDetail(node gjson.Result) usage.Detail { detail.ReasoningTokens, detail.TotalTokens, ) - return detail + return withUsageMeasurements(detail, node) } func parseInteractionsUsageDetail(node gjson.Result) usage.Detail { @@ -1023,7 +1051,7 @@ func parseInteractionsUsageDetail(node gjson.Result) usage.Detail { } if !okInput { detail.TokenBreakdown = invalidUsageTokenBreakdown(detail.TotalTokens) - return detail + return withUsageMeasurements(detail, node) } if !cacheRead.Exists() && detail.CachedTokens > 0 { detail.CacheReadTokens = detail.CachedTokens @@ -1034,7 +1062,7 @@ func parseInteractionsUsageDetail(node gjson.Result) usage.Detail { if !okTotal { detail.TotalTokens = 0 detail.TokenBreakdown = invalidUsageTokenBreakdown(0) - return detail + return withUsageMeasurements(detail, node) } } detail.TokenBreakdown = usage.NewSeparateReasoningTokenBreakdown( @@ -1045,7 +1073,7 @@ func parseInteractionsUsageDetail(node gjson.Result) usage.Detail { detail.ReasoningTokens, detail.TotalTokens, ) - return detail + return withUsageMeasurements(detail, node) } func hasUsageDetail(detail usage.Detail) bool { @@ -1108,7 +1136,7 @@ func ParseGeminiUsage(data []byte) usage.Detail { if !node.Exists() { return usage.Detail{} } - return parseGeminiFamilyUsageDetail(node) + return withResponseBilling(parseGeminiFamilyUsageDetail(node), usageNode) } func ParseGeminiStreamUsage(line []byte) (usage.Detail, bool) { @@ -1123,7 +1151,7 @@ func ParseGeminiStreamUsage(line []byte) (usage.Detail, bool) { if !node.Exists() { return usage.Detail{}, false } - detail := parseGeminiFamilyUsageDetail(node) + detail := withResponseBilling(parseGeminiFamilyUsageDetail(node), gjson.ParseBytes(payload)) if !hasNonZeroTokenUsage(detail) { return usage.Detail{}, false } @@ -1370,3 +1398,16 @@ func jsonPayload(line []byte) []byte { } return trimmed } + +func (r *UsageReporter) SetOperation(kind, endpoint string) { + if r != nil { + r.kind = kind + r.endpoint = endpoint + } +} + +func (r *UsageReporter) SetTransport(transport string) { + if r != nil { + r.transport = transport + } +} diff --git a/internal/runtime/executor/helps/usage_integrity_test.go b/internal/runtime/executor/helps/usage_integrity_test.go new file mode 100644 index 00000000000..3f702fee85a --- /dev/null +++ b/internal/runtime/executor/helps/usage_integrity_test.go @@ -0,0 +1,32 @@ +package helps + +import ( + "testing" +) + +func TestUsageIntegrityPreservesClaudeCacheTTL(t *testing.T) { + d := ParseClaudeUsage([]byte(`{"usage":{"input_tokens":10,"output_tokens":20,"cache_creation_input_tokens":100,"cache_creation":{"ephemeral_5m_input_tokens":40,"ephemeral_1h_input_tokens":60}}}`)) + if !d.UsageObserved || string(d.RawUsage) == "" || d.CacheCreation5mTokens != 40 || d.CacheCreation1hTokens != 60 { + t.Fatalf("lost TTL/measurement provenance: %+v", d) + } +} +func TestUsageIntegrityPreservesProviderCostWithoutTokens(t *testing.T) { + d := ParseOpenAIUsage([]byte(`{"usage":{"cost_in_usd_ticks":123456789}}`)) + if !d.UsageObserved || d.CostUSD == nil || *d.CostUSD != "0.0123456789" { + t.Fatalf("lost provider cost: %+v", d) + } +} +func TestUsageIntegrityMissingAndMeasuredZeroDiffer(t *testing.T) { + if ParseOpenAIUsage([]byte(`{}`)).UsageObserved { + t.Fatal("missing usage was marked measured") + } + if !ParseOpenAIUsage([]byte(`{"usage":{"prompt_tokens":0,"completion_tokens":0}}`)).UsageObserved { + t.Fatal("measured zero lost") + } +} +func TestUsageIntegrityDeepSeekLegacyCache(t *testing.T) { + d := ParseOpenAIUsage([]byte(`{"usage":{"prompt_tokens":100,"completion_tokens":10,"prompt_cache_hit_tokens":80}}`)) + if d.CacheReadTokens != 80 { + t.Fatalf("legacy cache lost: %+v", d) + } +} diff --git a/internal/runtime/executor/helps/usage_measurements.go b/internal/runtime/executor/helps/usage_measurements.go new file mode 100644 index 00000000000..e90e09c3be5 --- /dev/null +++ b/internal/runtime/executor/helps/usage_measurements.go @@ -0,0 +1,65 @@ +package helps + +import ( + "encoding/json" + "math/big" + + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" + "github.com/tidwall/gjson" +) + +func withUsageMeasurements(detail usage.Detail, node gjson.Result) usage.Detail { + detail.UsageObserved = node.Exists() && node.IsObject() + if detail.UsageObserved { + detail.RawUsage = node.Raw + } + detail.CacheCreation5mTokens = node.Get("cache_creation.ephemeral_5m_input_tokens").Int() + detail.CacheCreation1hTokens = node.Get("cache_creation.ephemeral_1h_input_tokens").Int() + ticks := node.Get("cost_in_usd_ticks") + if ticks.Exists() { + // Decimal arithmetic preserves the provider's 1e-10 USD billing unit. + if amount, ok := new(big.Int).SetString(ticks.String(), 10); ok && amount.Sign() >= 0 { + value := new(big.Rat).SetFrac(amount, big.NewInt(10000000000)).FloatString(10) + detail.CostUSD = &value + } + } + return detail +} + +// Keep billing dimensions outside the token object. These must not silently +// disappear when a provider adds hosted tools with independent charges. +func withResponseBilling(detail usage.Detail, root gjson.Result) usage.Detail { + var raw map[string]any + if errDecode := json.Unmarshal([]byte(detail.RawUsage), &raw); errDecode != nil { + raw = map[string]any{} + } + if extra := root.Get("tool_usage"); extra.Exists() { + raw["tool_usage"] = extra.Value() + } + for _, candidate := range root.Get("candidates").Array() { + if candidate.Get("groundingMetadata").Exists() { + raw["unpriced_server_tools"] = true + } + } + webSearch, fileSearch := 0, 0 + for _, item := range root.Get("output").Array() { + switch item.Get("type").String() { + case "web_search_call": + webSearch++ + case "file_search_call": + fileSearch++ + case "code_interpreter_call", "shell_call": + raw["unpriced_server_tools"] = true + } + } + if webSearch > 0 { + raw["web_search_calls"] = webSearch + } + if fileSearch > 0 { + raw["file_search_calls"] = fileSearch + } + if encoded, errEncode := json.Marshal(raw); errEncode == nil && len(raw) > 0 { + detail.RawUsage = string(encoded) + } + return detail +} diff --git a/internal/runtime/executor/openai_compat_executor_tool_results_test.go b/internal/runtime/executor/openai_compat_executor_tool_results_test.go index 7ab0e21232f..1a727bec9ed 100644 --- a/internal/runtime/executor/openai_compat_executor_tool_results_test.go +++ b/internal/runtime/executor/openai_compat_executor_tool_results_test.go @@ -5,6 +5,7 @@ import ( "io" "net/http" "net/http/httptest" + "strings" "testing" "github.com/router-for-me/CLIProxyAPI/v7/internal/config" @@ -84,17 +85,18 @@ func TestOpenAICompatExecutorToolResultContentByInputModalities(t *testing.T) { } toolContent := gjson.GetBytes(gotBody, "messages.1.content") + if toolContent.String() != "image inspected" { + t.Fatalf("tool text lost: %s", gotBody) + } + relayContent := gjson.GetBytes(gotBody, "messages.2.content") if tt.wantString { - if toolContent.Type != gjson.String { - t.Fatalf("tool content type = %s, want string; body=%s", toolContent.Type, string(gotBody)) - } - want := "image inspected\n\n[image omitted: unsupported by upstream]" - if toolContent.String() != want { - t.Fatalf("tool content = %q, want %q", toolContent.String(), want) + if relayContent.Type != gjson.String || !strings.Contains(relayContent.String(), "[image omitted: unsupported by upstream]") { + t.Fatalf("text-only model received image relay: %s", gotBody) } - } else if !toolContent.IsArray() { - t.Fatalf("tool content type = %s, want array; body=%s", toolContent.Type, string(gotBody)) + } else if !relayContent.IsArray() || relayContent.Get("1.image_url.url").String() != "data:image/png;base64,AA==" { + t.Fatalf("multimodal tool image lost: %s", gotBody) } + }) } } diff --git a/internal/runtime/executor/xai_executor_media.go b/internal/runtime/executor/xai_executor_media.go index e79b735d914..a89562264d3 100644 --- a/internal/runtime/executor/xai_executor_media.go +++ b/internal/runtime/executor/xai_executor_media.go @@ -66,7 +66,12 @@ func (e *XAIExecutor) executeImages(ctx context.Context, auth *cliproxyauth.Auth return resp, err } - reporter.EnsurePublished(ctx) + detail := helps.ParseOpenAIUsage(data) + if detail.UsageObserved { + reporter.Publish(ctx, detail) + } else { + reporter.EnsurePublished(ctx) + } return cliproxyexecutor.Response{Payload: data, Headers: httpResp.Header.Clone()}, nil } @@ -140,6 +145,19 @@ func (e *XAIExecutor) executeVideos(ctx context.Context, auth *cliproxyauth.Auth return resp, xaiStatusErr(httpResp.StatusCode, data) } - reporter.EnsurePublished(ctx) + detail := helps.ParseOpenAIUsage(data) + billingID := strings.TrimSpace(gjson.GetBytes(data, "request_id").String()) + if billingID == "" { + billingID = strings.TrimSpace(gjson.GetBytes(payload, "request_id").String()) + } + if billingID != "" { + detail.BillingID = "xai-video/" + billingID + detail.CostScope = "operation" + } + if detail.UsageObserved { + reporter.Publish(ctx, detail) + } else { + reporter.EnsurePublished(ctx) + } return cliproxyexecutor.Response{Payload: data, Headers: httpResp.Header.Clone()}, nil } diff --git a/internal/runtime/executor/xai_executor_test.go b/internal/runtime/executor/xai_executor_test.go index 711dcee4b58..9199184a030 100644 --- a/internal/runtime/executor/xai_executor_test.go +++ b/internal/runtime/executor/xai_executor_test.go @@ -3289,8 +3289,8 @@ func TestXAIExecutorExecuteImagesUsesImagesEndpointAndPublishesUsage(t *testing. if record.Failed { t.Fatalf("failed = true, want false; failure=%+v", record.Fail) } - if record.Detail != (usage.Detail{}) { - t.Fatalf("detail = %+v, want zero token usage", record.Detail) + if record.Detail.TotalTokens != 0 || record.Detail.CostUSD == nil || *record.Detail.CostUSD != "0.0000250000" { + t.Fatalf("expected provider media charge with zero text tokens, got %+v", record.Detail) } if record.TTFT <= 0 { t.Fatalf("ttft = %v, want positive duration", record.TTFT) diff --git a/sdk/api/handlers/openai/openai_responses_websocket.go b/sdk/api/handlers/openai/openai_responses_websocket.go index a265db068d4..22a5f804ceb 100644 --- a/sdk/api/handlers/openai/openai_responses_websocket.go +++ b/sdk/api/handlers/openai/openai_responses_websocket.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" "net" "net/http" "strings" @@ -569,6 +570,8 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) { lastResponseID = "" lastResponsePendingToolCallIDs = nil prewarmID, errWrite := writeResponsesWebsocketSyntheticPrewarm(c, writer, requestJSON, wsTimelineLog, passthroughSessionID) + prewarmCtx := usage.WithNewGeneration(context.WithValue(c.Request.Context(), "gin", c)) + usage.PublishRecord(prewarmCtx, usage.Record{Provider: "local", Model: gjson.GetBytes(requestJSON, "model").String(), Kind: "prewarm", Transport: "websocket", Generate: usage.GenerateFlag(false), RequestedAt: time.Now(), Failed: errWrite != nil, Detail: usage.Detail{UsageObserved: true, TokenBreakdown: usage.NewSubsetTokenBreakdown(0, 0, 0, 0, 0, 0)}}) if errWrite != nil { wsTerminateErr = errWrite return @@ -616,6 +619,7 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) { if pinnedAuthID != "" && !routeOverridesModelResolution { cliCtx = handlers.WithPinnedAuthID(cliCtx, pinnedAuthID) } + cliCtx = usage.WithNewGeneration(cliCtx) dataChan, _, errChan := h.ExecuteStreamWithAuthManager(cliCtx, h.HandlerType(), modelName, requestJSON, "") if !selectedAuthObserved { // Plugin/alternate routes bypass auth selection. Keep canonical HTTP-mode diff --git a/sdk/cliproxy/usage/manager.go b/sdk/cliproxy/usage/manager.go index 5fd800999e1..5aa5f30f703 100644 --- a/sdk/cliproxy/usage/manager.go +++ b/sdk/cliproxy/usage/manager.go @@ -20,6 +20,13 @@ const AutoServiceTier = "auto" // Record contains the usage statistics captured for a single provider request. type Record struct { + Transport string + GenerationID string + EventID string + AttemptID string + Kind string + Endpoint string + Provider string // BaseURL stores the configured upstream base URL when available. BaseURL string @@ -69,6 +76,16 @@ type Failure struct { // Detail holds the token usage breakdown. type Detail struct { + BillingID string + CostScope string + // UsageObserved distinguishes an explicitly reported zero from absent usage. + UsageObserved bool + RawUsage string + CacheCreation5mTokens int64 + CacheCreation1hTokens int64 + // CostUSD is an exact decimal reported by the provider, not a token estimate. + CostUSD *string + InputTokens int64 OutputTokens int64 ReasoningTokens int64 @@ -343,6 +360,26 @@ func (m *Manager) Publish(ctx context.Context, record Record) { if m == nil { return } + m.mu.Lock() + closed := m.closed + m.mu.Unlock() + if closed { + return + } + markPublished(ctx) + if record.GenerationID == "" { + record.GenerationID = GenerationFromContext(ctx) + } + // Durable sinks run before returning to the caller so the in-memory dispatch + // queue is not another loss window. Other plugins retain asynchronous dispatch. + m.pluginsMu.RLock() + synchronous := append([]Plugin(nil), m.plugins...) + m.pluginsMu.RUnlock() + for _, plugin := range synchronous { + if p, ok := plugin.(interface{ Synchronous() bool }); ok && p.Synchronous() { + safeInvoke(plugin, ctx, record) + } + } // ensure worker is running even if Start was not called explicitly m.Start(context.Background()) m.mu.Lock() @@ -384,6 +421,9 @@ func (m *Manager) dispatch(item queueItem) { if plugin == nil { continue } + if p, ok := plugin.(interface{ Synchronous() bool }); ok && p.Synchronous() { + continue + } safeInvoke(plugin, item.ctx, item.record) } } diff --git a/sdk/cliproxy/usage/scope.go b/sdk/cliproxy/usage/scope.go new file mode 100644 index 00000000000..c40b840232f --- /dev/null +++ b/sdk/cliproxy/usage/scope.go @@ -0,0 +1,35 @@ +package usage + +import ( + "context" + "sync/atomic" + + "github.com/google/uuid" +) + +type scopeKey struct{} +type generationKey struct{} +type AccountingScope struct{ published atomic.Bool } + +func WithAccountingScope(ctx context.Context) (context.Context, *AccountingScope) { + s := &AccountingScope{} + return context.WithValue(ctx, scopeKey{}, s), s +} +func (s *AccountingScope) Published() bool { return s != nil && s.published.Load() } +func WithNewGeneration(ctx context.Context) context.Context { + return context.WithValue(ctx, generationKey{}, uuid.NewString()) +} +func GenerationFromContext(ctx context.Context) string { + if ctx == nil { + return "" + } + id, _ := ctx.Value(generationKey{}).(string) + return id +} +func markPublished(ctx context.Context) { + if ctx != nil { + if s, ok := ctx.Value(scopeKey{}).(*AccountingScope); ok { + s.published.Store(true) + } + } +} diff --git a/sdk/cliproxy/usage/scope_test.go b/sdk/cliproxy/usage/scope_test.go new file mode 100644 index 00000000000..8c22c676055 --- /dev/null +++ b/sdk/cliproxy/usage/scope_test.go @@ -0,0 +1,38 @@ +package usage + +import ( + "context" + "sync/atomic" + "testing" +) + +type synchronousTestSink struct{ calls atomic.Int32 } + +func (s *synchronousTestSink) Synchronous() bool { return true } +func (s *synchronousTestSink) HandleUsage(context.Context, Record) { s.calls.Add(1) } + +func TestSynchronousSinkCompletesBeforePublishReturns(t *testing.T) { + m := NewManager(1) + sink := &synchronousTestSink{} + m.Register(sink) + ctx, scope := WithAccountingScope(context.Background()) + m.Publish(WithNewGeneration(ctx), Record{}) + if sink.calls.Load() != 1 || !scope.Published() { + t.Fatal("durable sink must complete synchronously") + } + m.Stop() + if sink.calls.Load() != 1 { + t.Fatal("asynchronous dispatch repeated synchronous sink") + } +} +func TestGenerationIdentityChangesWithoutLosingCoverageScope(t *testing.T) { + ctx, scope := WithAccountingScope(context.Background()) + first, second := WithNewGeneration(ctx), WithNewGeneration(ctx) + if GenerationFromContext(first) == "" || GenerationFromContext(first) == GenerationFromContext(second) { + t.Fatal("generations must be distinct") + } + markPublished(second) + if !scope.Published() { + t.Fatal("generation context lost request coverage") + } +} From bd2e42c1c10c714a8b6d5fdc255fcf37b0a6a04a Mon Sep 17 00:00:00 2001 From: Ramapitecus Date: Sat, 12 Sep 2026 18:21:34 +0500 Subject: [PATCH 2/9] docs: describe usage identity, journal delivery and coverage limits --- docs/usage-accounting.md | 38 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 38 insertions(+) create mode 100644 docs/usage-accounting.md diff --git a/docs/usage-accounting.md b/docs/usage-accounting.md new file mode 100644 index 00000000000..332a550ea43 --- /dev/null +++ b/docs/usage-accounting.md @@ -0,0 +1,38 @@ +# Usage accounting and replay contract + +Usage events describe provider attempts, tool charges, and observed control calls. They are not a count of HTTP requests, WebSocket frames, or invoices. Enable usage statistics to collect them. + +## Event identity and measurements + +- `event_id` identifies one emitted accounting event. Consumers deduplicate only this ID, not `request_id`. +- `generation_id` groups retries for a logical generation; a downstream Responses WebSocket receives a new generation ID for every execution. `attempt_id` groups a provider attempt and its separately billed tools. +- `kind` distinguishes `attempt`, `tool`, `prewarm`, `management_call`, and `unmeasured`. Endpoint and transport are preserved. Missing executor coverage produces an unmeasured record for supported API route prefixes; model-list/configuration traffic is not a billable generation. +- `token_breakdown` remains the authoritative v2 normalized token contract. `raw_usage`, `usage_observed`, cache creation lifetimes, and provider cost preserve billing dimensions which cannot be reconstructed from token totals. +- `cost_usd` is an exact decimal string, currently sourced from xAI `cost_in_usd_ticks / 10000000000`. A video job's `billing_id` and `cost_scope=operation` identify a cumulative charge; repeated polling must not charge the same cumulative amount again. +- Claude cache creation retains both 5-minute and 1-hour tokens. Gemini thinking remains a separate legacy field and is included in the v2 billable output. OpenAI cached/reasoning tokens remain subsets rather than additional input/output. + +Codex image tools emit their own model, event ID, and raw modality measurements on HTTP and all three WebSocket execution paths. Live/realtime observation accepts upstream terminal usage events, ignores client/delta frames, and deduplicates response IDs. Capture is bounded without truncating forwarded media. Opaque WebRTC sessions without provider usage remain unmeasured; duration does not establish their charge. + +Alpha Search and model-bearing management POST probes publish available usage. Token-count endpoints do not turn predicted tokens into consumed tokens. Plugin executors must publish usage through the SDK to expose their own billing dimensions; arbitrary plugin routes are not automatically an exhaustive usage source. + +## Durable local consumer + +When a config file is supplied, the local journal is stored in `usage-journal/` beside it. The usage queue plugin persists events synchronously before asynchronous plugins run. Each file is atomically renamed after flushing; POSIX also flushes the directory. File permissions are private. API keys and credential-valued sources are fingerprinted; response headers are excluded from the durable copy. Legacy wire behavior remains available for old clients. + +The existing authenticated management API exposes: + +1. `GET /v0/management/usage-journal?count=500` (1–1000) reads without deletion. +2. Commit the events to a local inbox/database with a unique nonempty `event_id` constraint. +3. `POST /v0/management/usage-journal/ack` with `{"event_ids":["..."]}` removes committed events. Repeating an ACK is safe. + +The journal supports one acknowledging accounting consumer. It is not a multi-consumer broker. Unacknowledged files have no expiry: an old client that never ACKs, or a stopped collector, requires disk monitoring. No model or request body is stored, but usage/source metadata can still be sensitive. A storage write failure is logged and latches an error returned by journal reads; investigate the storage failure and reconcile provider usage before restarting. This cannot guarantee recovery of an event that could not be written, or usage never reported by the upstream. Avoid treating legacy destructive queue delivery as equally durable. + +## Validation + +Regression tests cover cache lifetime preservation, xAI cost without tokens, explicit versus absent usage, queue overflow fallback, journal replay/ACK/restart and credential fingerprints, bounded live observation, and image-tool accounting through Execute, streaming, and downstream WebSocket paths. The WebSocket regression was verified red with the publication hooks removed, then green with the fix. + +Full `go test ./...`, the server build, and race tests for usage, redisqueue, executor helpers, and live relay pass locally. A pre-existing text-only tool-result regression exposed by the full suite is also fixed: translated image relay messages now honor text-only tool mode while multimodal mode preserves them. + +## Coordinated desktop consumer + +EasyCLIProxyAPI needs the companion accounting update to consume the journal, preserve snapshots, distinguish unknown amounts, and apply modality/provider-specific tariffs. The desktop release must bundle a core release containing this change. Deploying only one side does not provide the full end-to-end guarantees. From 10cc49888878bf312e23356e472242705d7392f6 Mon Sep 17 00:00:00 2001 From: Ramapitecus Date: Sat, 12 Sep 2026 18:42:38 +0500 Subject: [PATCH 3/9] fix: retain streamed billing metadata and mark incomplete usage --- docs/usage-accounting.md | 6 ++ internal/client/codex/live/sideband.go | 17 ++-- internal/client/codex/live/usage.go | 77 +++++++++++++-- internal/client/codex/live/usage_test.go | 61 +++++++++++- internal/redisqueue/plugin.go | 3 +- .../redisqueue/usage_completeness_test.go | 18 ++++ .../executor/antigravity_executor_execute.go | 4 +- .../executor/antigravity_executor_stream.go | 4 +- internal/runtime/executor/gemini_executor.go | 4 +- .../executor/gemini_usage_billing_test.go | 48 +++++++++ .../executor/helps/plugin_executor_usage.go | 1 + .../helps/usage_billing_integrity_test.go | 75 ++++++++++++++ .../runtime/executor/helps/usage_helpers.go | 33 ++++++- .../executor/helps/usage_measurements.go | 99 +++++++++++++++++++ 14 files changed, 424 insertions(+), 26 deletions(-) create mode 100644 internal/redisqueue/usage_completeness_test.go create mode 100644 internal/runtime/executor/gemini_usage_billing_test.go create mode 100644 internal/runtime/executor/helps/usage_billing_integrity_test.go diff --git a/docs/usage-accounting.md b/docs/usage-accounting.md index 332a550ea43..1869910234f 100644 --- a/docs/usage-accounting.md +++ b/docs/usage-accounting.md @@ -36,3 +36,9 @@ Full `go test ./...`, the server build, and race tests for usage, redisqueue, ex ## Coordinated desktop consumer EasyCLIProxyAPI needs the companion accounting update to consume the journal, preserve snapshots, distinguish unknown amounts, and apply modality/provider-specific tariffs. The desktop release must bundle a core release containing this change. Deploying only one side does not provide the full end-to-end guarantees. + +## Independent review corrections + +The stream accounting buffer now retains billing metadata independently of final token counters, including metadata observed before Gemini/Antigravity SSE filtering. Repeated cumulative tool counters merge without adding the same count twice. A failed provider attempt publishes `usage_complete=false`, preventing a partial token snapshot from being frozen as a complete estimate by the desktop consumer. + +Live terminal usage is captured on the upstream read side, including a complete frame whose downstream write fails. Transcription events use the session's transcription model and item/content-index identity rather than the realtime model; absent transcription configuration remains an unknown model. This improves attribution without inventing usage after an upstream read failure. diff --git a/internal/client/codex/live/sideband.go b/internal/client/codex/live/sideband.go index da472f8a297..7e1808a545e 100644 --- a/internal/client/codex/live/sideband.go +++ b/internal/client/codex/live/sideband.go @@ -658,19 +658,16 @@ func copyWebsocket(destination, source *websocket.Conn, observers ...func([]byte } writer, errWriter := destination.NextWriter(messageType) if errWriter != nil { + if messageType == websocket.TextMessage && len(observers) > 0 { + _ = forwardUsageFrame(io.Discard, reader, observers...) + } return errWriter } - capture := &usageFrameCapture{} - var target io.Writer = writer - if len(observers) > 0 && messageType == websocket.TextMessage { - target = io.MultiWriter(writer, capture) - } - _, errCopy := io.Copy(target, reader) - if errCopy == nil && !capture.overflow && len(capture.data) > 0 { - for _, observe := range observers { - observe(capture.data) - } + var frameObservers []func([]byte) + if messageType == websocket.TextMessage { + frameObservers = observers } + errCopy := forwardUsageFrame(writer, reader, frameObservers...) errClose := writer.Close() if errCopy != nil { return errCopy diff --git a/internal/client/codex/live/usage.go b/internal/client/codex/live/usage.go index 3aed1baae21..d0309a2d9a1 100644 --- a/internal/client/codex/live/usage.go +++ b/internal/client/codex/live/usage.go @@ -3,6 +3,7 @@ package live import ( "context" "errors" + "io" "github.com/router-for-me/CLIProxyAPI/v7/internal/runtime/executor/helps" "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" @@ -13,11 +14,12 @@ import ( // Observe only upstream data events. Client frames and repeated response.done // notifications must never manufacture or double-count provider usage. type liveUsageObserver struct { - ctx context.Context - auth *auth.Auth - model string - seen map[string]bool - pending map[string]*helps.UsageReporter + ctx context.Context + auth *auth.Auth + model string + transcriptionModel string + seen map[string]bool + pending map[string]*helps.UsageReporter } func newLiveUsageObserver(ctx context.Context, selected *auth.Auth, model string) *liveUsageObserver { @@ -39,13 +41,29 @@ func (o *liveUsageObserver) observe(payload []byte) { if model := root.Get("session.model").String(); model != "" { o.model = model } + for _, path := range []string{"session.audio.input.transcription", "session.input_audio_transcription"} { + if node := root.Get(path); node.Exists() { + o.transcriptionModel = node.Get("model").String() + break + } + } kind := root.Get("type").String() id := root.Get("response.id").String() if id == "" { id = root.Get("item_id").String() } + transcription := kind == "conversation.item.input_audio_transcription.completed" + if transcription && id != "" { + id = "transcription:" + id + ":" + root.Get("content_index").String() + } else if id != "" { + id = "response:" + id + } if kind == "response.created" && id != "" && !o.seen[id] { - r := helps.NewUsageReporter(usage.WithNewGeneration(o.ctx), "codex", o.model, o.auth) + model := root.Get("response.model").String() + if model == "" { + model = o.model + } + r := helps.NewUsageReporter(usage.WithNewGeneration(o.ctx), "codex", model, o.auth) r.SetStream(true) r.SetTransport("websocket") o.pending[id] = r @@ -60,12 +78,21 @@ func (o *liveUsageObserver) observe(payload []byte) { if model == "" { model = o.model } + if transcription { + model = o.transcriptionModel + if model == "" { + model = "unknown" + } + } r := o.pending[id] if r == nil { r = helps.NewUsageReporter(usage.WithNewGeneration(o.ctx), "codex", model, o.auth) r.SetStream(true) r.SetTransport("websocket") } + if transcription { + r.SetOperation("tool", "") + } ctx := usage.WithNewGeneration(o.ctx) status := root.Get("response.status").String() if status == "failed" || status == "cancelled" || status == "incomplete" { @@ -98,3 +125,41 @@ func (b *usageFrameCapture) Write(p []byte) (int, error) { } return len(p), nil } + +// Read-side capture must survive a downstream write failure: the provider has +// already incurred usage even when the client no longer receives the response. +func forwardUsageFrame(writer io.Writer, reader io.Reader, observers ...func([]byte)) error { + if len(observers) == 0 { + _, errCopy := io.Copy(writer, reader) + return errCopy + } + capture := &usageFrameCapture{} + observedReader := io.TeeReader(reader, capture) + destination := &usageForwardWriter{Writer: writer} + _, errCopy := io.Copy(destination, observedReader) + complete := errCopy == nil + if destination.err != nil { + _, errDrain := io.Copy(io.Discard, observedReader) + complete = errDrain == nil + } + if complete && !capture.overflow && len(capture.data) > 0 { + for _, observe := range observers { + observe(capture.data) + } + } + return errCopy +} + +type usageForwardWriter struct { + io.Writer + err error +} + +func (w *usageForwardWriter) Write(p []byte) (int, error) { + n, errWrite := w.Writer.Write(p) + if errWrite == nil && n != len(p) { + errWrite = io.ErrShortWrite + } + w.err = errWrite + return n, errWrite +} diff --git a/internal/client/codex/live/usage_test.go b/internal/client/codex/live/usage_test.go index ed3e11e1cc2..0c71a7170c3 100644 --- a/internal/client/codex/live/usage_test.go +++ b/internal/client/codex/live/usage_test.go @@ -1,6 +1,13 @@ package live -import "testing" +import ( + "context" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" + "io" + "strings" + "testing" +) func TestLiveUsageTerminalDimensions(t *testing.T) { d, ok := liveUsageDetail([]byte(`{"type":"response.done","response":{"usage":{"input_tokens":100,"output_tokens":20,"input_token_details":{"audio_tokens":60,"text_tokens":40},"output_token_details":{"audio_tokens":20}}}}`)) @@ -18,3 +25,55 @@ func TestLiveUsageCaptureDoesNotLimitForwarding(t *testing.T) { t.Fatal("capture must be bounded without limiting stream") } } + +type disconnectedUsageWriter struct{} + +func (disconnectedUsageWriter) Write([]byte) (int, error) { return 0, io.ErrClosedPipe } +func TestUsageFrameSurvivesDownstreamWriteFailure(t *testing.T) { + payload := `{"type":"response.done","response":{"usage":{"input_tokens":100,"output_tokens":20}}}` + observed := "" + err := forwardUsageFrame(disconnectedUsageWriter{}, strings.NewReader(payload), func(b []byte) { observed = string(b) }) + if err != io.ErrClosedPipe { + t.Fatalf("lost forwarding failure: %v", err) + } + if observed != payload { + t.Fatal("provider usage lost after downstream disconnect") + } +} + +type liveReviewSink struct { + authID string + records chan usage.Record +} + +func (s *liveReviewSink) Synchronous() bool { return true } +func (s *liveReviewSink) HandleUsage(_ context.Context, r usage.Record) { + if r.AuthID == s.authID { + select { + case s.records <- r: + default: + } + } +} +func TestTranscriptionUsesItsOwnModelAndContentIdentity(t *testing.T) { + sink := &liveReviewSink{authID: t.Name(), records: make(chan usage.Record, 10)} + usage.RegisterPlugin(sink) + o := newLiveUsageObserver(context.Background(), &auth.Auth{ID: t.Name()}, "gpt-realtime") + o.observe([]byte(`{"type":"session.updated","session":{"model":"gpt-realtime","audio":{"input":{"transcription":{"model":"gpt-4o-transcribe"}}}}}`)) + for _, payload := range []string{ + `{"type":"conversation.item.input_audio_transcription.completed","item_id":"item1","content_index":0,"usage":{"input_tokens":10,"output_tokens":2,"total_tokens":12}}`, + `{"type":"conversation.item.input_audio_transcription.completed","item_id":"item1","content_index":1,"usage":{"input_tokens":20,"output_tokens":3,"total_tokens":23}}`, + } { + o.observe([]byte(payload)) + o.observe([]byte(payload)) + } + if len(sink.records) != 2 { + t.Fatalf("want two content events without duplicate terminals, got %d", len(sink.records)) + } + for i := 0; i < 2; i++ { + r := <-sink.records + if r.Model != "gpt-4o-transcribe" || r.Kind != "tool" { + t.Fatalf("transcription attributed to session: model=%s kind=%s", r.Model, r.Kind) + } + } +} diff --git a/internal/redisqueue/plugin.go b/internal/redisqueue/plugin.go index 7c886663ea7..a16304c4d0e 100644 --- a/internal/redisqueue/plugin.go +++ b/internal/redisqueue/plugin.go @@ -156,7 +156,7 @@ func (p *usageQueuePlugin) HandleUsage(ctx context.Context, record coreusage.Rec } payload, err := json.Marshal(queuedUsageDetail{ BillingID: usageDetail.BillingID, CostScope: usageDetail.CostScope, GenerationID: record.GenerationID, EventID: eventID, AttemptID: record.AttemptID, Kind: kind, Transport: transport, BaseURL: record.BaseURL, - UsageObserved: usageDetail.UsageObserved, RawUsage: rawUsage, + UsageObserved: usageDetail.UsageObserved, UsageComplete: !record.Failed, RawUsage: rawUsage, CacheCreation5mTokens: usageDetail.CacheCreation5mTokens, CacheCreation1hTokens: usageDetail.CacheCreation1hTokens, CostUSD: usageDetail.CostUSD, @@ -185,6 +185,7 @@ func (p *usageQueuePlugin) HandleUsage(ctx context.Context, record coreusage.Rec } type queuedUsageDetail struct { + UsageComplete bool `json:"usage_complete"` Transport string `json:"transport"` BillingID string `json:"billing_id,omitempty"` CostScope string `json:"cost_scope,omitempty"` diff --git a/internal/redisqueue/usage_completeness_test.go b/internal/redisqueue/usage_completeness_test.go new file mode 100644 index 00000000000..6ca4198de91 --- /dev/null +++ b/internal/redisqueue/usage_completeness_test.go @@ -0,0 +1,18 @@ +package redisqueue + +import ( + "context" + "testing" + + coreusage "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" +) + +func TestFailedPartialUsageIsNotPublishedAsComplete(t *testing.T) { + withEnabledQueue(t, func() { + (&usageQueuePlugin{}).HandleUsage(context.Background(), coreusage.Record{Provider: "claude", Model: "claude-opus-5", Failed: true, Detail: coreusage.Detail{InputTokens: 100, OutputTokens: 1, TotalTokens: 101, UsageObserved: true}}) + payload := popSinglePayload(t) + if string(payload["usage_complete"]) != "false" { + t.Fatalf("partial usage advertised as complete: %v", payload["usage_complete"]) + } + }) +} diff --git a/internal/runtime/executor/antigravity_executor_execute.go b/internal/runtime/executor/antigravity_executor_execute.go index a49a3f9c40f..5585ed56b21 100644 --- a/internal/runtime/executor/antigravity_executor_execute.go +++ b/internal/runtime/executor/antigravity_executor_execute.go @@ -406,8 +406,10 @@ func (e *AntigravityExecutor) executeClaudeNonStream(ctx context.Context, auth * }() scanner := bufio.NewScanner(resp.Body) scanner.Buffer(nil, streamScannerBuffer) + var billing helps.UsageBillingMetadata for scanner.Scan() { line := scanner.Bytes() + billing.ObservePayload(line) helps.AppendAPIResponseChunk(ctx, e.cfg, line) if replayAccumulator != nil { replayAccumulator.ObserveSSELine(line) @@ -423,7 +425,7 @@ func (e *AntigravityExecutor) executeClaudeNonStream(ctx context.Context, auth * } if detail, ok := helps.ParseAntigravityStreamUsage(payload); ok { - reporter.Publish(ctx, detail) + reporter.Publish(ctx, billing.Apply(detail)) } out <- cliproxyexecutor.StreamChunk{Payload: payload} diff --git a/internal/runtime/executor/antigravity_executor_stream.go b/internal/runtime/executor/antigravity_executor_stream.go index 67cbd18a5c8..d7af31e38f5 100644 --- a/internal/runtime/executor/antigravity_executor_stream.go +++ b/internal/runtime/executor/antigravity_executor_stream.go @@ -201,10 +201,12 @@ func (e *AntigravityExecutor) ExecuteStream(ctx context.Context, auth *cliproxya }() scanner := bufio.NewScanner(resp.Body) scanner.Buffer(nil, streamScannerBuffer) + var billing helps.UsageBillingMetadata claudeInputTokens := helps.NewClaudeInputTokenState(from, to, responseFormat, originalPayload) var param any for scanner.Scan() { line := scanner.Bytes() + billing.ObservePayload(line) helps.AppendAPIResponseChunk(ctx, e.cfg, line) if replayAccumulator != nil { replayAccumulator.ObserveSSELine(line) @@ -220,7 +222,7 @@ func (e *AntigravityExecutor) ExecuteStream(ctx context.Context, auth *cliproxya } if detail, ok := helps.ParseAntigravityStreamUsage(payload); ok { - reporter.Publish(ctx, detail) + reporter.Publish(ctx, billing.Apply(detail)) } payload = e.resolveWebSearchGroundingURLs(ctx, auth, from, originalPayload, translated, payload) diff --git a/internal/runtime/executor/gemini_executor.go b/internal/runtime/executor/gemini_executor.go index ac7557968a4..71dcb2e54eb 100644 --- a/internal/runtime/executor/gemini_executor.go +++ b/internal/runtime/executor/gemini_executor.go @@ -356,10 +356,12 @@ func (e *GeminiExecutor) ExecuteStream(ctx context.Context, auth *cliproxyauth.A }() scanner := bufio.NewScanner(httpResp.Body) scanner.Buffer(nil, streamScannerBuffer) + var billing helps.UsageBillingMetadata claudeInputTokens := helps.NewClaudeInputTokenState(from, to, responseFormat, originalPayload) var param any for scanner.Scan() { line := scanner.Bytes() + billing.ObservePayload(line) helps.AppendAPIResponseChunk(ctx, e.cfg, line) filtered := helps.FilterSSEUsageMetadata(line) payload := helps.JSONPayload(filtered) @@ -367,7 +369,7 @@ func (e *GeminiExecutor) ExecuteStream(ctx context.Context, auth *cliproxyauth.A continue } if detail, ok := helps.ParseGeminiStreamUsage(payload); ok { - reporter.Publish(ctx, detail) + reporter.Publish(ctx, billing.Apply(detail)) } lines := helps.TranslateStreamWithClaudeInputTokens(ctx, to, responseFormat, req.Model, opts.OriginalRequest, body, bytes.Clone(payload), ¶m, claudeInputTokens) for i := range lines { diff --git a/internal/runtime/executor/gemini_usage_billing_test.go b/internal/runtime/executor/gemini_usage_billing_test.go new file mode 100644 index 00000000000..7d6c59e3113 --- /dev/null +++ b/internal/runtime/executor/gemini_usage_billing_test.go @@ -0,0 +1,48 @@ +package executor + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "github.com/router-for-me/CLIProxyAPI/v7/internal/config" + cliproxyauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + cliproxyexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" + sdktranslator "github.com/router-for-me/CLIProxyAPI/v7/sdk/translator" + "github.com/tidwall/gjson" +) + +func TestGeminiStreamPreservesGroundingBeforeTerminalUsage(t *testing.T) { + capture := &websocketUsageCapture{authID: t.Name()} + usage.RegisterNamedPlugin(t.Name(), capture) + defer usage.RegisterNamedPlugin(t.Name(), &websocketUsageCapture{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + _, _ = fmt.Fprint(w, "data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"weather\"}]},\"groundingMetadata\":{\"webSearchQueries\":[\"weather\"]}}]}\n\n") + _, _ = fmt.Fprint(w, "data: {\"candidates\":[{\"finishReason\":\"STOP\"}],\"usageMetadata\":{\"promptTokenCount\":100,\"candidatesTokenCount\":10,\"totalTokenCount\":110}}\n\n") + })) + defer server.Close() + auth := &cliproxyauth.Auth{ID: t.Name(), Provider: "gemini", Attributes: map[string]string{"api_key": "test", "base_url": server.URL}} + req := cliproxyexecutor.Request{Model: "gemini-2.5-flash", Payload: []byte(`{"contents":[{"role":"user","parts":[{"text":"weather"}]}]}`)} + result, errStream := NewGeminiExecutor(&config.Config{}).ExecuteStream(context.Background(), auth, req, cliproxyexecutor.Options{SourceFormat: sdktranslator.FormatGemini}) + if errStream != nil { + t.Fatal(errStream) + } + for chunk := range result.Chunks { + if chunk.Err != nil { + t.Fatal(chunk.Err) + } + } + capture.mu.Lock() + defer capture.mu.Unlock() + if len(capture.records) != 1 { + t.Fatalf("got %d usage records, want one terminal event", len(capture.records)) + } + detail := capture.records[0].Detail + if detail.TotalTokens != 110 || !gjson.Get(detail.RawUsage, "unpriced_server_tools").Bool() { + t.Fatalf("native stream discarded early grounding: %+v", detail) + } +} diff --git a/internal/runtime/executor/helps/plugin_executor_usage.go b/internal/runtime/executor/helps/plugin_executor_usage.go index 74e90eceb95..9a96896e793 100644 --- a/internal/runtime/executor/helps/plugin_executor_usage.go +++ b/internal/runtime/executor/helps/plugin_executor_usage.go @@ -38,6 +38,7 @@ func ObservePluginExecutorStreamUsage(protocol string, payload []byte, buffer *S if buffer == nil || len(payload) == 0 { return } + IterateStreamLines(payload, buffer.ObserveBillingPayload) switch strings.ToLower(strings.TrimSpace(protocol)) { case "claude": IterateStreamLines(payload, func(line []byte) { diff --git a/internal/runtime/executor/helps/usage_billing_integrity_test.go b/internal/runtime/executor/helps/usage_billing_integrity_test.go new file mode 100644 index 00000000000..c36dea86ee5 --- /dev/null +++ b/internal/runtime/executor/helps/usage_billing_integrity_test.go @@ -0,0 +1,75 @@ +package helps + +import ( + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" + "github.com/tidwall/gjson" + "testing" +) + +func TestReviewAntigravityGroundingBilling(t *testing.T) { + detail := ParseAntigravityUsage([]byte(`{"response":{"candidates":[{"groundingMetadata":{"webSearchQueries":["weather"]}}],"usageMetadata":{"promptTokenCount":100,"candidatesTokenCount":10,"totalTokenCount":110}}}`)) + if !gjson.Get(detail.RawUsage, "unpriced_server_tools").Bool() { + t.Fatalf("grounding charge dropped: raw=%s", detail.RawUsage) + } +} + +func TestReviewGeminiStreamRetainsPriorGrounding(t *testing.T) { + var buffer StreamUsageBuffer + ObservePluginExecutorStreamUsage("gemini", []byte("data: {\"candidates\":[{\"groundingMetadata\":{\"webSearchQueries\":[\"weather\"]}}],\"usageMetadata\":{\"promptTokenCount\":100,\"candidatesTokenCount\":5,\"totalTokenCount\":105}}\n\n"), &buffer) + ObservePluginExecutorStreamUsage("gemini", []byte("data: {\"candidates\":[{\"finishReason\":\"STOP\"}],\"usageMetadata\":{\"promptTokenCount\":100,\"candidatesTokenCount\":10,\"totalTokenCount\":110}}\n\n"), &buffer) + detail, _ := buffer.Detail() + if !gjson.Get(detail.RawUsage, "unpriced_server_tools").Bool() { + t.Fatalf("earlier grounding charge dropped: raw=%s", detail.RawUsage) + } +} + +func TestReviewOpenAIStreamRetainsToolBilling(t *testing.T) { + detail, ok := ParseOpenAIStreamUsage([]byte(`data: {"usage":{"prompt_tokens":100,"completion_tokens":10,"total_tokens":110},"tool_usage":{"web_search":1}}`)) + if !ok || !gjson.Get(detail.RawUsage, "tool_usage.web_search").Exists() { + t.Fatalf("stream root tool billing dropped: raw=%s", detail.RawUsage) + } +} + +func TestBillingMetadataSurvivesNativeGeminiFiltering(t *testing.T) { + for _, wrapped := range []bool{false, true} { + var billing UsageBillingMetadata + early := `{"candidates":[{"groundingMetadata":{"webSearchQueries":["weather"]}}]}` + final := `{"candidates":[{"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":100,"candidatesTokenCount":10,"totalTokenCount":110}}` + if wrapped { + early = `{"response":` + early + `}` + final = `{"response":` + final + `}` + } + billing.ObservePayload([]byte("data: " + early)) + unmeasured := billing.Apply(usage.Detail{}) + if unmeasured.UsageObserved || unmeasured.TotalTokens != 0 { + t.Fatal("billing metadata manufactured token usage") + } + billing.ObservePayload([]byte("data: " + final)) + filtered := FilterSSEUsageMetadata([]byte("data: " + final)) + detail, ok := ParseGeminiStreamUsage(filtered) + if wrapped { + detail, ok = ParseAntigravityStreamUsage(filtered) + } + if !ok { + t.Fatal("missing terminal token usage") + } + detail = billing.Apply(detail) + if detail.TotalTokens != 110 || !gjson.Get(detail.RawUsage, "unpriced_server_tools").Bool() { + t.Fatalf("filtered native stream lost billing dimensions: %+v", detail) + } + } +} + +func TestOpenAIStreamBillingOnlyFramesRetainFinalCounters(t *testing.T) { + var buffer StreamUsageBuffer + buffer.ObserveOpenAIStream([]byte(`data: {"tool_usage":{"web_search":1}}`)) + if _, ok := buffer.Detail(); ok { + t.Fatal("billing-only frame manufactured observed token usage") + } + buffer.ObserveOpenAIStream([]byte(`data: {"usage":{"prompt_tokens":100,"completion_tokens":10,"total_tokens":110}}`)) + buffer.ObserveOpenAIStream([]byte(`data: {"tool_usage":{"file_search":1}}`)) + detail, ok := buffer.Detail() + if !ok || detail.TotalTokens != 110 || !gjson.Get(detail.RawUsage, "tool_usage.web_search").Exists() || !gjson.Get(detail.RawUsage, "tool_usage.file_search").Exists() { + t.Fatalf("metadata-only frames corrupted accounting: %+v", detail) + } +} diff --git a/internal/runtime/executor/helps/usage_helpers.go b/internal/runtime/executor/helps/usage_helpers.go index 84cc646a72b..8541510ed02 100644 --- a/internal/runtime/executor/helps/usage_helpers.go +++ b/internal/runtime/executor/helps/usage_helpers.go @@ -668,8 +668,9 @@ func resolveUsageAuthType(auth *cliproxyauth.Auth) string { // StreamUsageBuffer keeps the latest usage detail observed in a stream. type StreamUsageBuffer struct { - detail usage.Detail - ok bool + detail usage.Detail + ok bool + billing UsageBillingMetadata } var ( @@ -682,6 +683,7 @@ func (b *StreamUsageBuffer) Observe(detail usage.Detail, ok bool) { if b == nil || !ok { return } + detail = b.billing.Apply(detail) responseServiceTier := strings.TrimSpace(detail.ResponseServiceTier) if responseServiceTier == "" || hasNonZeroTokenUsage(detail) { preservedTier := b.detail.ResponseServiceTier @@ -695,12 +697,25 @@ func (b *StreamUsageBuffer) Observe(detail usage.Detail, ok bool) { b.ok = true } +// ObserveBillingPayload preserves billing metadata even when this frame has no +// token usage, or the translator later removes provisional usage metadata. +func (b *StreamUsageBuffer) ObserveBillingPayload(payload []byte) { + if b == nil { + return + } + b.billing.ObservePayload(payload) + if b.ok { + b.detail = b.billing.Apply(b.detail) + } +} + // ObserveOpenAIStream records response-tier state and the latest usage from an // OpenAI-style stream while avoiding JSON parsing for irrelevant chunks. func (b *StreamUsageBuffer) ObserveOpenAIStream(line []byte) { if b == nil { return } + b.ObserveBillingPayload(line) payload := jsonPayload(line) if len(payload) == 0 { return @@ -921,7 +936,7 @@ func ParseOpenAIStreamUsage(line []byte) (usage.Detail, bool) { } detail := parseOpenAIStyleUsageNode(usageNode) detail.ResponseServiceTier = responseServiceTier - return detail, true + return withResponseBilling(detail, gjson.ParseBytes(payload)), true } func ParseClaudeUsage(data []byte) usage.Detail { @@ -1203,7 +1218,11 @@ func ParseAntigravityUsage(data []byte) usage.Detail { if !node.Exists() { return usage.Detail{} } - return parseGeminiFamilyUsageDetail(node) + root := usageNode + if response := root.Get("response"); response.IsObject() { + root = response + } + return withResponseBilling(parseGeminiFamilyUsageDetail(node), root) } func ParseAntigravityStreamUsage(line []byte) (usage.Detail, bool) { @@ -1221,7 +1240,11 @@ func ParseAntigravityStreamUsage(line []byte) (usage.Detail, bool) { if !node.Exists() { return usage.Detail{}, false } - return parseGeminiFamilyUsageDetail(node), true + root := gjson.ParseBytes(payload) + if response := root.Get("response"); response.IsObject() { + root = response + } + return withResponseBilling(parseGeminiFamilyUsageDetail(node), root), true } var stopChunkWithoutUsage sync.Map diff --git a/internal/runtime/executor/helps/usage_measurements.go b/internal/runtime/executor/helps/usage_measurements.go index e90e09c3be5..4e47f0381ac 100644 --- a/internal/runtime/executor/helps/usage_measurements.go +++ b/internal/runtime/executor/helps/usage_measurements.go @@ -1,6 +1,7 @@ package helps import ( + "bytes" "encoding/json" "math/big" @@ -63,3 +64,101 @@ func withResponseBilling(detail usage.Detail, root gjson.Result) usage.Detail { } return detail } + +// UsageBillingMetadata retains billing dimensions which may arrive in a +// different frame from the final token counters. Observing metadata never +// turns an unmeasured frame into a measured token response. +type UsageBillingMetadata struct { + fields map[string]any +} + +func (b *UsageBillingMetadata) ObservePayload(payload []byte) { + // Most stream frames contain only text/audio deltas. Avoid parsing their + // potentially large bodies when no billing dimension can be present. + candidate := false + for _, marker := range []string{"\"tool_usage\"", "\"server_tool_use\"", "\"groundingMetadata\"", "\"web_search_call\"", "\"file_search_call\"", "\"code_interpreter_call\"", "\"shell_call\""} { + if bytes.Contains(payload, []byte(marker)) { + candidate = true + break + } + } + if !candidate { + return + } + payload = ExtractStreamJSONPayload(payload) + if !gjson.ValidBytes(payload) { + return + } + root := gjson.ParseBytes(payload) + for _, node := range []gjson.Result{root, root.Get("response")} { + if !node.IsObject() { + continue + } + detail := withResponseBilling(usage.Detail{}, node) + b.observeDetail(detail) + for _, path := range []string{"usage", "message.usage"} { + if raw := node.Get(path); raw.IsObject() { + b.observeDetail(usage.Detail{RawUsage: raw.Raw}) + } + } + } +} + +func (b *UsageBillingMetadata) observeDetail(detail usage.Detail) { + var raw map[string]any + if json.Unmarshal([]byte(detail.RawUsage), &raw) != nil { + return + } + if b.fields == nil { + b.fields = make(map[string]any) + } + for _, key := range []string{"unpriced_server_tools", "tool_usage", "server_tool_use", "web_search_calls", "file_search_calls"} { + if value, ok := raw[key]; ok { + b.fields[key] = mergeBillingMetadata(b.fields[key], value) + } + } +} + +// Provider stream counters are cumulative snapshots, not additive deltas. +// Keep known keys and the largest count so repeats cannot charge twice and +// later metadata-less frames cannot erase already observed tool usage. +func mergeBillingMetadata(previous, next any) any { + switch value := next.(type) { + case map[string]any: + merged, _ := previous.(map[string]any) + if merged == nil { + merged = make(map[string]any) + } + for key, child := range value { + merged[key] = mergeBillingMetadata(merged[key], child) + } + return merged + case float64: + if old, ok := previous.(float64); ok && old > value { + return old + } + case bool: + if old, ok := previous.(bool); ok && old { + return true + } + } + return next +} + +func (b *UsageBillingMetadata) Apply(detail usage.Detail) usage.Detail { + b.observeDetail(detail) + if len(b.fields) == 0 { + return detail + } + var raw map[string]any + if json.Unmarshal([]byte(detail.RawUsage), &raw) != nil || raw == nil { + raw = make(map[string]any) + } + for key, value := range b.fields { + raw[key] = value + } + if encoded, errEncode := json.Marshal(raw); errEncode == nil { + detail.RawUsage = string(encoded) + } + return detail +} From 771f30123ebf3985df1747413d45f2f093cd50ec Mon Sep 17 00:00:00 2001 From: Ramapitecus Date: Sat, 12 Sep 2026 18:52:03 +0500 Subject: [PATCH 4/9] fix: address review feedback on accounting context and delivery identity --- docs/usage-accounting.md | 6 ++ internal/api/handlers/management/api_tools.go | 49 ++++++++---- .../management/api_tools_usage_test.go | 66 ++++++++++++++++ internal/api/usage_coverage.go | 6 ++ internal/api/usage_coverage_test.go | 78 +++++++++++++++++++ internal/redisqueue/plugin.go | 2 +- .../redisqueue/usage_completeness_test.go | 13 ++++ sdk/api/handlers/handlers.go | 1 + sdk/cliproxy/usage/manager.go | 4 + sdk/cliproxy/usage/scope.go | 23 ++++++ sdk/cliproxy/usage/scope_test.go | 30 +++++++ 11 files changed, 264 insertions(+), 14 deletions(-) create mode 100644 internal/api/handlers/management/api_tools_usage_test.go create mode 100644 internal/api/usage_coverage_test.go diff --git a/docs/usage-accounting.md b/docs/usage-accounting.md index 1869910234f..b2704e5cfb6 100644 --- a/docs/usage-accounting.md +++ b/docs/usage-accounting.md @@ -42,3 +42,9 @@ EasyCLIProxyAPI needs the companion accounting update to consume the journal, pr The stream accounting buffer now retains billing metadata independently of final token counters, including metadata observed before Gemini/Antigravity SSE filtering. Repeated cumulative tool counters merge without adding the same count twice. A failed provider attempt publishes `usage_complete=false`, preventing a partial token snapshot from being frozen as a complete estimate by the desktop consumer. Live terminal usage is captured on the upstream read side, including a complete frame whose downstream write fails. Transcription events use the session's transcription model and item/content-index identity rather than the realtime model; absent transcription configuration remains an unknown model. This improves attribution without inventing usage after an upstream read failure. + +## PR review corrections + +The common HTTP handler context bridge now copies accounting scope and generation identity while preserving the execution context's existing cancellation behavior. Measured HTTP operations therefore suppress the unmeasured fallback correctly. The usage manager assigns a missing event ID once before either synchronous or asynchronous plugins run, preserving caller-supplied identities. + +Coverage skips unmatched routes and locally rejected authorization, so unauthenticated requests cannot create permanent journal entries through that fallback. `usage_complete` uses the same resolved failure state as the emitted `failed` field. Model-bearing management calls select the provider/endpoint protocol parser and merge SSE usage; partial response bodies are retained when reading fails. diff --git a/internal/api/handlers/management/api_tools.go b/internal/api/handlers/management/api_tools.go index 9572f42f2e0..ac9e25ba19c 100644 --- a/internal/api/handlers/management/api_tools.go +++ b/internal/api/handlers/management/api_tools.go @@ -4,8 +4,6 @@ import ( "context" "encoding/json" "fmt" - "github.com/router-for-me/CLIProxyAPI/v7/internal/runtime/executor/helps" - "github.com/tidwall/gjson" "io" "net/http" "net/url" @@ -14,9 +12,12 @@ import ( "github.com/gin-gonic/gin" "github.com/router-for-me/CLIProxyAPI/v7/internal/config" + "github.com/router-for-me/CLIProxyAPI/v7/internal/runtime/executor/helps" coreauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + coreusage "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" "github.com/router-for-me/CLIProxyAPI/v7/sdk/proxyutil" log "github.com/sirupsen/logrus" + "github.com/tidwall/gjson" ) const defaultAPICallTimeout = 60 * time.Second @@ -195,13 +196,13 @@ func (h *Handler) APICall(c *gin.Context) { var reporter *helps.UsageReporter usageCtx := context.WithValue(c.Request.Context(), "gin", c) model := gjson.Get(body.Data, "model").String() + provider := "openai-compatible" + if auth != nil { + provider = auth.Provider + } else if req.URL.Hostname() == "api.x.ai" { + provider = "xai" + } if method == http.MethodPost && model != "" { - provider := "openai-compatible" - if auth != nil { - provider = auth.Provider - } else if req.URL.Hostname() == "api.x.ai" { - provider = "xai" - } reporter = helps.NewUsageReporter(usageCtx, provider, model, auth) reporter.SetOperation("management_call", method+" "+req.URL.Path) defer reporter.EnsurePublished(usageCtx) @@ -220,17 +221,14 @@ func (h *Handler) APICall(c *gin.Context) { }() respBody, errReadAll := io.ReadAll(resp.Body) + detail := parseManagementResponseUsage(provider, req.URL.Path, resp.Header.Get("Content-Type"), respBody) if errReadAll != nil { - reporter.PublishFailure(usageCtx, errReadAll) + reporter.PublishFailureWithDetail(usageCtx, detail, errReadAll) c.JSON(http.StatusBadGateway, gin.H{"error": "failed to read response"}) return } if reporter != nil { - detail := helps.ParseOpenAIUsage(respBody) - if strings.HasSuffix(req.URL.Path, "/messages") { - detail = helps.ParseClaudeUsage(respBody) - } if resp.StatusCode >= 400 { reporter.PublishFailureWithDetail(usageCtx, detail, fmt.Errorf("upstream HTTP %d", resp.StatusCode)) } else { @@ -691,3 +689,28 @@ func buildProxyTransport(proxyStr string) *http.Transport { } return transport } + +// Management calls bypass the executor dispatcher, so select the same protocol +// parsers explicitly and merge streaming start/delta usage when appropriate. +func parseManagementResponseUsage(provider, path, contentType string, payload []byte) coreusage.Detail { + protocol := "openai" + switch { + case strings.EqualFold(provider, "antigravity"): + protocol = "antigravity" + case strings.HasSuffix(path, "/interactions"): + protocol = "interactions" + case strings.HasSuffix(path, "/messages") || strings.EqualFold(provider, "claude") || strings.EqualFold(provider, "anthropic"): + protocol = "claude" + case strings.Contains(path, ":generateContent") || strings.Contains(path, ":streamGenerateContent") || strings.EqualFold(provider, "gemini") || strings.EqualFold(provider, "vertex") || strings.EqualFold(provider, "aistudio"): + protocol = "gemini" + case strings.HasSuffix(path, "/responses") || strings.EqualFold(provider, "codex"): + protocol = "openai-response" + } + if strings.Contains(strings.ToLower(contentType), "text/event-stream") { + var buffer helps.StreamUsageBuffer + helps.ObservePluginExecutorStreamUsage(protocol, payload, &buffer) + detail, _ := buffer.Detail() + return detail + } + return helps.ParsePluginExecutorResponseUsage(protocol, payload) +} diff --git a/internal/api/handlers/management/api_tools_usage_test.go b/internal/api/handlers/management/api_tools_usage_test.go new file mode 100644 index 00000000000..43119213ff8 --- /dev/null +++ b/internal/api/handlers/management/api_tools_usage_test.go @@ -0,0 +1,66 @@ +package management + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gin-gonic/gin" + coreauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" +) + +type managementUsageReviewSink struct { + authID string + records []usage.Record +} + +func (*managementUsageReviewSink) Synchronous() bool { return true } +func (s *managementUsageReviewSink) HandleUsage(_ context.Context, r usage.Record) { + if r.AuthID == s.authID { + s.records = append(s.records, r) + } +} +func TestManagementUsageUsesProviderProtocol(t *testing.T) { + for _, tc := range []struct{ name, provider, path, body string }{ + {"gemini", "gemini", "/v1beta/models/test:generateContent", `{"usageMetadata":{"promptTokenCount":100,"candidatesTokenCount":20,"thoughtsTokenCount":10,"totalTokenCount":130}}`}, + {"antigravity", "antigravity", "/v1internal:generateContent", `{"response":{"usageMetadata":{"promptTokenCount":100,"candidatesTokenCount":20,"thoughtsTokenCount":10,"totalTokenCount":130}}}`}, + {"claude_stream", "claude", "/v1/messages", "data: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":100,\"output_tokens\":1}}}\n\ndata: {\"type\":\"message_delta\",\"usage\":{\"output_tokens\":30}}\n\ndata: {\"type\":\"message_stop\"}\n\n"}, + } { + t.Run(tc.name, func(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if strings.Contains(tc.body, "data:") { + w.Header().Set("Content-Type", "text/event-stream") + } + _, _ = w.Write([]byte(tc.body)) + })) + defer upstream.Close() + manager := coreauth.NewManager(nil, nil, nil) + auth := &coreauth.Auth{ID: t.Name(), Provider: tc.provider} + auth.EnsureIndex() + if _, err := manager.Register(context.Background(), auth); err != nil { + t.Fatal(err) + } + sink := &managementUsageReviewSink{authID: t.Name()} + usage.RegisterPlugin(sink) + h := &Handler{authManager: manager} + router := gin.New() + router.POST("/", h.APICall) + body, _ := json.Marshal(map[string]any{"method": "POST", "url": upstream.URL + tc.path, "auth_index": auth.Index, "data": `{"model":"test-model"}`}) + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(string(body))) + req.Header.Set("Content-Type", "application/json") + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, req) + if len(sink.records) != 1 { + t.Fatalf("want one usage event, got %d (HTTP %d)", len(sink.records), recorder.Code) + } + d := sink.records[0].Detail + if !d.UsageObserved || d.TotalTokens != 130 { + t.Fatalf("lost provider usage: %+v", d) + } + }) + } +} diff --git a/internal/api/usage_coverage.go b/internal/api/usage_coverage.go index b51d2b7961d..ba2770e736e 100644 --- a/internal/api/usage_coverage.go +++ b/internal/api/usage_coverage.go @@ -29,6 +29,12 @@ func usageCoverageMiddleware() gin.HandlerFunc { if scope.Published() { return } + // Do not turn rejected credentials or unknown routes into unexpired + // journal files. Coverage is for accepted, registered API operations. + status := c.Writer.Status() + if c.FullPath() == "" || c.IsAborted() && (status == http.StatusUnauthorized || status == http.StatusForbidden) { + return + } usage.PublishRecord(context.WithValue(ctx, "gin", c), usage.Record{Provider: "unknown", Model: "unknown", ExecutorType: "EndpointCoverage", Endpoint: c.Request.Method + " " + path, Kind: "unmeasured", Generate: usage.GenerateFlag(false), RequestedAt: started, Latency: time.Since(started), Failed: c.Writer.Status() >= http.StatusBadRequest, Fail: usage.Failure{StatusCode: c.Writer.Status()}}) } } diff --git a/internal/api/usage_coverage_test.go b/internal/api/usage_coverage_test.go new file mode 100644 index 00000000000..0870187a382 --- /dev/null +++ b/internal/api/usage_coverage_test.go @@ -0,0 +1,78 @@ +package api + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gin-gonic/gin" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/api/handlers" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" + sdkconfig "github.com/router-for-me/CLIProxyAPI/v7/sdk/config" +) + +type coverageReviewSink struct { + model string + records []usage.Record +} + +func (*coverageReviewSink) Synchronous() bool { return true } +func (s *coverageReviewSink) HandleUsage(_ context.Context, r usage.Record) { + if r.Model == s.model || r.ExecutorType == "EndpointCoverage" && strings.HasSuffix(r.Endpoint, "/"+s.model) { + s.records = append(s.records, r) + } +} +func TestCoverageSurvivesNormalHandlerContext(t *testing.T) { + sink := &coverageReviewSink{model: t.Name()} + usage.RegisterPlugin(sink) + router := gin.New() + router.Use(usageCoverageMiddleware()) + router.POST("/v1/"+t.Name(), func(c *gin.Context) { + h := &handlers.BaseAPIHandler{Cfg: &sdkconfig.SDKConfig{}} + ctx, cancel := h.GetContextWithCancel(nil, c, context.Background()) + defer cancel() + usage.PublishRecord(ctx, usage.Record{Model: t.Name()}) + c.Status(http.StatusOK) + }) + router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodPost, "/v1/"+t.Name(), nil)) + if len(sink.records) != 1 { + t.Fatalf("measured operation emitted %d records, want 1", len(sink.records)) + } + if sink.records[0].GenerationID == "" { + t.Fatal("executor lost generation identity") + } +} + +func TestCoverageSkipsRejectedAndUnmatchedRequests(t *testing.T) { + for _, tc := range []struct { + name string + status int + matched bool + want int + }{ + {"unauthorized", 401, true, 0}, {"forbidden", 403, true, 0}, {"unmatched", 404, false, 0}, {"accepted_control", 200, true, 1}, + } { + t.Run(tc.name, func(t *testing.T) { + sink := &coverageReviewSink{model: t.Name()} + usage.RegisterPlugin(sink) + router := gin.New() + router.Use(usageCoverageMiddleware()) + path := "/v1/" + t.Name() + if tc.matched { + router.POST(path, func(c *gin.Context) { + if tc.status >= 400 { + c.AbortWithStatus(tc.status) + } else { + c.Status(tc.status) + } + }) + } + router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodPost, path, nil)) + if len(sink.records) != tc.want { + t.Fatalf("record count=%d, want %d", len(sink.records), tc.want) + } + }) + } +} diff --git a/internal/redisqueue/plugin.go b/internal/redisqueue/plugin.go index a16304c4d0e..cb2bc394169 100644 --- a/internal/redisqueue/plugin.go +++ b/internal/redisqueue/plugin.go @@ -156,7 +156,7 @@ func (p *usageQueuePlugin) HandleUsage(ctx context.Context, record coreusage.Rec } payload, err := json.Marshal(queuedUsageDetail{ BillingID: usageDetail.BillingID, CostScope: usageDetail.CostScope, GenerationID: record.GenerationID, EventID: eventID, AttemptID: record.AttemptID, Kind: kind, Transport: transport, BaseURL: record.BaseURL, - UsageObserved: usageDetail.UsageObserved, UsageComplete: !record.Failed, RawUsage: rawUsage, + UsageObserved: usageDetail.UsageObserved, UsageComplete: !failed, RawUsage: rawUsage, CacheCreation5mTokens: usageDetail.CacheCreation5mTokens, CacheCreation1hTokens: usageDetail.CacheCreation1hTokens, CostUSD: usageDetail.CostUSD, diff --git a/internal/redisqueue/usage_completeness_test.go b/internal/redisqueue/usage_completeness_test.go index 6ca4198de91..8daabad184e 100644 --- a/internal/redisqueue/usage_completeness_test.go +++ b/internal/redisqueue/usage_completeness_test.go @@ -2,6 +2,7 @@ package redisqueue import ( "context" + "github.com/router-for-me/CLIProxyAPI/v7/internal/logging" "testing" coreusage "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" @@ -16,3 +17,15 @@ func TestFailedPartialUsageIsNotPublishedAsComplete(t *testing.T) { } }) } + +func TestResolvedFailureMarksUsageIncomplete(t *testing.T) { + withEnabledQueue(t, func() { + ctx := logging.WithResponseStatusHolder(context.Background()) + logging.SetResponseStatus(ctx, 500) + (&usageQueuePlugin{}).HandleUsage(ctx, coreusage.Record{Provider: "claude", Model: "claude-opus-5", Detail: coreusage.Detail{InputTokens: 100, OutputTokens: 1, TotalTokens: 101, UsageObserved: true}}) + payload := popSinglePayload(t) + if string(payload["failed"]) != "true" || string(payload["usage_complete"]) != "false" { + t.Fatalf("inconsistent completeness: failed=%s complete=%s", payload["failed"], payload["usage_complete"]) + } + }) +} diff --git a/sdk/api/handlers/handlers.go b/sdk/api/handlers/handlers.go index 32b21f8c0d2..6d8e404be8c 100644 --- a/sdk/api/handlers/handlers.go +++ b/sdk/api/handlers/handlers.go @@ -487,6 +487,7 @@ func (h *BaseAPIHandler) GetContextWithCancel(handler interfaces.APIHandler, c * parentCtx = logging.WithRequestID(parentCtx, requestID) } } + parentCtx = coreusage.WithAccountingContext(parentCtx, requestCtx) newCtx, cancel := context.WithCancel(parentCtx) endpoint := "" diff --git a/sdk/cliproxy/usage/manager.go b/sdk/cliproxy/usage/manager.go index 5aa5f30f703..ff5aebd21c5 100644 --- a/sdk/cliproxy/usage/manager.go +++ b/sdk/cliproxy/usage/manager.go @@ -7,6 +7,7 @@ import ( "sync" "time" + "github.com/google/uuid" log "github.com/sirupsen/logrus" ) @@ -366,6 +367,9 @@ func (m *Manager) Publish(ctx context.Context, record Record) { if closed { return } + if strings.TrimSpace(record.EventID) == "" { + record.EventID = uuid.NewString() + } markPublished(ctx) if record.GenerationID == "" { record.GenerationID = GenerationFromContext(ctx) diff --git a/sdk/cliproxy/usage/scope.go b/sdk/cliproxy/usage/scope.go index c40b840232f..49991185c36 100644 --- a/sdk/cliproxy/usage/scope.go +++ b/sdk/cliproxy/usage/scope.go @@ -33,3 +33,26 @@ func markPublished(ctx context.Context) { } } } + +// WithAccountingContext carries accounting identity across an execution context +// boundary without changing its cancellation or copying unrelated request values. +// An explicitly chosen execution generation takes precedence over the request. +func WithAccountingContext(target, source context.Context) context.Context { + if target == nil { + target = context.Background() + } + if source == nil { + return target + } + if target.Value(scopeKey{}) == nil { + if scope, ok := source.Value(scopeKey{}).(*AccountingScope); ok { + target = context.WithValue(target, scopeKey{}, scope) + } + } + if GenerationFromContext(target) == "" { + if id := GenerationFromContext(source); id != "" { + target = context.WithValue(target, generationKey{}, id) + } + } + return target +} diff --git a/sdk/cliproxy/usage/scope_test.go b/sdk/cliproxy/usage/scope_test.go index 8c22c676055..b4294f0c55b 100644 --- a/sdk/cliproxy/usage/scope_test.go +++ b/sdk/cliproxy/usage/scope_test.go @@ -36,3 +36,33 @@ func TestGenerationIdentityChangesWithoutLosingCoverageScope(t *testing.T) { t.Fatal("generation context lost request coverage") } } + +type identityReviewSink struct { + synchronous bool + records chan Record +} + +func (s *identityReviewSink) Synchronous() bool { return s.synchronous } +func (s *identityReviewSink) HandleUsage(_ context.Context, r Record) { s.records <- r } +func TestEventIdentityIsSharedBySynchronousAndAsyncPlugins(t *testing.T) { + m := NewManager(1) + defer m.Stop() + a := &identityReviewSink{synchronous: true, records: make(chan Record, 2)} + b := &identityReviewSink{records: make(chan Record, 2)} + m.Register(a) + m.Register(b) + for _, provided := range []string{"", "caller-supplied"} { + m.Publish(context.Background(), Record{EventID: provided}) + syncRecord := <-a.records + if syncRecord.EventID == "" { + t.Fatal("manager did not assign an event ID") + } + asyncRecord := <-b.records + if syncRecord.EventID != asyncRecord.EventID { + t.Fatal("plugins received different IDs") + } + if provided != "" && syncRecord.EventID != provided { + t.Fatal("caller identity was replaced") + } + } +} From 84b18c0e4cdd48710fd55eb6e2a9002012e9eaf8 Mon Sep 17 00:00:00 2001 From: Ramapitecus Date: Sat, 12 Sep 2026 19:06:16 +0500 Subject: [PATCH 5/9] fix usage for URL models and opaque journal identities --- docs/usage-accounting.md | 2 + internal/api/handlers/management/api_tools.go | 22 +++- .../management/api_tools_usage_test.go | 61 ++++++++++ internal/redisqueue/journal.go | 60 +++++++++- internal/redisqueue/journal_test.go | 113 +++++++++++++++++- 5 files changed, 252 insertions(+), 6 deletions(-) diff --git a/docs/usage-accounting.md b/docs/usage-accounting.md index b2704e5cfb6..2a58929eef6 100644 --- a/docs/usage-accounting.md +++ b/docs/usage-accounting.md @@ -48,3 +48,5 @@ Live terminal usage is captured on the upstream read side, including a complete The common HTTP handler context bridge now copies accounting scope and generation identity while preserving the execution context's existing cancellation behavior. Measured HTTP operations therefore suppress the unmeasured fallback correctly. The usage manager assigns a missing event ID once before either synchronous or asynchronous plugins run, preserving caller-supplied identities. Coverage skips unmatched routes and locally rejected authorization, so unauthenticated requests cannot create permanent journal entries through that fallback. `usage_complete` uses the same resolved failure state as the emitted `failed` field. Model-bearing management calls select the provider/endpoint protocol parser and merge SSE usage; partial response bodies are retained when reading fails. + +Management Gemini/Vertex generation calls also resolve models from `/models/{model}:generateContent` and `:streamGenerateContent` paths when the request body omits `model`; explicit body models remain authoritative. Opaque event IDs are preserved in payloads and mapped to SHA-256 filenames for durable storage and ACK. Previously written journal files remain readable and acknowledgeable. diff --git a/internal/api/handlers/management/api_tools.go b/internal/api/handlers/management/api_tools.go index ac9e25ba19c..e306d95137a 100644 --- a/internal/api/handlers/management/api_tools.go +++ b/internal/api/handlers/management/api_tools.go @@ -195,7 +195,10 @@ func (h *Handler) APICall(c *gin.Context) { // Management model probes bypass normal executors and need their own record. var reporter *helps.UsageReporter usageCtx := context.WithValue(c.Request.Context(), "gin", c) - model := gjson.Get(body.Data, "model").String() + model := strings.TrimSpace(gjson.Get(body.Data, "model").String()) + if model == "" { + model = managementModelFromPath(req.URL.Path) + } provider := "openai-compatible" if auth != nil { provider = auth.Provider @@ -242,6 +245,23 @@ func (h *Handler) APICall(c *gin.Context) { }) } +// Gemini and Vertex encode generation models in the URL instead of the body. +func managementModelFromPath(path string) string { + prefix, action, ok := strings.Cut(path, ":") + if !ok || (action != "generateContent" && action != "streamGenerateContent") { + return "" + } + index := strings.LastIndex(prefix, "/models/") + if index < 0 { + return "" + } + model := prefix[index+len("/models/"):] + if strings.Contains(model, "/") { + return "" + } + return strings.TrimSpace(model) +} + func firstNonEmptyString(values ...*string) string { for _, v := range values { if v == nil { diff --git a/internal/api/handlers/management/api_tools_usage_test.go b/internal/api/handlers/management/api_tools_usage_test.go index 43119213ff8..e716c59f1ec 100644 --- a/internal/api/handlers/management/api_tools_usage_test.go +++ b/internal/api/handlers/management/api_tools_usage_test.go @@ -64,3 +64,64 @@ func TestManagementUsageUsesProviderProtocol(t *testing.T) { }) } } + +func TestManagementUsageDerivesModelFromURL(t *testing.T) { + for _, tc := range []struct { + name, provider, path, data, wantModel string + }{ + {"gemini", "gemini", "/v1beta/models/gemini-2.5-pro:generateContent", `{"contents":[]}`, "gemini-2.5-pro"}, + {"gemini_stream", "gemini", "/v1beta/models/gemini-2.5-flash:streamGenerateContent?alt=sse", `{"contents":[]}`, "gemini-2.5-flash"}, + {"vertex", "vertex", "/v1/projects/demo/locations/global/publishers/google/models/gemini-2.5-pro:generateContent", `{"contents":[]}`, "gemini-2.5-pro"}, + {"antigravity", "antigravity", "/v1beta/models/gemini-2.5-pro:generateContent", `{"contents":[]}`, "gemini-2.5-pro"}, + {"without_provider", "", "/v1beta/models/gemini-2.5-pro:generateContent", `{"contents":[]}`, "gemini-2.5-pro"}, + {"explicit_model", "gemini", "/v1beta/models/path-model:generateContent", `{"model":"body-model"}`, "body-model"}, + {"unrelated_path", "gemini", "/v1/models/model:delete", `{}`, ""}, + } { + t.Run(tc.name, func(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + payload := `{"usageMetadata":{"promptTokenCount":100,"candidatesTokenCount":30,"totalTokenCount":130}}` + if tc.provider == "antigravity" { + payload = `{"response":` + payload + `}` + } + if strings.Contains(tc.path, "streamGenerateContent") { + w.Header().Set("Content-Type", "text/event-stream") + payload = "data: " + payload + "\n\n" + } + _, _ = w.Write([]byte(payload)) + })) + defer upstream.Close() + manager := coreauth.NewManager(nil, nil, nil) + auth := &coreauth.Auth{ID: t.Name(), Provider: tc.provider} + auth.EnsureIndex() + if _, err := manager.Register(context.Background(), auth); err != nil { + t.Fatal(err) + } + sink := &managementUsageReviewSink{authID: t.Name()} + usage.RegisterPlugin(sink) + h := &Handler{authManager: manager} + router := gin.New() + router.POST("/", h.APICall) + body, _ := json.Marshal(map[string]any{"method": "POST", "url": upstream.URL + tc.path, "auth_index": auth.Index, "data": tc.data}) + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(string(body))) + req.Header.Set("Content-Type", "application/json") + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, req) + if recorder.Code != http.StatusOK { + t.Fatalf("HTTP %d: %s", recorder.Code, recorder.Body.String()) + } + if tc.wantModel == "" { + if len(sink.records) != 0 { + t.Fatalf("unrelated URL generated usage: %+v", sink.records) + } + return + } + if len(sink.records) != 1 { + t.Fatalf("want one usage event, got %d", len(sink.records)) + } + r := sink.records[0] + if r.Model != tc.wantModel || !r.Detail.UsageObserved || r.Detail.TotalTokens != 130 { + t.Fatalf("lost model or usage: %+v", r) + } + }) + } +} diff --git a/internal/redisqueue/journal.go b/internal/redisqueue/journal.go index d476521c314..afd1155ce14 100644 --- a/internal/redisqueue/journal.go +++ b/internal/redisqueue/journal.go @@ -34,6 +34,17 @@ func ConfigureUsageJournal(dir string) { journal.lastErr = nil } func validEventID(id string) bool { + return id != "" +} + +// Public IDs are opaque; they must never be used as filesystem paths. The dot +// in this prefix separates new names from the legacy filename alphabet. +func journalEventFilename(id string) string { + sum := sha256.Sum256([]byte(id)) + return "sha256." + hex.EncodeToString(sum[:]) + ".json" +} + +func legacyJournalEventID(id string) bool { if len(id) == 0 || len(id) > 128 { return false } @@ -42,8 +53,35 @@ func validEventID(id string) bool { return false } } - return true + // On Windows this also rejects reserved device names such as CON.json. + // Those names could never have been ordinary legacy journal files. + return filepath.IsLocal(id + ".json") } + +// Check payload identity as well as the legacy filename. Case-insensitive or +// Unicode-normalizing filesystems can resolve a different public ID to the +// same legacy path; that event must not be deduplicated or acknowledged. +func (j *usageJournal) matchingLegacyEventPath(id string) (string, error) { + if !legacyJournalEventID(id) { + return "", nil + } + path := filepath.Join(j.dir, id+".json") + payload, errRead := os.ReadFile(path) + if os.IsNotExist(errRead) { + return "", nil + } + if errRead != nil { + return "", errRead + } + var event struct { + EventID string `json:"event_id"` + } + if json.Unmarshal(payload, &event) != nil || event.EventID != id { + return "", nil + } + return path, nil +} + func (j *usageJournal) append(payload []byte) error { j.mu.Lock() defer j.mu.Unlock() @@ -89,9 +127,16 @@ func (j *usageJournal) append(payload []byte) error { if err := os.MkdirAll(j.dir, 0700); err != nil { return err } - target := filepath.Join(j.dir, event.EventID+".json") + target := filepath.Join(j.dir, journalEventFilename(event.EventID)) if _, err := os.Stat(target); err == nil { return nil + } else if !os.IsNotExist(err) { + return err + } + if legacy, errLegacy := j.matchingLegacyEventPath(event.EventID); errLegacy != nil { + return errLegacy + } else if legacy != "" { + return nil } file, err := os.CreateTemp(j.dir, ".pending-") if err != nil { @@ -164,9 +209,18 @@ func (j *usageJournal) ack(ids []string) error { } } for _, id := range ids { - if err := os.Remove(filepath.Join(j.dir, id+".json")); err != nil && !os.IsNotExist(err) { + if err := os.Remove(filepath.Join(j.dir, journalEventFilename(id))); err != nil && !os.IsNotExist(err) { return err } + legacy, errLegacy := j.matchingLegacyEventPath(id) + if errLegacy != nil { + return errLegacy + } + if legacy != "" { + if errRemove := os.Remove(legacy); errRemove != nil && !os.IsNotExist(errRemove) { + return errRemove + } + } } if len(ids) == 0 { return nil diff --git a/internal/redisqueue/journal_test.go b/internal/redisqueue/journal_test.go index f78543da18d..41264fbfca4 100644 --- a/internal/redisqueue/journal_test.go +++ b/internal/redisqueue/journal_test.go @@ -1,7 +1,12 @@ package redisqueue import ( + "crypto/sha256" + "encoding/hex" "encoding/json" + "os" + "path/filepath" + "strings" "testing" ) @@ -33,8 +38,112 @@ func TestUsageJournalReplayRequiresAcknowledgement(t *testing.T) { if err != nil || len(items) != 1 { t.Fatalf("ack: %d %v", len(items), err) } - if err := restarted.ack([]string{"../../other"}); err == nil { - t.Fatal("unsafe id accepted") + if err := restarted.ack([]string{""}); err == nil { + t.Fatal("empty event ID accepted") + } +} + +func TestUsageJournalOpaqueEventIDsRoundTrip(t *testing.T) { + root := t.TempDir() + j := &usageJournal{dir: filepath.Join(root, "journal")} + outside := filepath.Join(root, "outside.json") + if err := os.WriteFile(outside, []byte("keep"), 0600); err != nil { + t.Fatal(err) + } + ids := []string{"usage:123", "evt.123", "../outside", `C:\event\123`, "Case", "case", "事件/123", "event\x00suffix", "CON", "NUL", "COM1", "LPT1", strings.Repeat("long-id:", 100)} + for _, id := range ids { + payload, err := json.Marshal(map[string]string{"event_id": id, "request_id": "same"}) + if err != nil { + t.Fatal(err) + } + for attempt := 0; attempt < 2; attempt++ { + if err := j.append(payload); err != nil { + t.Fatalf("opaque ID %q rejected: %v", id, err) + } + } + } + restarted := &usageJournal{dir: j.dir} + items, err := restarted.read(100) + if err != nil || len(items) != len(ids) { + t.Fatalf("replay: got %d items, want %d: %v", len(items), len(ids), err) + } + seen := make(map[string]bool) + for _, payload := range items { + var event struct { + EventID string `json:"event_id"` + } + if err := json.Unmarshal(payload, &event); err != nil { + t.Fatal(err) + } + seen[event.EventID] = true + } + for _, id := range ids { + if !seen[id] { + t.Fatalf("public ID changed: %q", id) + } + if err := restarted.ack([]string{id}); err != nil { + t.Fatalf("ACK %q failed: %v", id, err) + } + } + if items, err := restarted.read(100); err != nil || len(items) != 0 { + t.Fatalf("ACK left %d items: %v", len(items), err) + } + if data, err := os.ReadFile(outside); err != nil || string(data) != "keep" { + t.Fatalf("opaque traversal ID changed an outside file: %q %v", data, err) + } +} + +func TestUsageJournalLegacyFilesRemainReplayableAndAcknowledged(t *testing.T) { + j := &usageJournal{dir: t.TempDir()} + legacy := []byte(`{"event_id":"Legacy-ID","request_id":"legacy"}`) + if err := os.WriteFile(filepath.Join(j.dir, "Legacy-ID.json"), legacy, 0600); err != nil { + t.Fatal(err) + } + if err := j.append(legacy); err != nil { + t.Fatal(err) + } + if items, err := j.read(10); err != nil || len(items) != 1 { + t.Fatalf("legacy replay duplicated: %d %v", len(items), err) + } + if err := j.append([]byte(`{"event_id":"legacy-id","request_id":"new"}`)); err != nil { + t.Fatal(err) + } + if items, err := j.read(10); err != nil || len(items) != 2 { + t.Fatalf("distinct case-sensitive public IDs collided: %d %v", len(items), err) + } + if err := j.ack([]string{"legacy-id"}); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(j.dir, "Legacy-ID.json")); err != nil { + t.Fatalf("ACK removed a different legacy ID: %v", err) + } + if err := j.ack([]string{"Legacy-ID"}); err != nil { + t.Fatal(err) + } + if items, err := j.read(10); err != nil || len(items) != 0 { + t.Fatalf("legacy ACK left %d items: %v", len(items), err) + } +} + +func TestUsageJournalHashedNamesCannotCollideWithLegacyIDs(t *testing.T) { + j := &usageJournal{dir: t.TempDir()} + sum := sha256.Sum256([]byte("usage:123")) + legacyID := hex.EncodeToString(sum[:]) + legacy, _ := json.Marshal(map[string]string{"event_id": legacyID}) + if err := os.WriteFile(filepath.Join(j.dir, legacyID+".json"), legacy, 0600); err != nil { + t.Fatal(err) + } + if err := j.append([]byte(`{"event_id":"usage:123"}`)); err != nil { + t.Fatal(err) + } + if items, err := j.read(10); err != nil || len(items) != 2 { + t.Fatalf("hash collided with legacy file: %d %v", len(items), err) + } + if err := j.ack([]string{"usage:123"}); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(j.dir, legacyID+".json")); err != nil { + t.Fatalf("hashed ACK removed legacy event: %v", err) } } From 0a77686175e424a427a8e010cb21e49e598933bd Mon Sep 17 00:00:00 2001 From: Ramapitecus Date: Sat, 12 Sep 2026 19:24:05 +0500 Subject: [PATCH 6/9] fix admitted request coverage and streaming usage edge cases --- docs/usage-accounting.md | 2 + internal/api/server_alpha_usage_test.go | 136 ++++++++++++++++++ internal/api/server_middleware.go | 3 + internal/api/server_routes.go | 30 ++-- internal/api/server_test.go | 11 +- internal/api/usage_coverage.go | 9 +- internal/api/usage_coverage_test.go | 41 +++++- .../runtime/executor/helps/usage_helpers.go | 6 +- .../executor/helps/usage_helpers_test.go | 14 +- .../executor/helps/usage_integrity_test.go | 45 ++++++ 10 files changed, 266 insertions(+), 31 deletions(-) create mode 100644 internal/api/server_alpha_usage_test.go diff --git a/docs/usage-accounting.md b/docs/usage-accounting.md index 2a58929eef6..ba374b2454e 100644 --- a/docs/usage-accounting.md +++ b/docs/usage-accounting.md @@ -50,3 +50,5 @@ The common HTTP handler context bridge now copies accounting scope and generatio Coverage skips unmatched routes and locally rejected authorization, so unauthenticated requests cannot create permanent journal entries through that fallback. `usage_complete` uses the same resolved failure state as the emitted `failed` field. Model-bearing management calls select the provider/endpoint protocol parser and merge SSE usage; partial response bodies are retained when reading fails. Management Gemini/Vertex generation calls also resolve models from `/models/{model}:generateContent` and `:streamGenerateContent` paths when the request body omits `model`; explicit body models remain authoritative. Opaque event IDs are preserved in payloads and mapped to SHA-256 filenames for durable storage and ACK. Previously written journal files remain readable and acknowledgeable. + +Coverage fallback requires successful route admission, including legacy authentication-disabled and realtime client-secret paths. Requests blocked by Home heartbeat or other gates before admission do not create journal files; admitted handler failures remain visible. Alpha Search creates its reporter before the upstream attempt so transport/read failures and latency retain their source attribution. Gemini and Interactions streams preserve explicitly reported zero usage, including when a response tier is present. diff --git a/internal/api/server_alpha_usage_test.go b/internal/api/server_alpha_usage_test.go new file mode 100644 index 00000000000..193ffd9e71e --- /dev/null +++ b/internal/api/server_alpha_usage_test.go @@ -0,0 +1,136 @@ +package api + +import ( + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/router-for-me/CLIProxyAPI/v7/internal/config" + "github.com/router-for-me/CLIProxyAPI/v7/internal/registry" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executionregistry" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" +) + +type alphaSearchUsageCapture struct { + mu sync.Mutex + authID string + records []usage.Record +} + +func (*alphaSearchUsageCapture) Synchronous() bool { return true } +func (p *alphaSearchUsageCapture) HandleUsage(_ context.Context, record usage.Record) { + p.mu.Lock() + defer p.mu.Unlock() + if record.AuthID == p.authID { + p.records = append(p.records, record) + } +} + +func TestAlphaSearchUsageCoversAttemptAndReadFailures(t *testing.T) { + const response = `{"results":[],"usage":{"input_tokens":100,"output_tokens":10,"total_tokens":110}}` + for _, scenario := range []struct { + name string + prepareErr error + httpErr error + readErr bool + status int + wantStatus int + wantFailed bool + wantTokens int64 + }{ + {name: "success", status: 200, wantStatus: 200, wantTokens: 110}, + {name: "upstream_error_usage", status: 500, wantStatus: 500, wantFailed: true, wantTokens: 110}, + {name: "request_preparation", prepareErr: errors.New("prepare failed"), wantStatus: 502, wantFailed: true}, + {name: "connection_error", httpErr: errors.New("upstream connection failed"), wantStatus: 502, wantFailed: true}, + {name: "body_read_error_with_usage", status: 200, readErr: true, wantStatus: 502, wantFailed: true, wantTokens: 110}, + } { + t.Run(scenario.name, func(t *testing.T) { + server := newTestServer(t) + capture := &alphaSearchUsageCapture{authID: t.Name()} + usage.RegisterNamedPlugin(t.Name(), capture) + t.Cleanup(func() { usage.RegisterNamedPlugin(t.Name(), &alphaSearchUsageCapture{}) }) + var upstreamReached time.Time + executor := &codexSearchCaptureExecutor{prepareErr: scenario.prepareErr, httpErr: scenario.httpErr, statuses: []int{scenario.status}, responseBody: io.NopCloser(strings.NewReader(response)), beforeReturn: func() { upstreamReached = time.Now() }} + if scenario.readErr { + executor.responseBody = &errorSearchResponseBody{payload: []byte(response)} + } + server.handlers.AuthManager.RegisterExecutor(executor) + credential := &auth.Auth{ID: t.Name(), Provider: "codex", Status: auth.StatusActive, Metadata: map[string]any{"access_token": "test-token"}} + if _, err := server.handlers.AuthManager.Register(context.Background(), credential); err != nil { + t.Fatal(err) + } + registry.GetGlobalRegistry().RegisterClient(credential.ID, "codex", []*registry.ModelInfo{{ID: "gpt-5.6-sol"}}) + t.Cleanup(func() { registry.GetGlobalRegistry().UnregisterClient(credential.ID) }) + rr := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rr) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/alpha/search", strings.NewReader(`{"model":"gpt-5.6-sol","query":"test"}`)) + server.codexAlphaSearch(c) + if rr.Code != scenario.wantStatus { + t.Fatalf("HTTP status=%d want %d: %s", rr.Code, scenario.wantStatus, rr.Body.String()) + } + capture.mu.Lock() + defer capture.mu.Unlock() + if len(capture.records) != 1 { + t.Fatalf("got %d selected-credential usage records, want 1", len(capture.records)) + } + r := capture.records[0] + if r.Failed != scenario.wantFailed || r.Detail.TotalTokens != scenario.wantTokens || r.Model != "gpt-5.6-sol" { + t.Fatalf("incorrect attempt usage: %+v", r) + } + if !upstreamReached.IsZero() && (r.RequestedAt.After(upstreamReached) || r.Latency < upstreamReached.Sub(r.RequestedAt)) { + t.Fatalf("usage timer started after upstream attempt: requested=%v upstream=%v latency=%v", r.RequestedAt, upstreamReached, r.Latency) + } + }) + } +} + +func TestAlphaSearchHomeUnauthorizedPublishesOneAttempt(t *testing.T) { + const response = `{"error":{"message":"access token expired"},"usage":{"input_tokens":100,"output_tokens":10,"total_tokens":110}}` + for _, mode := range []string{"response", "read_failure", "bind_failure"} { + t.Run(mode, func(t *testing.T) { + server := newTestServer(t) + capture := &alphaSearchUsageCapture{authID: t.Name()} + usage.RegisterNamedPlugin(t.Name(), capture) + t.Cleanup(func() { usage.RegisterNamedPlugin(t.Name(), &alphaSearchUsageCapture{}) }) + attemptRegistry := executionregistry.New() + server.handlers.AuthManager.SetConfig(&config.Config{Home: config.HomeConfig{Enabled: true}}) + server.handlers.AuthManager.PublishHomeDispatch(&codexSearchHomeDispatcher{authID: t.Name()}, attemptRegistry, 1) + executor := &codexSearchCaptureExecutor{statuses: []int{http.StatusUnauthorized}, responseBody: io.NopCloser(strings.NewReader(response))} + if mode == "read_failure" { + executor.responseBody = &errorSearchResponseBody{payload: []byte(response)} + } + if mode == "bind_failure" { + executor.beforeReturn = func() { _ = attemptRegistry.Close() } + } + server.handlers.AuthManager.RegisterExecutor(executor) + rr := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rr) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/alpha/search", strings.NewReader(`{"model":"gpt-5.6-sol","query":"test"}`)) + server.codexAlphaSearch(c) + capture.mu.Lock() + defer capture.mu.Unlock() + if len(capture.records) != 1 { + t.Fatalf("one upstream attempt generated %d accounting events", len(capture.records)) + } + r := capture.records[0] + wantBody, wantTokens := response, int64(110) + if mode == "bind_failure" { + wantBody, wantTokens = "upstream unauthorized", 0 + } + if !r.Failed || r.Fail.StatusCode != http.StatusUnauthorized || r.Fail.Body != wantBody || r.Detail.TotalTokens != wantTokens { + t.Fatalf("incorrect Home 401 failure: %+v", r) + } + if r.AuthIndex == "" || r.AccessTokenSHA256 == "" { + t.Fatal("Home result lost credential attribution") + } + }) + } +} diff --git a/internal/api/server_middleware.go b/internal/api/server_middleware.go index 447f280cf82..c413853e45c 100644 --- a/internal/api/server_middleware.go +++ b/internal/api/server_middleware.go @@ -159,6 +159,7 @@ func realtimeStandardAuthMiddleware(manager *sdkaccess.Manager) gin.HandlerFunc func accessAuthMiddleware(manager *sdkaccess.Manager, realtimeError bool) gin.HandlerFunc { return func(c *gin.Context) { if manager == nil { + c.Set(usageRequestAcceptedKey, true) c.Next() return } @@ -172,6 +173,7 @@ func accessAuthMiddleware(manager *sdkaccess.Manager, realtimeError bool) gin.Ha c.Set("accessMetadata", result.Metadata) } } + c.Set(usageRequestAcceptedKey, true) c.Next() return } @@ -228,6 +230,7 @@ func realtimeAuthMiddleware(manager *sdkaccess.Manager, handler *codexlive.Handl c.Set("accessProvider", provider) c.Set(codexlive.ClientSecretSessionContextKey, authorization.Session) c.Set(codexlive.ClientSecretPrincipalContextKey, authorization.Principal) + c.Set(usageRequestAcceptedKey, true) c.Next() } } diff --git a/internal/api/server_routes.go b/internal/api/server_routes.go index a28b0ab1063..56fd006a541 100644 --- a/internal/api/server_routes.go +++ b/internal/api/server_routes.go @@ -391,6 +391,8 @@ func (s *Server) codexAlphaSearch(c *gin.Context) { defer releaseAttempt() } logging.SetGinCPATraceID(c, selected.EnsureIndex()) + reporter := helps.NewUsageReporter(ctx, "codex", selectionModel, selected) + defer reporter.EnsurePublished(ctx) baseHeaders := make(http.Header) baseHeaders.Set("Content-Type", "application/json") @@ -449,6 +451,7 @@ func (s *Server) codexAlphaSearch(c *gin.Context) { } if errCtx := ctx.Err(); errCtx != nil { + reporter.PublishFailure(ctx, errCtx) if selection != nil { selection.End("attempt_canceled") } @@ -457,6 +460,7 @@ func (s *Server) codexAlphaSearch(c *gin.Context) { } resp, err := performRequest(selected) if err != nil { + reporter.PublishFailure(ctx, err) if errors.Is(err, errMissingBaseURL) { if selection != nil { selection.End("missing_base_url") @@ -478,11 +482,22 @@ func (s *Server) codexAlphaSearch(c *gin.Context) { } return errClose } - if selection != nil { - if errBind := selection.Bind(closeResponseBody); errBind != nil { + responseFailure := func(body []byte, fallback error) error { + if resp.StatusCode < http.StatusBadRequest { + return fallback + } + message := strings.TrimSpace(string(body)) + if message == "" { + message = "search returned HTTP " + strconv.Itoa(resp.StatusCode) if resp.StatusCode == http.StatusUnauthorized { - s.handlers.AuthManager.ReportHomeUnauthorized(ctx, selected, "codex", selectionModel) + message = "upstream unauthorized" } + } + return &auth.Error{HTTPStatus: resp.StatusCode, Message: message} + } + if selection != nil { + if errBind := selection.Bind(closeResponseBody); errBind != nil { + reporter.PublishFailure(ctx, responseFailure(nil, errBind)) selection.End("response_bind_failed") c.JSON(http.StatusServiceUnavailable, gin.H{"error": errBind.Error()}) return @@ -494,24 +509,20 @@ func (s *Server) codexAlphaSearch(c *gin.Context) { helps.RecordAPIResponseMetadata(ctx, s.cfg, resp.StatusCode, resp.Header.Clone()) upstreamBody, err := io.ReadAll(io.LimitReader(resp.Body, 32<<20)) if err != nil { + reporter.PublishFailureWithDetail(ctx, helps.ParseOpenAIUsage(upstreamBody), responseFailure(upstreamBody, err)) helps.AppendAPIResponseChunk(ctx, s.cfg, upstreamBody) - if selection != nil && resp.StatusCode == http.StatusUnauthorized { - s.handlers.AuthManager.ReportHomeUnauthorized(ctx, selected, "codex", selectionModel, upstreamBody) - } helps.RecordAPIResponseError(ctx, s.cfg, err) c.JSON(clienterror.HTTPStatusFromErrorOr(err, http.StatusBadGateway), gin.H{"error": "Failed to read Codex search response"}) return } - reporter := helps.NewUsageReporter(ctx, "codex", selectionModel, selected) detail := helps.ParseOpenAIUsage(upstreamBody) if resp.StatusCode >= 400 { - reporter.PublishFailureWithDetail(ctx, detail, fmt.Errorf("search returned HTTP %d", resp.StatusCode)) + reporter.PublishFailureWithDetail(ctx, detail, responseFailure(upstreamBody, nil)) } else { reporter.Publish(ctx, detail) } helps.AppendAPIResponseChunk(ctx, s.cfg, upstreamBody) if selection != nil && resp.StatusCode == http.StatusUnauthorized { - s.handlers.AuthManager.ReportHomeUnauthorized(ctx, selected, "codex", selectionModel, upstreamBody) log.WithField("status", resp.StatusCode).Warnf("codex alpha search upstream request failed: %s", logging.SafeDiagnosticForLog(string(upstreamBody))) } if contentType := resp.Header.Get("Content-Type"); contentType != "" { @@ -545,6 +556,7 @@ func (s *Server) AttachWebsocketRoute(path string, handler http.Handler) { authMiddleware := AuthMiddleware(s.accessManager) conditionalAuth := func(c *gin.Context) { if !s.wsAuthEnabled.Load() { + c.Set(usageRequestAcceptedKey, true) c.Next() return } diff --git a/internal/api/server_test.go b/internal/api/server_test.go index 6eb8342cd3e..889ec31ec1a 100644 --- a/internal/api/server_test.go +++ b/internal/api/server_test.go @@ -488,7 +488,9 @@ func TestHomeCodexAlphaSearchReportsUnauthorizedBeforeEarlyReturn(t *testing.T) executor.beforeReturn = func() { test.beforeReturn(registry) } } server.handlers.AuthManager.RegisterExecutor(executor) - usageCapture := registerHomeUnauthorizedUsageCapture(t, t.Name(), testAuthID) + usageCapture := &alphaSearchUsageCapture{authID: testAuthID} + coreusage.RegisterNamedPlugin(t.Name(), usageCapture) + t.Cleanup(func() { coreusage.RegisterNamedPlugin(t.Name(), &alphaSearchUsageCapture{}) }) recorder := httptest.NewRecorder() request := httptest.NewRequest(http.MethodPost, "/v1/alpha/search", strings.NewReader(`{"model":"gpt-5-codex","query":"test"}`)) @@ -498,7 +500,12 @@ func TestHomeCodexAlphaSearchReportsUnauthorizedBeforeEarlyReturn(t *testing.T) if recorder.Code != test.wantStatus { t.Fatalf("status = %d, want %d; body=%s", recorder.Code, test.wantStatus, recorder.Body.String()) } - record := usageCapture.wait(t) + usageCapture.mu.Lock() + defer usageCapture.mu.Unlock() + if len(usageCapture.records) != 1 { + t.Fatalf("got %d Home unauthorized records, want 1", len(usageCapture.records)) + } + record := usageCapture.records[0] if record.Fail.StatusCode != http.StatusUnauthorized || record.Fail.Body != test.wantFailBody { t.Fatalf("Home unauthorized failure = status %d body %q, want status 401 body %q", record.Fail.StatusCode, record.Fail.Body, test.wantFailBody) } diff --git a/internal/api/usage_coverage.go b/internal/api/usage_coverage.go index ba2770e736e..0095e5b45a1 100644 --- a/internal/api/usage_coverage.go +++ b/internal/api/usage_coverage.go @@ -10,6 +10,8 @@ import ( "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" ) +const usageRequestAcceptedKey = "cpa.usage.requestAccepted" + // coverageMiddleware records operations whose handlers have no token reporter. // Control-plane calls and opaque WebRTC media are explicitly unmeasured; token // estimates (for example count_tokens) are never recorded as consumed tokens. @@ -29,10 +31,9 @@ func usageCoverageMiddleware() gin.HandlerFunc { if scope.Published() { return } - // Do not turn rejected credentials or unknown routes into unexpired - // journal files. Coverage is for accepted, registered API operations. - status := c.Writer.Status() - if c.FullPath() == "" || c.IsAborted() && (status == http.StatusUnauthorized || status == http.StatusForbidden) { + // Only requests admitted by route authentication may create fallback + // journal files. Earlier gates can reject requests with any status. + if c.FullPath() == "" || !c.GetBool(usageRequestAcceptedKey) { return } usage.PublishRecord(context.WithValue(ctx, "gin", c), usage.Record{Provider: "unknown", Model: "unknown", ExecutorType: "EndpointCoverage", Endpoint: c.Request.Method + " " + path, Kind: "unmeasured", Generate: usage.GenerateFlag(false), RequestedAt: started, Latency: time.Since(started), Failed: c.Writer.Status() >= http.StatusBadRequest, Fail: usage.Failure{StatusCode: c.Writer.Status()}}) diff --git a/internal/api/usage_coverage_test.go b/internal/api/usage_coverage_test.go index 0870187a382..7c6287a0870 100644 --- a/internal/api/usage_coverage_test.go +++ b/internal/api/usage_coverage_test.go @@ -8,6 +8,8 @@ import ( "testing" "github.com/gin-gonic/gin" + "github.com/router-for-me/CLIProxyAPI/v7/internal/config" + "github.com/router-for-me/CLIProxyAPI/v7/internal/home" "github.com/router-for-me/CLIProxyAPI/v7/sdk/api/handlers" "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" sdkconfig "github.com/router-for-me/CLIProxyAPI/v7/sdk/config" @@ -28,7 +30,7 @@ func TestCoverageSurvivesNormalHandlerContext(t *testing.T) { sink := &coverageReviewSink{model: t.Name()} usage.RegisterPlugin(sink) router := gin.New() - router.Use(usageCoverageMiddleware()) + router.Use(usageCoverageMiddleware(), AuthMiddleware(nil)) router.POST("/v1/"+t.Name(), func(c *gin.Context) { h := &handlers.BaseAPIHandler{Cfg: &sdkconfig.SDKConfig{}} ctx, cancel := h.GetContextWithCancel(nil, c, context.Background()) @@ -52,7 +54,7 @@ func TestCoverageSkipsRejectedAndUnmatchedRequests(t *testing.T) { matched bool want int }{ - {"unauthorized", 401, true, 0}, {"forbidden", 403, true, 0}, {"unmatched", 404, false, 0}, {"accepted_control", 200, true, 1}, + {"unauthorized", 401, true, 0}, {"forbidden", 403, true, 0}, {"unmatched", 404, false, 0}, {"gate_unavailable", 503, true, 0}, {"gate_rate_limited", 429, true, 0}, {"accepted_control", 200, true, 1}, } { t.Run(tc.name, func(t *testing.T) { sink := &coverageReviewSink{model: t.Name()} @@ -64,10 +66,8 @@ func TestCoverageSkipsRejectedAndUnmatchedRequests(t *testing.T) { router.POST(path, func(c *gin.Context) { if tc.status >= 400 { c.AbortWithStatus(tc.status) - } else { - c.Status(tc.status) } - }) + }, AuthMiddleware(nil), func(c *gin.Context) { c.Status(tc.status) }) } router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodPost, path, nil)) if len(sink.records) != tc.want { @@ -76,3 +76,34 @@ func TestCoverageSkipsRejectedAndUnmatchedRequests(t *testing.T) { }) } } + +func TestCoverageSkipsHomeHeartbeatRejection(t *testing.T) { + previous := home.Current() + home.SetCurrent(nil) + t.Cleanup(func() { home.SetCurrent(previous) }) + sink := &coverageReviewSink{model: t.Name()} + usage.RegisterPlugin(sink) + server := &Server{cfg: &config.Config{Home: config.HomeConfig{Enabled: true}}} + router := gin.New() + router.Use(usageCoverageMiddleware(), server.homeHeartbeatMiddleware(), AuthMiddleware(nil)) + path := "/v1/" + t.Name() + router.POST(path, func(c *gin.Context) { t.Fatal("handler ran while heartbeat was unavailable") }) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, path, nil)) + if recorder.Code != http.StatusServiceUnavailable || len(sink.records) != 0 { + t.Fatalf("heartbeat rejection: status=%d records=%d", recorder.Code, len(sink.records)) + } +} + +func TestCoverageRetainsAcceptedHandlerFailures(t *testing.T) { + sink := &coverageReviewSink{model: t.Name()} + usage.RegisterPlugin(sink) + router := gin.New() + router.Use(usageCoverageMiddleware(), AuthMiddleware(nil)) + path := "/v1/" + t.Name() + router.POST(path, func(c *gin.Context) { c.AbortWithStatus(http.StatusServiceUnavailable) }) + router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodPost, path, nil)) + if len(sink.records) != 1 || !sink.records[0].Failed { + t.Fatalf("accepted handler failure lost: %+v", sink.records) + } +} diff --git a/internal/runtime/executor/helps/usage_helpers.go b/internal/runtime/executor/helps/usage_helpers.go index 8541510ed02..5e8ce962412 100644 --- a/internal/runtime/executor/helps/usage_helpers.go +++ b/internal/runtime/executor/helps/usage_helpers.go @@ -685,7 +685,7 @@ func (b *StreamUsageBuffer) Observe(detail usage.Detail, ok bool) { } detail = b.billing.Apply(detail) responseServiceTier := strings.TrimSpace(detail.ResponseServiceTier) - if responseServiceTier == "" || hasNonZeroTokenUsage(detail) { + if responseServiceTier == "" || hasUsageDetail(detail) { preservedTier := b.detail.ResponseServiceTier b.detail = detail if b.detail.ResponseServiceTier == "" { @@ -1092,7 +1092,7 @@ func parseInteractionsUsageDetail(node gjson.Result) usage.Detail { } func hasUsageDetail(detail usage.Detail) bool { - return hasNonZeroTokenUsage(detail) + return detail.UsageObserved || hasNonZeroTokenUsage(detail) } func ParseInteractionsUsage(data []byte) usage.Detail { @@ -1167,7 +1167,7 @@ func ParseGeminiStreamUsage(line []byte) (usage.Detail, bool) { return usage.Detail{}, false } detail := withResponseBilling(parseGeminiFamilyUsageDetail(node), gjson.ParseBytes(payload)) - if !hasNonZeroTokenUsage(detail) { + if !detail.UsageObserved { return usage.Detail{}, false } return detail, true diff --git a/internal/runtime/executor/helps/usage_helpers_test.go b/internal/runtime/executor/helps/usage_helpers_test.go index 7283cf13073..54acbb03f4a 100644 --- a/internal/runtime/executor/helps/usage_helpers_test.go +++ b/internal/runtime/executor/helps/usage_helpers_test.go @@ -431,24 +431,22 @@ func TestParseGeminiUsageIncludesToolUsePromptTokens(t *testing.T) { } } -func TestParseGeminiStreamUsageSkipsZeroPlaceholder(t *testing.T) { +func TestParseGeminiStreamUsageFinalUsageReplacesZeroPlaceholder(t *testing.T) { lines := [][]byte{ []byte(`data: {"usageMetadata":{"promptTokenCount":0,"candidatesTokenCount":0,"thoughtsTokenCount":0,"totalTokenCount":0}}`), []byte(`data: {"usageMetadata":{"promptTokenCount":17984,"candidatesTokenCount":2668,"thoughtsTokenCount":1028,"totalTokenCount":21680}}`), } - accepted := make([]usage.Detail, 0, len(lines)) + var buffer StreamUsageBuffer for _, line := range lines { detail, ok := ParseGeminiStreamUsage(line) - if ok { - accepted = append(accepted, detail) - } + buffer.Observe(detail, ok) } - if len(accepted) != 1 { - t.Fatalf("accepted usage count = %d, want 1", len(accepted)) + detail, ok := buffer.Detail() + if !ok || !detail.UsageObserved { + t.Fatalf("final usage missing: %+v, ok=%v", detail, ok) } - detail := accepted[0] if detail.InputTokens != 17984 || detail.OutputTokens != 2668 || detail.ReasoningTokens != 1028 || detail.TotalTokens != 21680 { t.Fatalf("accepted usage detail = %+v", detail) } diff --git a/internal/runtime/executor/helps/usage_integrity_test.go b/internal/runtime/executor/helps/usage_integrity_test.go index 3f702fee85a..43b7c7d1099 100644 --- a/internal/runtime/executor/helps/usage_integrity_test.go +++ b/internal/runtime/executor/helps/usage_integrity_test.go @@ -30,3 +30,48 @@ func TestUsageIntegrityDeepSeekLegacyCache(t *testing.T) { t.Fatalf("legacy cache lost: %+v", d) } } + +func TestGeminiFamilyStreamPreservesMeasuredZero(t *testing.T) { + tests := []struct { + name string + protocol string + payload string + }{ + {"gemini", "gemini", `data: {"candidates":[{"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":0,"candidatesTokenCount":0,"totalTokenCount":0}}`}, + {"gemini snake case", "gemini", `data: {"usage_metadata":{"promptTokenCount":0,"candidatesTokenCount":0,"totalTokenCount":0}}`}, + {"interactions", "interactions", `data: {"type":"interaction.completed","interaction":{"usage":{"input_tokens":0,"output_tokens":0,"total_tokens":0}}}`}, + {"interactions Gemini fields", "interactions", `data: {"type":"interaction.completed","interaction":{"usage":{"promptTokenCount":0,"candidatesTokenCount":0,"totalTokenCount":0}}}`}, + {"antigravity", "antigravity", `data: {"response":{"usageMetadata":{"promptTokenCount":0,"candidatesTokenCount":0,"totalTokenCount":0}}}`}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var buffer StreamUsageBuffer + ObservePluginExecutorStreamUsage(tt.protocol, []byte(tt.payload), &buffer) + detail, ok := buffer.Detail() + if !ok || !detail.UsageObserved || detail.RawUsage == "" || detail.TotalTokens != 0 { + t.Fatalf("measured zero was lost: detail=%+v ok=%v", detail, ok) + } + var missing StreamUsageBuffer + ObservePluginExecutorStreamUsage(tt.protocol, []byte(`data: {"candidates":[{"finishReason":"STOP"}]}`), &missing) + if detail, ok := missing.Detail(); ok || detail.UsageObserved { + t.Fatalf("missing usage became measured: detail=%+v ok=%v", detail, ok) + } + }) + } +} + +func TestStreamUsageBufferPreservesMeasuredZeroWithTier(t *testing.T) { + var buffer StreamUsageBuffer + buffer.ObserveOpenAIStream([]byte(`data: {"service_tier":"default","usage":{"input_tokens":2,"output_tokens":3,"total_tokens":5}}`)) + buffer.ObserveOpenAIStream([]byte(`data: {"service_tier":"priority","usage":{"input_tokens":0,"output_tokens":0,"total_tokens":0}}`)) + detail, ok := buffer.Detail() + if !ok || !detail.UsageObserved || detail.RawUsage == "" || detail.TotalTokens != 0 || detail.ResponseServiceTier != "priority" { + t.Fatalf("final measured zero with tier was lost: detail=%+v ok=%v", detail, ok) + } + var zeroOnly StreamUsageBuffer + zeroOnly.ObserveOpenAIStream([]byte(`data: {"service_tier":"priority","usage":{"input_tokens":0,"output_tokens":0,"total_tokens":0}}`)) + detail, ok = zeroOnly.Detail() + if !ok || !detail.UsageObserved || detail.RawUsage == "" || detail.TotalTokens != 0 { + t.Fatalf("only measured zero with tier was lost: detail=%+v ok=%v", detail, ok) + } +} From 30c7821850386bd58fb5af889027384dd6002674 Mon Sep 17 00:00:00 2001 From: Ramapitecus Date: Sat, 12 Sep 2026 19:38:20 +0500 Subject: [PATCH 7/9] preserve streamed usage and video identity without persisting error bodies --- docs/usage-accounting.md | 4 +- internal/api/handlers/management/api_tools.go | 19 +++ .../management/api_tools_usage_test.go | 53 ++++++++ internal/redisqueue/journal.go | 10 ++ internal/redisqueue/journal_test.go | 42 +++++++ .../runtime/executor/xai_executor_media.go | 7 +- .../runtime/executor/xai_executor_test.go | 7 +- .../runtime/executor/xai_video_usage_test.go | 113 ++++++++++++++++++ 8 files changed, 247 insertions(+), 8 deletions(-) create mode 100644 internal/runtime/executor/xai_video_usage_test.go diff --git a/docs/usage-accounting.md b/docs/usage-accounting.md index ba374b2454e..a9933fbd7e2 100644 --- a/docs/usage-accounting.md +++ b/docs/usage-accounting.md @@ -17,7 +17,7 @@ Alpha Search and model-bearing management POST probes publish available usage. T ## Durable local consumer -When a config file is supplied, the local journal is stored in `usage-journal/` beside it. The usage queue plugin persists events synchronously before asynchronous plugins run. Each file is atomically renamed after flushing; POSIX also flushes the directory. File permissions are private. API keys and credential-valued sources are fingerprinted; response headers are excluded from the durable copy. Legacy wire behavior remains available for old clients. +When a config file is supplied, the local journal is stored in `usage-journal/` beside it. The usage queue plugin persists events synchronously before asynchronous plugins run. Each file is atomically renamed after flushing; POSIX also flushes the directory. File permissions are private. API keys and credential-valued sources are fingerprinted; response headers and upstream failure bodies are excluded from the durable copy. Legacy wire behavior remains available for old clients. The existing authenticated management API exposes: @@ -52,3 +52,5 @@ Coverage skips unmatched routes and locally rejected authorization, so unauthent Management Gemini/Vertex generation calls also resolve models from `/models/{model}:generateContent` and `:streamGenerateContent` paths when the request body omits `model`; explicit body models remain authoritative. Opaque event IDs are preserved in payloads and mapped to SHA-256 filenames for durable storage and ACK. Previously written journal files remain readable and acknowledgeable. Coverage fallback requires successful route admission, including legacy authentication-disabled and realtime client-secret paths. Requests blocked by Home heartbeat or other gates before admission do not create journal files; admitted handler failures remain visible. Alpha Search creates its reporter before the upstream attempt so transport/read failures and latency retain their source attribution. Gemini and Interactions streams preserve explicitly reported zero usage, including when a response tier is present. + +The durable sanitizer removes `fail.body` because provider errors can echo request content or credentials; the status and token attribution remain, and the legacy memory queue retains diagnostics. Management Gemini JSON-array streams merge elements using the same accounting buffer as SSE. xAI video responses preserve cumulative billing identity even when the provider has not yet reported usage. diff --git a/internal/api/handlers/management/api_tools.go b/internal/api/handlers/management/api_tools.go index e306d95137a..18b45785088 100644 --- a/internal/api/handlers/management/api_tools.go +++ b/internal/api/handlers/management/api_tools.go @@ -732,5 +732,24 @@ func parseManagementResponseUsage(provider, path, contentType string, payload [] detail, _ := buffer.Detail() return detail } + // Google REST streams use a JSON array unless the request selects SSE. + // Preserve element boundaries (including pretty-printed objects), keep + // earlier billing metadata, and let the final measured counters win. + if protocol == "gemini" && gjson.ValidBytes(payload) { + if root := gjson.ParseBytes(payload); root.IsArray() { + var buffer helps.StreamUsageBuffer + for _, item := range root.Array() { + if !item.IsObject() { + continue + } + frame := []byte(item.Raw) + buffer.ObserveBillingPayload(frame) + detail := helps.ParsePluginExecutorResponseUsage(protocol, frame) + buffer.Observe(detail, detail.UsageObserved) + } + detail, _ := buffer.Detail() + return detail + } + } return helps.ParsePluginExecutorResponseUsage(protocol, payload) } diff --git a/internal/api/handlers/management/api_tools_usage_test.go b/internal/api/handlers/management/api_tools_usage_test.go index e716c59f1ec..d8bd552f836 100644 --- a/internal/api/handlers/management/api_tools_usage_test.go +++ b/internal/api/handlers/management/api_tools_usage_test.go @@ -125,3 +125,56 @@ func TestManagementUsageDerivesModelFromURL(t *testing.T) { }) } } + +func TestManagementUsageParsesGeminiJSONStream(t *testing.T) { + for _, tc := range []struct { + name, provider, path, response string + tokens int64 + observed, grounding bool + }{ + {"gemini", "gemini", "/v1beta/models/gemini-2.5-flash:streamGenerateContent", `[{"usageMetadata":{"promptTokenCount":100,"totalTokenCount":100}},{"usageMetadata":{"promptTokenCount":100,"candidatesTokenCount":20,"thoughtsTokenCount":10,"totalTokenCount":130}}]`, 130, true, false}, + {"vertex", "vertex", "/v1/projects/demo/locations/global/publishers/google/models/gemini-2.5-pro:streamGenerateContent", `[ + {"candidates":[{"groundingMetadata":{"webSearchQueries":["query"]}}]}, + {"usageMetadata":{ + "promptTokenCount":100,"candidatesTokenCount":30,"totalTokenCount":130 + }} + ]`, 130, true, true}, + {"final_zero", "gemini", "/v1beta/models/gemini-2.5-flash:streamGenerateContent", `[{"usageMetadata":{"promptTokenCount":100,"totalTokenCount":100}},{"usageMetadata":{"promptTokenCount":0,"candidatesTokenCount":0,"totalTokenCount":0}}]`, 0, true, false}, + {"no_usage", "gemini", "/v1beta/models/gemini-2.5-flash:streamGenerateContent", `[{"candidates":[{"content":{"parts":[{"text":"hello"}]}}]}]`, 0, false, false}, + {"empty", "gemini", "/v1beta/models/gemini-2.5-flash:streamGenerateContent", `[]`, 0, false, false}, + } { + t.Run(tc.name, func(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + _, _ = w.Write([]byte(tc.response)) + })) + defer upstream.Close() + manager := coreauth.NewManager(nil, nil, nil) + auth := &coreauth.Auth{ID: t.Name(), Provider: tc.provider} + auth.EnsureIndex() + if _, err := manager.Register(context.Background(), auth); err != nil { + t.Fatal(err) + } + sink := &managementUsageReviewSink{authID: t.Name()} + usage.RegisterPlugin(sink) + h := &Handler{authManager: manager} + router := gin.New() + router.POST("/", h.APICall) + body, _ := json.Marshal(map[string]any{"method": "POST", "url": upstream.URL + tc.path, "auth_index": auth.Index, "data": `{"contents":[]}`}) + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(string(body))) + req.Header.Set("Content-Type", "application/json") + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, req) + if recorder.Code != http.StatusOK || len(sink.records) != 1 { + t.Fatalf("HTTP %d, usage records=%d: %s", recorder.Code, len(sink.records), recorder.Body.String()) + } + detail := sink.records[0].Detail + if detail.UsageObserved != tc.observed || detail.TotalTokens != tc.tokens { + t.Fatalf("lost array stream usage: %+v", detail) + } + if tc.grounding && !strings.Contains(detail.RawUsage, `"unpriced_server_tools":true`) { + t.Fatalf("lost earlier grounding metadata: %s", detail.RawUsage) + } + }) + } +} diff --git a/internal/redisqueue/journal.go b/internal/redisqueue/journal.go index afd1155ce14..fff87044167 100644 --- a/internal/redisqueue/journal.go +++ b/internal/redisqueue/journal.go @@ -116,6 +116,16 @@ func (j *usageJournal) append(payload []byte) error { } delete(fields, "api_key") delete(fields, "response_headers") + // Provider failures may echo prompts or credentials. Preserve the status + // for accounting, but keep response bodies only in the legacy queue. + if rawFail, exists := fields["fail"]; exists { + var failure map[string]json.RawMessage + if errDecode := json.Unmarshal(rawFail, &failure); errDecode != nil { + return errDecode + } + delete(failure, "body") + fields["fail"], _ = json.Marshal(failure) + } var errEncode error payload, errEncode = json.Marshal(fields) if errEncode != nil { diff --git a/internal/redisqueue/journal_test.go b/internal/redisqueue/journal_test.go index 41264fbfca4..0f85fd54e47 100644 --- a/internal/redisqueue/journal_test.go +++ b/internal/redisqueue/journal_test.go @@ -1,9 +1,11 @@ package redisqueue import ( + "context" "crypto/sha256" "encoding/hex" "encoding/json" + coreusage "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" "os" "path/filepath" "strings" @@ -185,3 +187,43 @@ func TestUsageJournalFingerprintsCredentialSource(t *testing.T) { t.Fatal("upstream credential persisted in source") } } + +func TestUsageJournalStripsFailureBodyWithoutChangingLegacyQueue(t *testing.T) { + withEnabledQueue(t, func() { + ConfigureUsageJournal(t.TempDir()) + defer ConfigureUsageJournal("") + body := strings.Repeat("private-prompt-and-echoed-credential", 1000) + (&usageQueuePlugin{}).HandleUsage(context.Background(), coreusage.Record{ + EventID: "failure-redaction", Provider: "codex", Model: "test", Failed: true, + Fail: coreusage.Failure{StatusCode: 401, Body: body}, + Detail: coreusage.Detail{InputTokens: 10, TotalTokens: 10, UsageObserved: true}, + }) + legacy := popSinglePayload(t) + var legacyFail failDetail + if err := json.Unmarshal(legacy["fail"], &legacyFail); err != nil { + t.Fatal(err) + } + if legacyFail.Body != body { + t.Fatal("legacy failure diagnostics were changed") + } + items, err := ReadUsageJournal(10) + if err != nil || len(items) != 1 { + t.Fatalf("journal read: count=%d err=%v", len(items), err) + } + var durable struct { + EventID string `json:"event_id"` + Failed bool `json:"failed"` + Fail map[string]json.RawMessage `json:"fail"` + Tokens tokenStats `json:"tokens"` + } + if err := json.Unmarshal(items[0], &durable); err != nil { + t.Fatal(err) + } + if _, exists := durable.Fail["body"]; exists || strings.Contains(string(items[0]), "private-prompt") { + t.Fatal("durable journal retained sensitive upstream failure body") + } + if durable.EventID != "failure-redaction" || !durable.Failed || string(durable.Fail["status_code"]) != "401" || durable.Tokens.TotalTokens != 10 { + t.Fatalf("failure attribution or usage was lost: %+v", durable) + } + }) +} diff --git a/internal/runtime/executor/xai_executor_media.go b/internal/runtime/executor/xai_executor_media.go index a89562264d3..d6108783752 100644 --- a/internal/runtime/executor/xai_executor_media.go +++ b/internal/runtime/executor/xai_executor_media.go @@ -154,10 +154,7 @@ func (e *XAIExecutor) executeVideos(ctx context.Context, auth *cliproxyauth.Auth detail.BillingID = "xai-video/" + billingID detail.CostScope = "operation" } - if detail.UsageObserved { - reporter.Publish(ctx, detail) - } else { - reporter.EnsurePublished(ctx) - } + // Creation and polling responses can identify the operation without reporting usage. + reporter.Publish(ctx, detail) return cliproxyexecutor.Response{Payload: data, Headers: httpResp.Header.Clone()}, nil } diff --git a/internal/runtime/executor/xai_executor_test.go b/internal/runtime/executor/xai_executor_test.go index 9199184a030..9d0889bf893 100644 --- a/internal/runtime/executor/xai_executor_test.go +++ b/internal/runtime/executor/xai_executor_test.go @@ -3657,8 +3657,11 @@ func TestXAIExecutorExecuteVideosCreate(t *testing.T) { if record.Failed { t.Fatalf("failed = true, want false; failure=%+v", record.Fail) } - if record.Detail != (usage.Detail{}) { - t.Fatalf("detail = %+v, want zero token usage", record.Detail) + if record.Detail.BillingID != "xai-video/vid_123" || record.Detail.CostScope != "operation" { + t.Fatalf("detail = %+v, want video operation identity", record.Detail) + } + if record.Detail.UsageObserved || record.Detail.TotalTokens != 0 || record.Detail.InputTokens != 0 || record.Detail.OutputTokens != 0 || record.Detail.CostUSD != nil { + t.Fatalf("detail = %+v, want unreported token usage and cost", record.Detail) } if record.TTFT <= 0 { t.Fatalf("ttft = %v, want positive duration", record.TTFT) diff --git a/internal/runtime/executor/xai_video_usage_test.go b/internal/runtime/executor/xai_video_usage_test.go new file mode 100644 index 00000000000..a1399ecac5f --- /dev/null +++ b/internal/runtime/executor/xai_video_usage_test.go @@ -0,0 +1,113 @@ +package executor + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "github.com/router-for-me/CLIProxyAPI/v7/internal/config" + cliproxyauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + cliproxyexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" + sdktranslator "github.com/router-for-me/CLIProxyAPI/v7/sdk/translator" +) + +func TestXAIExecutorVideoBillingWithoutTokens(t *testing.T) { + for _, tc := range []struct { + name, path, payload, response, billingID, cost string + observed bool + }{ + {name: "create", path: "/v1/videos/generations", payload: `{"model":"grok-imagine-video"}`, response: `{"request_id":"vid_123"}`, billingID: "vid_123"}, + {name: "edit", path: "/v1/videos/edits", payload: `{"model":"grok-imagine-video"}`, response: `{"request_id":"vid_123"}`, billingID: "vid_123"}, + {name: "extend", path: "/v1/videos/extensions", payload: `{"model":"grok-imagine-video"}`, response: `{"request_id":"vid_123"}`, billingID: "vid_123"}, + {name: "poll pending", payload: `{"request_id":"vid_123"}`, response: `{"status":"pending"}`, billingID: "vid_123"}, + {name: "poll done", payload: `{"request_id":"vid_123"}`, response: `{"status":"done","video":{"url":"https://example.com/video.mp4"}}`, billingID: "vid_123"}, + {name: "response identity wins", payload: `{"request_id":"vid_123"}`, response: `{"request_id":"vid_456","status":"done"}`, billingID: "vid_456"}, + {name: "create cost only", path: "/v1/videos/generations", payload: `{"model":"grok-imagine-video"}`, response: `{"request_id":"vid_123","usage":{"cost_in_usd_ticks":250000}}`, billingID: "vid_123", cost: "0.0000250000", observed: true}, + {name: "poll cost only", payload: `{"request_id":"vid_123"}`, response: `{"status":"done","usage":{"cost_in_usd_ticks":250000}}`, billingID: "vid_123", cost: "0.0000250000", observed: true}, + {name: "reported zero cost", payload: `{"request_id":"vid_123"}`, response: `{"status":"done","usage":{"cost_in_usd_ticks":0}}`, billingID: "vid_123", cost: "0.0000000000", observed: true}, + {name: "unknown identity and cost", payload: `{"model":"grok-imagine-video"}`, response: `{}`}, + } { + t.Run(tc.name, func(t *testing.T) { + capture := &websocketUsageCapture{authID: t.Name()} + usage.RegisterNamedPlugin(t.Name(), capture) + defer usage.RegisterNamedPlugin(t.Name(), &websocketUsageCapture{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprint(w, tc.response) + })) + defer server.Close() + auth := &cliproxyauth.Auth{ID: t.Name(), Provider: "xai", Attributes: map[string]string{"base_url": server.URL}, Metadata: map[string]any{"access_token": "test"}} + _, err := NewXAIExecutor(&config.Config{}).Execute(context.Background(), auth, cliproxyexecutor.Request{ + Model: "grok-imagine-video", Payload: []byte(tc.payload), + }, cliproxyexecutor.Options{SourceFormat: sdktranslator.FromString("openai-video"), Metadata: map[string]any{cliproxyexecutor.RequestPathMetadataKey: tc.path}}) + if err != nil { + t.Fatal(err) + } + capture.mu.Lock() + defer capture.mu.Unlock() + if len(capture.records) != 1 { + t.Fatalf("got %d events, want 1", len(capture.records)) + } + record := capture.records[0] + detail := record.Detail + wantID, wantScope := "", "" + if tc.billingID != "" { + wantID, wantScope = "xai-video/"+tc.billingID, "operation" + } + if record.Failed || detail.BillingID != wantID || detail.CostScope != wantScope { + t.Errorf("failed=%v, billing ID=%q, scope=%q; want false, %q, %q", record.Failed, detail.BillingID, detail.CostScope, wantID, wantScope) + } + if detail.UsageObserved != tc.observed || detail.InputTokens != 0 || detail.OutputTokens != 0 || detail.TotalTokens != 0 { + t.Errorf("unexpected token usage: %+v", detail) + } + if tc.cost == "" { + if detail.CostUSD != nil { + t.Errorf("missing cost became reported cost %q", *detail.CostUSD) + } + } else if detail.CostUSD == nil || *detail.CostUSD != tc.cost { + t.Errorf("cost = %v, want %q", detail.CostUSD, tc.cost) + } + }) + } +} + +func TestXAIExecutorVideoPollsShareBillingIdentityAndKeepDistinctEvents(t *testing.T) { + capture := &websocketUsageCapture{authID: t.Name()} + usage.RegisterNamedPlugin(t.Name(), capture) + defer usage.RegisterNamedPlugin(t.Name(), &websocketUsageCapture{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + if r.Method == http.MethodPost { + _, _ = fmt.Fprint(w, `{"request_id":"vid_123"}`) + return + } + _, _ = fmt.Fprint(w, `{"status":"done","usage":{"cost_in_usd_ticks":250000}}`) + })) + defer server.Close() + auth := &cliproxyauth.Auth{ID: t.Name(), Provider: "xai", Attributes: map[string]string{"base_url": server.URL}, Metadata: map[string]any{"access_token": "test"}} + executor := NewXAIExecutor(&config.Config{}) + for _, payload := range []string{`{"model":"grok-imagine-video"}`, `{"request_id":"vid_123"}`, `{"request_id":"vid_123"}`} { + _, err := executor.Execute(context.Background(), auth, cliproxyexecutor.Request{Model: "grok-imagine-video", Payload: []byte(payload)}, cliproxyexecutor.Options{SourceFormat: sdktranslator.FromString("openai-video")}) + if err != nil { + t.Fatal(err) + } + } + capture.mu.Lock() + defer capture.mu.Unlock() + if len(capture.records) != 3 { + t.Fatalf("got %d events, want 3 HTTP attempts", len(capture.records)) + } + ids := make(map[string]bool) + for _, record := range capture.records { + if record.EventID == "" || ids[record.EventID] { + t.Errorf("missing or duplicate event ID %q", record.EventID) + } + ids[record.EventID] = true + if record.Detail.BillingID != "xai-video/vid_123" || record.Detail.CostScope != "operation" { + t.Errorf("lost shared operation identity: %+v", record.Detail) + } + } +} From 1386ec7d23e1b84108f29dbf875f1bc7d7c8ad08 Mon Sep 17 00:00:00 2001 From: Ramapitecus Date: Sat, 12 Sep 2026 19:56:11 +0500 Subject: [PATCH 8/9] retain billing-only streams and recover mixed subscriber overflow --- docs/usage-accounting.md | 2 + .../management/api_tools_usage_test.go | 1 + internal/redisqueue/queue.go | 8 +- internal/redisqueue/usage_integrity_test.go | 45 ++++++ .../runtime/executor/aistudio_executor.go | 15 +- .../executor/antigravity_executor_execute.go | 12 +- .../executor/antigravity_executor_stream.go | 12 +- internal/runtime/executor/gemini_executor.go | 10 +- .../helps/usage_billing_integrity_test.go | 68 ++++++++- .../runtime/executor/helps/usage_helpers.go | 16 +- .../executor/native_billing_only_test.go | 142 ++++++++++++++++++ 11 files changed, 305 insertions(+), 26 deletions(-) create mode 100644 internal/runtime/executor/native_billing_only_test.go diff --git a/docs/usage-accounting.md b/docs/usage-accounting.md index a9933fbd7e2..afe2e55c2c8 100644 --- a/docs/usage-accounting.md +++ b/docs/usage-accounting.md @@ -54,3 +54,5 @@ Management Gemini/Vertex generation calls also resolve models from `/models/{mod Coverage fallback requires successful route admission, including legacy authentication-disabled and realtime client-secret paths. Requests blocked by Home heartbeat or other gates before admission do not create journal files; admitted handler failures remain visible. Alpha Search creates its reporter before the upstream attempt so transport/read failures and latency retain their source attribution. Gemini and Interactions streams preserve explicitly reported zero usage, including when a response tier is present. The durable sanitizer removes `fail.body` because provider errors can echo request content or credentials; the status and token attribution remain, and the legacy memory queue retains diagnostics. Management Gemini JSON-array streams merge elements using the same accounting buffer as SSE. xAI video responses preserve cumulative billing identity even when the provider has not yet reported usage. + +Legacy fan-out falls back to the queue if any subscriber overflows, even when other subscribers receive the event. Consumers must deduplicate repeated delivery by event ID. Streaming billing dimensions are retained independently of token measurements; a billing-only event keeps `usage_observed:false` instead of inventing measured token counters. diff --git a/internal/api/handlers/management/api_tools_usage_test.go b/internal/api/handlers/management/api_tools_usage_test.go index d8bd552f836..712bec413aa 100644 --- a/internal/api/handlers/management/api_tools_usage_test.go +++ b/internal/api/handlers/management/api_tools_usage_test.go @@ -141,6 +141,7 @@ func TestManagementUsageParsesGeminiJSONStream(t *testing.T) { ]`, 130, true, true}, {"final_zero", "gemini", "/v1beta/models/gemini-2.5-flash:streamGenerateContent", `[{"usageMetadata":{"promptTokenCount":100,"totalTokenCount":100}},{"usageMetadata":{"promptTokenCount":0,"candidatesTokenCount":0,"totalTokenCount":0}}]`, 0, true, false}, {"no_usage", "gemini", "/v1beta/models/gemini-2.5-flash:streamGenerateContent", `[{"candidates":[{"content":{"parts":[{"text":"hello"}]}}]}]`, 0, false, false}, + {"grounding_without_tokens", "gemini", "/v1beta/models/gemini-2.5-flash:streamGenerateContent", `[{"candidates":[{"groundingMetadata":{"webSearchQueries":["query"]}}]}]`, 0, false, true}, {"empty", "gemini", "/v1beta/models/gemini-2.5-flash:streamGenerateContent", `[]`, 0, false, false}, } { t.Run(tc.name, func(t *testing.T) { diff --git a/internal/redisqueue/queue.go b/internal/redisqueue/queue.go index 8c5c6269e52..b594dc28a7c 100644 --- a/internal/redisqueue/queue.go +++ b/internal/redisqueue/queue.go @@ -146,19 +146,21 @@ func (q *queue) publishToSubscribers(payload []byte) bool { return false } - delivered := false + allDelivered := true for id, subscriber := range q.subscribers { cloned := append([]byte(nil), payload...) select { case subscriber <- cloned: - delivered = true default: + allDelivered = false delete(q.subscribers, id) close(subscriber) } } - return delivered + // Any disconnected subscriber needs the event in the fallback queue, + // even when other subscribers accepted their live copy. + return allDelivered } func (q *queue) subscribe(buffer int, initialPayload []byte) (<-chan []byte, func()) { diff --git a/internal/redisqueue/usage_integrity_test.go b/internal/redisqueue/usage_integrity_test.go index 138222eb995..c10d5557e5e 100644 --- a/internal/redisqueue/usage_integrity_test.go +++ b/internal/redisqueue/usage_integrity_test.go @@ -16,3 +16,48 @@ func TestUsageIntegrityOverflowFallsBack(t *testing.T) { t.Fatalf("lost overflow event: got %d", len(got)) } } + +func TestUsageIntegrityMixedSubscriberOverflowFallsBack(t *testing.T) { + withEnabledQueue(t, func() { + slow, cancelSlow := SubscribeUsage() + defer cancelSlow() + fast, cancelFast := SubscribeUsage() + defer cancelFast() + requireUsageSubscriberPayload(t, fast, usageSupportRefreshPayload) + // Keep one subscriber writable while filling the other's buffer. + for i := 0; i < usageSubscriberBuffer-1; i++ { + Enqueue([]byte(`{"event_id":"fill"}`)) + requireUsageSubscriberPayload(t, fast, `{"event_id":"fill"}`) + } + const overflow = `{"event_id":"overflow"}` + Enqueue([]byte(overflow)) + requireUsageSubscriberPayload(t, fast, overflow) + got := PopOldest(10) + if len(got) != 1 || string(got[0]) != overflow { + t.Fatalf("disconnected subscriber cannot recover overflow: %q", got) + } + for i := 0; i < usageSubscriberBuffer; i++ { + select { + case _, ok := <-slow: + if !ok { + t.Fatal("slow subscriber lost buffered records") + } + default: + t.Fatal("slow subscriber buffer unexpectedly empty") + } + } + select { + case _, ok := <-slow: + if ok { + t.Fatal("slow subscriber remained open") + } + default: + t.Fatal("slow subscriber was not disconnected") + } + Enqueue([]byte(`{"event_id":"next"}`)) + requireUsageSubscriberPayload(t, fast, `{"event_id":"next"}`) + if got := PopOldest(10); len(got) != 0 { + t.Fatalf("healthy subscriber delivery was unnecessarily queued: %q", got) + } + }) +} diff --git a/internal/runtime/executor/aistudio_executor.go b/internal/runtime/executor/aistudio_executor.go index b84a9f68bcd..8cfbe579944 100644 --- a/internal/runtime/executor/aistudio_executor.go +++ b/internal/runtime/executor/aistudio_executor.go @@ -303,7 +303,8 @@ func (e *AIStudioExecutor) ExecuteStream(ctx context.Context, auth *cliproxyauth out := make(chan cliproxyexecutor.StreamChunk) go func(first wsrelay.StreamEvent) { defer close(out) - defer reporter.EnsurePublished(ctx) + var usageBuffer helps.StreamUsageBuffer + defer usageBuffer.EnsurePublished(ctx, reporter) responseFormat := cliproxyexecutor.ResponseFormatOrSource(opts) originalRequest := opts.OriginalRequest if len(originalRequest) == 0 { @@ -313,9 +314,10 @@ func (e *AIStudioExecutor) ExecuteStream(ctx context.Context, auth *cliproxyauth var param any metadataLogged := false processEvent := func(event wsrelay.StreamEvent) bool { + helps.IterateStreamLines(event.Payload, usageBuffer.ObserveBillingPayload) if event.Err != nil { helps.RecordAPIResponseError(ctx, e.cfg, event.Err) - reporter.PublishFailure(ctx, event.Err) + usageBuffer.PublishFailure(ctx, reporter, event.Err) select { case out <- cliproxyexecutor.StreamChunk{Err: fmt.Errorf("wsrelay: %v", event.Err)}: case <-ctx.Done(): @@ -335,7 +337,8 @@ func (e *AIStudioExecutor) ExecuteStream(ctx context.Context, auth *cliproxyauth helps.AppendAPIResponseChunk(ctx, e.cfg, event.Payload) filtered := helps.FilterSSEUsageMetadata(event.Payload) if detail, ok := helps.ParseGeminiStreamUsage(filtered); ok { - reporter.Publish(ctx, detail) + usageBuffer.Observe(detail, true) + usageBuffer.Publish(ctx, reporter) } lines := helps.TranslateStreamWithClaudeInputTokens(ctx, body.toFormat, responseFormat, req.Model, opts.OriginalRequest, translatedReq, filtered, ¶m, claudeInputTokens) for i := range lines { @@ -367,11 +370,13 @@ func (e *AIStudioExecutor) ExecuteStream(ctx context.Context, auth *cliproxyauth return false } } - reporter.Publish(ctx, helps.ParseGeminiUsage(event.Payload)) + detail := helps.ParseGeminiUsage(event.Payload) + usageBuffer.Observe(detail, detail.UsageObserved) + usageBuffer.Publish(ctx, reporter) return false case wsrelay.MessageTypeError: helps.RecordAPIResponseError(ctx, e.cfg, event.Err) - reporter.PublishFailure(ctx, event.Err) + usageBuffer.PublishFailure(ctx, reporter, event.Err) select { case out <- cliproxyexecutor.StreamChunk{Err: fmt.Errorf("wsrelay: %v", event.Err)}: case <-ctx.Done(): diff --git a/internal/runtime/executor/antigravity_executor_execute.go b/internal/runtime/executor/antigravity_executor_execute.go index 5585ed56b21..098815ff147 100644 --- a/internal/runtime/executor/antigravity_executor_execute.go +++ b/internal/runtime/executor/antigravity_executor_execute.go @@ -406,10 +406,11 @@ func (e *AntigravityExecutor) executeClaudeNonStream(ctx context.Context, auth * }() scanner := bufio.NewScanner(resp.Body) scanner.Buffer(nil, streamScannerBuffer) - var billing helps.UsageBillingMetadata + var usageBuffer helps.StreamUsageBuffer + defer usageBuffer.EnsurePublished(ctx, reporter) for scanner.Scan() { line := scanner.Bytes() - billing.ObservePayload(line) + usageBuffer.ObserveBillingPayload(line) helps.AppendAPIResponseChunk(ctx, e.cfg, line) if replayAccumulator != nil { replayAccumulator.ObserveSSELine(line) @@ -425,20 +426,21 @@ func (e *AntigravityExecutor) executeClaudeNonStream(ctx context.Context, auth * } if detail, ok := helps.ParseAntigravityStreamUsage(payload); ok { - reporter.Publish(ctx, billing.Apply(detail)) + usageBuffer.Observe(detail, true) + usageBuffer.Publish(ctx, reporter) } out <- cliproxyexecutor.StreamChunk{Payload: payload} } if errScan := scanner.Err(); errScan != nil { helps.RecordAPIResponseError(ctx, e.cfg, errScan) - reporter.PublishFailure(ctx, errScan) + usageBuffer.PublishFailure(ctx, reporter, errScan) out <- cliproxyexecutor.StreamChunk{Err: errScan} } else { if replayAccumulator != nil { replayAccumulator.Commit(ctx) } - reporter.EnsurePublished(ctx) + usageBuffer.EnsurePublished(ctx, reporter) } }(httpResp) diff --git a/internal/runtime/executor/antigravity_executor_stream.go b/internal/runtime/executor/antigravity_executor_stream.go index d7af31e38f5..5e7ed55ccb3 100644 --- a/internal/runtime/executor/antigravity_executor_stream.go +++ b/internal/runtime/executor/antigravity_executor_stream.go @@ -201,12 +201,13 @@ func (e *AntigravityExecutor) ExecuteStream(ctx context.Context, auth *cliproxya }() scanner := bufio.NewScanner(resp.Body) scanner.Buffer(nil, streamScannerBuffer) - var billing helps.UsageBillingMetadata + var usageBuffer helps.StreamUsageBuffer + defer usageBuffer.EnsurePublished(ctx, reporter) claudeInputTokens := helps.NewClaudeInputTokenState(from, to, responseFormat, originalPayload) var param any for scanner.Scan() { line := scanner.Bytes() - billing.ObservePayload(line) + usageBuffer.ObserveBillingPayload(line) helps.AppendAPIResponseChunk(ctx, e.cfg, line) if replayAccumulator != nil { replayAccumulator.ObserveSSELine(line) @@ -222,7 +223,8 @@ func (e *AntigravityExecutor) ExecuteStream(ctx context.Context, auth *cliproxya } if detail, ok := helps.ParseAntigravityStreamUsage(payload); ok { - reporter.Publish(ctx, billing.Apply(detail)) + usageBuffer.Observe(detail, true) + usageBuffer.Publish(ctx, reporter) } payload = e.resolveWebSearchGroundingURLs(ctx, auth, from, originalPayload, translated, payload) @@ -237,7 +239,7 @@ func (e *AntigravityExecutor) ExecuteStream(ctx context.Context, auth *cliproxya } if errScan := scanner.Err(); errScan != nil { helps.RecordAPIResponseError(ctx, e.cfg, errScan) - reporter.PublishFailure(ctx, errScan) + usageBuffer.PublishFailure(ctx, reporter, errScan) select { case out <- cliproxyexecutor.StreamChunk{Err: errScan}: case <-ctx.Done(): @@ -257,7 +259,7 @@ func (e *AntigravityExecutor) ExecuteStream(ctx context.Context, auth *cliproxya if replayAccumulator != nil { replayAccumulator.Commit(ctx) } - reporter.EnsurePublished(ctx) + usageBuffer.EnsurePublished(ctx, reporter) } }(httpResp) return &cliproxyexecutor.StreamResult{Headers: httpResp.Header.Clone(), Chunks: out}, nil diff --git a/internal/runtime/executor/gemini_executor.go b/internal/runtime/executor/gemini_executor.go index 71dcb2e54eb..b8c01f89611 100644 --- a/internal/runtime/executor/gemini_executor.go +++ b/internal/runtime/executor/gemini_executor.go @@ -356,12 +356,13 @@ func (e *GeminiExecutor) ExecuteStream(ctx context.Context, auth *cliproxyauth.A }() scanner := bufio.NewScanner(httpResp.Body) scanner.Buffer(nil, streamScannerBuffer) - var billing helps.UsageBillingMetadata + var usageBuffer helps.StreamUsageBuffer + defer usageBuffer.EnsurePublished(ctx, reporter) claudeInputTokens := helps.NewClaudeInputTokenState(from, to, responseFormat, originalPayload) var param any for scanner.Scan() { line := scanner.Bytes() - billing.ObservePayload(line) + usageBuffer.ObserveBillingPayload(line) helps.AppendAPIResponseChunk(ctx, e.cfg, line) filtered := helps.FilterSSEUsageMetadata(line) payload := helps.JSONPayload(filtered) @@ -369,7 +370,8 @@ func (e *GeminiExecutor) ExecuteStream(ctx context.Context, auth *cliproxyauth.A continue } if detail, ok := helps.ParseGeminiStreamUsage(payload); ok { - reporter.Publish(ctx, billing.Apply(detail)) + usageBuffer.Observe(detail, true) + usageBuffer.Publish(ctx, reporter) } lines := helps.TranslateStreamWithClaudeInputTokens(ctx, to, responseFormat, req.Model, opts.OriginalRequest, body, bytes.Clone(payload), ¶m, claudeInputTokens) for i := range lines { @@ -390,7 +392,7 @@ func (e *GeminiExecutor) ExecuteStream(ctx context.Context, auth *cliproxyauth.A } if errScan := scanner.Err(); errScan != nil { helps.RecordAPIResponseError(ctx, e.cfg, errScan) - reporter.PublishFailure(ctx, errScan) + usageBuffer.PublishFailure(ctx, reporter, errScan) select { case out <- cliproxyexecutor.StreamChunk{Err: errScan}: case <-ctx.Done(): diff --git a/internal/runtime/executor/helps/usage_billing_integrity_test.go b/internal/runtime/executor/helps/usage_billing_integrity_test.go index c36dea86ee5..2597c762aef 100644 --- a/internal/runtime/executor/helps/usage_billing_integrity_test.go +++ b/internal/runtime/executor/helps/usage_billing_integrity_test.go @@ -1,8 +1,11 @@ package helps import ( + "context" + "errors" "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" "github.com/tidwall/gjson" + "sync" "testing" ) @@ -63,8 +66,8 @@ func TestBillingMetadataSurvivesNativeGeminiFiltering(t *testing.T) { func TestOpenAIStreamBillingOnlyFramesRetainFinalCounters(t *testing.T) { var buffer StreamUsageBuffer buffer.ObserveOpenAIStream([]byte(`data: {"tool_usage":{"web_search":1}}`)) - if _, ok := buffer.Detail(); ok { - t.Fatal("billing-only frame manufactured observed token usage") + if detail, ok := buffer.Detail(); !ok || detail.UsageObserved || detail.TotalTokens != 0 { + t.Fatalf("billing-only frame must remain unmeasured: %+v, ok=%v", detail, ok) } buffer.ObserveOpenAIStream([]byte(`data: {"usage":{"prompt_tokens":100,"completion_tokens":10,"total_tokens":110}}`)) buffer.ObserveOpenAIStream([]byte(`data: {"tool_usage":{"file_search":1}}`)) @@ -73,3 +76,64 @@ func TestOpenAIStreamBillingOnlyFramesRetainFinalCounters(t *testing.T) { t.Fatalf("metadata-only frames corrupted accounting: %+v", detail) } } + +type billingOnlyCapture struct { + mu sync.Mutex + model string + records []usage.Record +} + +func (*billingOnlyCapture) Synchronous() bool { return true } +func (s *billingOnlyCapture) HandleUsage(_ context.Context, r usage.Record) { + if r.Model == s.model { + s.mu.Lock() + defer s.mu.Unlock() + s.records = append(s.records, r) + } +} +func TestStreamBillingWithoutTokenUsageIsPublished(t *testing.T) { + for _, protocol := range []string{"gemini", "antigravity", "openai"} { + for _, failed := range []bool{false, true} { + name := protocol + "/success" + if failed { + name = protocol + "/failure" + } + t.Run(name, func(t *testing.T) { + payload := `{"candidates":[{"groundingMetadata":{"webSearchQueries":["weather"]}}],"tool_usage":{"web_search":1}}` + if protocol == "antigravity" { + payload = `{"response":` + payload + `}` + } + var buffer StreamUsageBuffer + ObservePluginExecutorStreamUsage(protocol, []byte("data: "+payload), &buffer) + detail, ok := buffer.Detail() + if !ok || detail.UsageObserved || detail.TotalTokens != 0 || !gjson.Get(detail.RawUsage, "unpriced_server_tools").Bool() { + t.Fatalf("billing-only detail was lost or marked measured: %+v, ok=%v", detail, ok) + } + capture := &billingOnlyCapture{model: t.Name()} + usage.RegisterNamedPlugin(t.Name(), capture) + defer usage.RegisterNamedPlugin(t.Name(), &billingOnlyCapture{}) + reporter := NewUsageReporter(context.Background(), protocol, t.Name(), nil) + if failed { + buffer.PublishFailure(context.Background(), reporter, errors.New("truncated stream")) + } else { + buffer.Publish(context.Background(), reporter) + } + reporter.EnsurePublished(context.Background()) + capture.mu.Lock() + defer capture.mu.Unlock() + if len(capture.records) != 1 { + t.Fatalf("records=%d", len(capture.records)) + } + record := capture.records[0] + if record.Failed != failed || record.Detail.UsageObserved || !gjson.Get(record.Detail.RawUsage, "tool_usage.web_search").Exists() { + t.Fatalf("published billing lost: %+v", record) + } + }) + } + } + var empty StreamUsageBuffer + empty.ObserveBillingPayload([]byte(`data: {"candidates":[{"content":{"parts":[{"text":"hello"}]}}]}`)) + if _, ok := empty.Detail(); ok { + t.Fatal("unrelated payload manufactured billing metadata") + } +} diff --git a/internal/runtime/executor/helps/usage_helpers.go b/internal/runtime/executor/helps/usage_helpers.go index 5e8ce962412..d974f46ec9c 100644 --- a/internal/runtime/executor/helps/usage_helpers.go +++ b/internal/runtime/executor/helps/usage_helpers.go @@ -704,9 +704,13 @@ func (b *StreamUsageBuffer) ObserveBillingPayload(payload []byte) { return } b.billing.ObservePayload(payload) - if b.ok { - b.detail = b.billing.Apply(b.detail) + if len(b.billing.fields) == 0 { + return } + b.detail = b.billing.Apply(b.detail) + // Billing metadata is publishable even when token measurements are absent. + // UsageObserved remains the separate indicator for measured token counters. + b.ok = true } // ObserveOpenAIStream records response-tier state and the latest usage from an @@ -765,6 +769,14 @@ func (b *StreamUsageBuffer) Publish(ctx context.Context, reporter *UsageReporter return true } +// EnsurePublished retains any billing metadata before falling back to an +// unmeasured request record when the stream supplies no accounting details. +func (b *StreamUsageBuffer) EnsurePublished(ctx context.Context, reporter *UsageReporter) { + if reporter != nil && !b.Publish(ctx, reporter) { + reporter.EnsurePublished(ctx) + } +} + // PublishFailure emits the latest observed usage detail together with failure details. func (b *StreamUsageBuffer) PublishFailure(ctx context.Context, reporter *UsageReporter, errs ...error) bool { if b == nil || reporter == nil { diff --git a/internal/runtime/executor/native_billing_only_test.go b/internal/runtime/executor/native_billing_only_test.go new file mode 100644 index 00000000000..93591549adf --- /dev/null +++ b/internal/runtime/executor/native_billing_only_test.go @@ -0,0 +1,142 @@ +package executor + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gorilla/websocket" + "github.com/router-for-me/CLIProxyAPI/v7/internal/config" + "github.com/router-for-me/CLIProxyAPI/v7/internal/wsrelay" + cliproxyauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" + cliproxyexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" + sdktranslator "github.com/router-for-me/CLIProxyAPI/v7/sdk/translator" + "github.com/tidwall/gjson" +) + +const nativeBillingOnlyPayload = `{"candidates":[{"content":{"role":"model","parts":[{"text":"weather"}]},"groundingMetadata":{"webSearchQueries":["weather"]},"finishReason":"STOP"}]}` + +func TestNativeStreamPublishesBillingWithoutTokens(t *testing.T) { + for _, provider := range []string{"gemini", "antigravity", "antigravity_nonstream"} { + for _, failed := range []bool{false, true} { + name := provider + "/success" + if failed { + name = provider + "/truncated" + } + t.Run(name, func(t *testing.T) { + capture := &websocketUsageCapture{authID: t.Name()} + usage.RegisterNamedPlugin(t.Name(), capture) + defer usage.RegisterNamedPlugin(t.Name(), &websocketUsageCapture{}) + payload := nativeBillingOnlyPayload + if provider != "gemini" { + payload = `{"response":` + payload + `}` + } + payload = "data: " + payload + "\n\n" + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + if failed { + w.Header().Set("Content-Length", fmt.Sprint(len(payload)+100)) + } + _, _ = w.Write([]byte(payload)) + })) + defer server.Close() + auth := &cliproxyauth.Auth{ID: t.Name(), Provider: provider, Attributes: map[string]string{"api_key": "test", "base_url": server.URL}, Metadata: map[string]any{"access_token": "test", "expired": time.Now().Add(time.Hour).Format(time.RFC3339), "project_id": "test"}} + req := cliproxyexecutor.Request{Model: "gemini-2.5-flash", Payload: []byte(`{"contents":[{"role":"user","parts":[{"text":"weather"}]}]}`)} + opts := cliproxyexecutor.Options{SourceFormat: sdktranslator.FormatGemini} + var executionErr error + if provider == "antigravity_nonstream" { + _, executionErr = NewAntigravityExecutor(&config.Config{}).executeClaudeNonStream(context.Background(), auth, req, opts) + } else { + var result *cliproxyexecutor.StreamResult + if provider == "gemini" { + result, executionErr = NewGeminiExecutor(&config.Config{}).ExecuteStream(context.Background(), auth, req, opts) + } else { + result, executionErr = NewAntigravityExecutor(&config.Config{}).ExecuteStream(context.Background(), auth, req, opts) + } + if executionErr == nil { + for chunk := range result.Chunks { + if chunk.Err != nil { + executionErr = chunk.Err + } + } + } + } + if (executionErr != nil) != failed { + t.Fatalf("execution error=%v, want failure=%v", executionErr, failed) + } + capture.mu.Lock() + defer capture.mu.Unlock() + if len(capture.records) != 1 { + t.Fatalf("records=%d", len(capture.records)) + } + record := capture.records[0] + if record.Failed != failed || record.Detail.UsageObserved || !gjson.Get(record.Detail.RawUsage, "unpriced_server_tools").Bool() { + t.Fatalf("native billing-only record lost: %+v", record) + } + }) + } + } +} + +func TestAIStudioStreamPublishesBillingWithoutTokens(t *testing.T) { + authID := t.Name() + connected := make(chan struct{}) + relay := wsrelay.NewManager(wsrelay.Options{ProviderFactory: func(*http.Request) (string, error) { return authID, nil }, OnConnected: func(string) { close(connected) }}) + server := httptest.NewServer(relay.Handler()) + defer server.Close() + defer func() { _ = relay.Stop(context.Background()) }() + conn, _, errDial := websocket.DefaultDialer.Dial("ws"+strings.TrimPrefix(server.URL, "http")+relay.Path(), nil) + if errDial != nil { + t.Fatal(errDial) + } + defer func() { _ = conn.Close() }() + <-connected + clientDone := make(chan error, 1) + go func() { + var msg wsrelay.Message + if errRead := conn.ReadJSON(&msg); errRead != nil { + clientDone <- errRead + return + } + for _, event := range []wsrelay.Message{ + {ID: msg.ID, Type: wsrelay.MessageTypeStreamStart, Payload: map[string]any{"status": float64(http.StatusOK), "headers": map[string]any{"Content-Type": "text/event-stream"}}}, + {ID: msg.ID, Type: wsrelay.MessageTypeStreamChunk, Payload: map[string]any{"data": "data: " + nativeBillingOnlyPayload + "\n\n"}}, + {ID: msg.ID, Type: wsrelay.MessageTypeStreamEnd}, + } { + if errWrite := conn.WriteJSON(event); errWrite != nil { + clientDone <- errWrite + return + } + } + clientDone <- nil + }() + capture := &websocketUsageCapture{authID: authID} + usage.RegisterNamedPlugin(t.Name(), capture) + defer usage.RegisterNamedPlugin(t.Name(), &websocketUsageCapture{}) + result, errStream := NewAIStudioExecutor(&config.Config{}, "aistudio", relay).ExecuteStream(context.Background(), &cliproxyauth.Auth{ID: authID, Provider: "aistudio"}, cliproxyexecutor.Request{Model: "gemini-2.5-flash", Payload: []byte(`{"contents":[{"role":"user","parts":[{"text":"weather"}]}]}`)}, cliproxyexecutor.Options{SourceFormat: sdktranslator.FormatGemini}) + if errStream != nil { + t.Fatal(errStream) + } + for chunk := range result.Chunks { + if chunk.Err != nil { + t.Fatal(chunk.Err) + } + } + if errClient := <-clientDone; errClient != nil { + t.Fatal(errClient) + } + capture.mu.Lock() + defer capture.mu.Unlock() + if len(capture.records) != 1 { + t.Fatalf("records=%d", len(capture.records)) + } + detail := capture.records[0].Detail + if detail.UsageObserved || !gjson.Get(detail.RawUsage, "unpriced_server_tools").Bool() { + t.Fatalf("relay lost billing-only metadata: %+v", detail) + } +} From fa7946482c43d0bfed1a9995b48ed5b7027d58e3 Mon Sep 17 00:00:00 2001 From: Ramapitecus Date: Sat, 12 Sep 2026 20:20:09 +0500 Subject: [PATCH 9/9] preserve terminal usage from oversized realtime frames --- docs/usage-accounting.md | 2 + internal/client/codex/live/usage.go | 44 +- .../codex/live/usage_frame_projection.go | 488 ++++++++++++++++++ .../codex/live/usage_frame_projection_test.go | 113 ++++ .../usage_frame_projection_validation_test.go | 165 ++++++ internal/client/codex/live/usage_test.go | 97 ++++ 6 files changed, 902 insertions(+), 7 deletions(-) create mode 100644 internal/client/codex/live/usage_frame_projection.go create mode 100644 internal/client/codex/live/usage_frame_projection_test.go create mode 100644 internal/client/codex/live/usage_frame_projection_validation_test.go diff --git a/docs/usage-accounting.md b/docs/usage-accounting.md index afe2e55c2c8..2cada3c48ef 100644 --- a/docs/usage-accounting.md +++ b/docs/usage-accounting.md @@ -56,3 +56,5 @@ Coverage fallback requires successful route admission, including legacy authenti The durable sanitizer removes `fail.body` because provider errors can echo request content or credentials; the status and token attribution remain, and the legacy memory queue retains diagnostics. Management Gemini JSON-array streams merge elements using the same accounting buffer as SSE. xAI video responses preserve cumulative billing identity even when the provider has not yet reported usage. Legacy fan-out falls back to the queue if any subscriber overflows, even when other subscribers receive the event. Consumers must deduplicate repeated delivery by event ID. Streaming billing dimensions are retained independently of token measurements; a billing-only event keeps `usage_observed:false` instead of inventing measured token counters. + +Oversized realtime text frames are parsed incrementally while the original bytes continue downstream. The observer retains terminal identity, status, models, usage and aggregated hosted-tool counts without retaining response text/audio. The metadata projection has a 256 KiB budget and nesting limit of 256; malformed frames or metadata exceeding these bounds are not presented as measured usage. Downstream write failures still drain and account for a complete upstream frame. diff --git a/internal/client/codex/live/usage.go b/internal/client/codex/live/usage.go index d0309a2d9a1..32b69c533a3 100644 --- a/internal/client/codex/live/usage.go +++ b/internal/client/codex/live/usage.go @@ -29,7 +29,11 @@ func liveUsageDetail(payload []byte) (usage.Detail, bool) { root := gjson.ParseBytes(payload) switch root.Get("type").String() { case "response.done", "response.completed", "response.incomplete": - return helps.ParseCodexUsage(payload) + var buffer helps.StreamUsageBuffer + buffer.ObserveBillingPayload(payload) + detail, ok := helps.ParseCodexUsage(payload) + buffer.Observe(detail, ok) + return buffer.Detail() case "conversation.item.input_audio_transcription.completed": d := helps.ParseOpenAIUsage(payload) return d, d.UsageObserved @@ -113,16 +117,22 @@ func (o *liveUsageObserver) close() { // Bounded capture never limits forwarding of large media frames. type usageFrameCapture struct { - data []byte - overflow bool + data []byte + overflow bool + projection io.Writer } func (b *usageFrameCapture) Write(p []byte) (int, error) { if len(b.data)+len(p) > 1<<20 { b.overflow = true + b.data = nil } else if !b.overflow { b.data = append(b.data, p...) } + if b.projection != nil { + // Accounting parse failures must never interrupt frame forwarding. + _, _ = b.projection.Write(p) + } return len(p), nil } @@ -133,7 +143,19 @@ func forwardUsageFrame(writer io.Writer, reader io.Reader, observers ...func([]b _, errCopy := io.Copy(writer, reader) return errCopy } - capture := &usageFrameCapture{} + projectionReader, projectionWriter := io.Pipe() + defer func() { _ = projectionWriter.Close() }() + type projectionResult struct { + payload []byte + err error + } + projected := make(chan projectionResult, 1) + go func() { + payload, err := projectUsageFrame(projectionReader) + _ = projectionReader.CloseWithError(err) + projected <- projectionResult{payload: payload, err: err} + }() + capture := &usageFrameCapture{projection: projectionWriter} observedReader := io.TeeReader(reader, capture) destination := &usageForwardWriter{Writer: writer} _, errCopy := io.Copy(destination, observedReader) @@ -142,9 +164,17 @@ func forwardUsageFrame(writer io.Writer, reader io.Reader, observers ...func([]b _, errDrain := io.Copy(io.Discard, observedReader) complete = errDrain == nil } - if complete && !capture.overflow && len(capture.data) > 0 { - for _, observe := range observers { - observe(capture.data) + _ = projectionWriter.Close() + result := <-projected + if complete && result.err == nil { + payload := capture.data + if capture.overflow { + payload = result.payload + } + if len(payload) > 0 { + for _, observe := range observers { + observe(payload) + } } } return errCopy diff --git a/internal/client/codex/live/usage_frame_projection.go b/internal/client/codex/live/usage_frame_projection.go new file mode 100644 index 00000000000..3b712b39f51 --- /dev/null +++ b/internal/client/codex/live/usage_frame_projection.go @@ -0,0 +1,488 @@ +package live + +import ( + "bufio" + "encoding/json" + "errors" + "io" + "strconv" +) + +const usageProjectionBudget = 256 << 10 +const usageProjectionDepth = 256 + +// projectUsageFrame validates the entire JSON frame while retaining only fields +// used by the usage observer. Skipped strings, numbers and containers are scanned +// byte by byte, so a large transcript or media item never becomes a JSON token +// allocation. Relevant metadata has a separate, explicit size limit. +func projectUsageFrame(reader io.Reader) ([]byte, error) { + p := &usageProjectionParser{reader: bufio.NewReaderSize(reader, 16<<10), remaining: usageProjectionBudget} + projected, err := p.object("", 0) + if err != nil { + return nil, err + } + if err = p.space(); err != nil { + return nil, err + } + if _, err = p.peek(); err != io.EOF { + if err == nil { + err = errors.New("usage projection: trailing JSON data") + } + return nil, err + } + return projected, nil +} + +type usageProjectionParser struct { + reader *bufio.Reader + remaining int + capturing bool + capture []byte + readErr error +} + +func (p *usageProjectionParser) peek() (byte, error) { + if p.readErr != nil { + return 0, p.readErr + } + b, err := p.reader.Peek(1) + if err != nil { + p.readErr = err + return 0, err + } + return b[0], nil +} + +func (p *usageProjectionParser) take() (byte, error) { + if p.readErr != nil { + return 0, p.readErr + } + b, err := p.reader.ReadByte() + if err != nil { + p.readErr = err + return 0, err + } + if p.capturing { + if p.remaining == 0 { + return 0, errors.New("usage projection: accounting metadata exceeds size limit") + } + p.remaining-- + p.capture = append(p.capture, b) + } + return b, nil +} + +func (p *usageProjectionParser) expect(want byte) error { + got, err := p.take() + if err != nil { + return err + } + if got != want { + return errors.New("usage projection: invalid JSON delimiter") + } + return nil +} + +func (p *usageProjectionParser) space() error { + for { + b, err := p.peek() + if err == io.EOF { + return nil + } + if err != nil { + return err + } + if b != ' ' && b != '\n' && b != '\r' && b != '\t' { + return nil + } + if _, err = p.take(); err != nil { + return err + } + } +} + +// stringToken retains a bounded key/type token, or no bytes when limit is zero. +// An oversized skipped key cannot match the short accounting field names. +func (p *usageProjectionParser) stringToken(limit int) ([]byte, error) { + if err := p.expect('"'); err != nil { + return nil, err + } + var token []byte + if limit > 0 { + token = append(token, '"') + } + appendToken := func(b byte) { + if limit > 0 { + if len(token) == limit { + limit, token = 0, nil + } else { + token = append(token, b) + } + } + } + for { + b, err := p.take() + if err != nil { + return nil, err + } + appendToken(b) + if b == '"' { + return token, nil + } + if b < 0x20 { + return nil, errors.New("usage projection: control byte in JSON string") + } + if b != '\\' { + continue + } + b, err = p.take() + if err != nil { + return nil, err + } + appendToken(b) + switch b { + case '"', '\\', '/', 'b', 'f', 'n', 'r', 't': + case 'u': + for range 4 { + b, err = p.take() + if err != nil { + return nil, err + } + appendToken(b) + if !(b >= '0' && b <= '9' || b >= 'a' && b <= 'f' || b >= 'A' && b <= 'F') { + return nil, errors.New("usage projection: invalid JSON unicode escape") + } + } + default: + return nil, errors.New("usage projection: invalid JSON escape") + } + } +} + +func (p *usageProjectionParser) number() error { + b, _ := p.peek() + if b == '-' { + if _, err := p.take(); err != nil { + return err + } + b, _ = p.peek() + } + if b == '0' { + if _, err := p.take(); err != nil { + return err + } + } else if b >= '1' && b <= '9' { + if err := p.digits(); err != nil { + return err + } + } else { + return errors.New("usage projection: invalid JSON number") + } + b, _ = p.peek() + if b == '.' { + if _, err := p.take(); err != nil { + return err + } + if err := p.digits(); err != nil { + return err + } + } + b, _ = p.peek() + if b == 'e' || b == 'E' { + if _, err := p.take(); err != nil { + return err + } + b, _ = p.peek() + if b == '+' || b == '-' { + if _, err := p.take(); err != nil { + return err + } + } + return p.digits() + } + return nil +} + +func (p *usageProjectionParser) digits() error { + seen := false + for { + b, err := p.peek() + if err != nil && err != io.EOF { + return err + } + if err == io.EOF || b < '0' || b > '9' { + if !seen { + return errors.New("usage projection: missing JSON number digits") + } + return nil + } + if _, err = p.take(); err != nil { + return err + } + seen = true + } +} + +func (p *usageProjectionParser) value(depth int) error { + if depth > usageProjectionDepth { + return errors.New("usage projection: JSON nesting exceeds limit") + } + if err := p.space(); err != nil { + return err + } + b, err := p.peek() + if err != nil { + return err + } + switch b { + case '"': + _, err = p.stringToken(0) + return err + case '{': + return p.members(func(string) error { return p.value(depth + 1) }, false) + case '[': + return p.elements(func() error { return p.value(depth + 1) }) + case 't', 'f', 'n': + literal := map[byte]string{'t': "true", 'f': "false", 'n': "null"}[b] + for i := range literal { + if err := p.expect(literal[i]); err != nil { + return err + } + } + return nil + default: + return p.number() + } +} + +func (p *usageProjectionParser) members(visit func(string) error, keys bool) error { + if err := p.expect('{'); err != nil { + return err + } + return p.sequence('}', func() error { + limit := 0 + if keys { + limit = 256 + } + raw, err := p.stringToken(limit) + if err != nil { + return err + } + var key string + if len(raw) > 0 { + if err = json.Unmarshal(raw, &key); err != nil { + return err + } + } + if err = p.space(); err != nil { + return err + } + if err = p.expect(':'); err != nil { + return err + } + if err = p.space(); err != nil { + return err + } + return visit(key) + }) +} + +func (p *usageProjectionParser) elements(visit func() error) error { + if err := p.expect('['); err != nil { + return err + } + return p.sequence(']', visit) +} + +func (p *usageProjectionParser) sequence(end byte, visit func() error) error { + if err := p.space(); err != nil { + return err + } + if b, _ := p.peek(); b == end { + return p.expect(end) + } + for { + if err := visit(); err != nil { + return err + } + if err := p.space(); err != nil { + return err + } + b, err := p.take() + if err != nil { + return err + } + if b == end { + return nil + } + if b != ',' { + return errors.New("usage projection: invalid JSON separator") + } + if err = p.space(); err != nil { + return err + } + } +} + +func (p *usageProjectionParser) raw(depth int) (json.RawMessage, error) { + p.capturing, p.capture = true, nil + err := p.value(depth) + raw := p.capture + p.capturing, p.capture = false, nil + return raw, err +} + +func usageProjectionField(path, key string) (keep, object bool) { + switch path { + case "": + switch key { + case "type", "item_id", "content_index", "usage", "service_tier", "tool_usage": + return true, false + case "response", "session": + return true, true + } + case "response": + switch key { + case "id", "model", "status", "service_tier", "usage", "tool_usage": + return true, false + } + case "session": + if key == "model" { + return true, false + } + if key == "audio" || key == "input_audio_transcription" { + return true, true + } + case "session.audio": + return key == "input", true + case "session.audio.input": + return key == "transcription", true + case "session.input_audio_transcription", "session.audio.input.transcription": + return key == "model", false + } + return false, false +} + +func (p *usageProjectionParser) object(path string, depth int) ([]byte, error) { + if err := p.space(); err != nil { + return nil, err + } + if b, _ := p.peek(); b != '{' { + if path == "" { + return nil, errors.New("usage projection: frame must be a JSON object") + } + return p.raw(depth) + } + fields := make(map[string]json.RawMessage) + var counts usageProjectionTools + err := p.members(func(key string) error { + if key == "output" && (path == "response" || path == "") { + counts = usageProjectionTools{} + return p.tools(&counts, depth+1) + } + keep, nested := usageProjectionField(path, key) + if !keep { + return p.value(depth + 1) + } + var raw []byte + var err error + if nested { + child := key + if path != "" { + child = path + "." + key + } + raw, err = p.object(child, depth+1) + } else { + raw, err = p.raw(depth + 1) + } + if err == nil { + fields[key] = raw + } + return err + }, true) + if err != nil { + return nil, err + } + if err = counts.apply(fields); err != nil { + return nil, err + } + return json.Marshal(fields) +} + +type usageProjectionTools struct { + web, file int64 + unpriced bool +} + +func (p *usageProjectionParser) tools(counts *usageProjectionTools, depth int) error { + if b, _ := p.peek(); b != '[' { + return p.value(depth) + } + return p.elements(func() error { + if b, _ := p.peek(); b != '{' { + return p.value(depth + 1) + } + var kind string + if err := p.members(func(key string) error { + if key == "type" { + kind = "" + } + if b, _ := p.peek(); key != "type" || b != '"' { + return p.value(depth + 2) + } + raw, err := p.stringToken(256) + if err == nil && len(raw) > 0 { + err = json.Unmarshal(raw, &kind) + } + return err + }, true); err != nil { + return err + } + switch kind { + case "web_search_call": + counts.web++ + case "file_search_call": + counts.file++ + case "code_interpreter_call", "shell_call": + counts.unpriced = true + } + return nil + }) +} + +func (c usageProjectionTools) apply(fields map[string]json.RawMessage) error { + if c.web == 0 && c.file == 0 && !c.unpriced { + return nil + } + // Never invent a usage object: without measured usage, retain billing-only + // counts in tool_usage so consumers cannot mistake them for measured zero. + key := "usage" + var target map[string]json.RawMessage + if json.Unmarshal(fields[key], &target) != nil || target == nil { + key = "tool_usage" + if len(fields[key]) > 0 && string(fields[key]) != "null" { + if err := json.Unmarshal(fields[key], &target); err != nil { + return errors.New("usage projection: tool_usage must be an object") + } + } + } + if target == nil { + target = make(map[string]json.RawMessage) + } + for key, n := range map[string]int64{"web_search_calls": c.web, "file_search_calls": c.file} { + if n == 0 { + continue + } + // Compare without rewriting the provider's numeric representation. + previous, _ := strconv.ParseFloat(string(target[key]), 64) + if previous < float64(n) { + target[key] = json.RawMessage(strconv.FormatInt(n, 10)) + } + } + if c.unpriced { + target["unpriced_server_tools"] = json.RawMessage("true") + } + encoded, err := json.Marshal(target) + fields[key] = encoded + return err +} diff --git a/internal/client/codex/live/usage_frame_projection_test.go b/internal/client/codex/live/usage_frame_projection_test.go new file mode 100644 index 00000000000..298b6c8a6f0 --- /dev/null +++ b/internal/client/codex/live/usage_frame_projection_test.go @@ -0,0 +1,113 @@ +package live + +import ( + "bytes" + "encoding/json" + "io" + "strings" + "testing" + + "github.com/tidwall/gjson" +) + +func TestProjectUsageFrameSkipsLargeOutputAndKeepsAccounting(t *testing.T) { + reader := io.MultiReader(strings.NewReader(`{"type":"response.done","response":{"output":[{"type":"message","content":[{"text":"`), io.LimitReader(repeatedProjectionByte('a'), 8<<20), strings.NewReader(`"}]},{"type":"web_search_call"},{"type":"file_search_call"},{"type":"web_search_call"},{"type":"code_interpreter_call"}],"id":"r1","model":"gpt-5.4","status":"completed","service_tier":"priority","usage":{"input_tokens":10,"output_tokens":2,"total_tokens":12},"tool_usage":{"image_gen":{"input_tokens":3}}}}`)) + projected, err := projectUsageFrame(reader) + if err != nil { + t.Fatal(err) + } + if len(projected) > 4096 || !json.Valid(projected) { + t.Fatalf("invalid or oversized projection: %d bytes", len(projected)) + } + for path, want := range map[string]string{"type": "response.done", "response.id": "r1", "response.model": "gpt-5.4", "response.status": "completed", "response.service_tier": "priority", "response.usage.total_tokens": "12", "response.tool_usage.image_gen.input_tokens": "3", "response.usage.web_search_calls": "2", "response.usage.file_search_calls": "1", "response.usage.unpriced_server_tools": "true"} { + if got := gjson.GetBytes(projected, path).String(); got != want { + t.Errorf("%s = %q, want %q", path, got, want) + } + } +} + +func TestProjectUsageFrameBillingOnlyDoesNotInventUsage(t *testing.T) { + got, err := projectUsageFrame(strings.NewReader(`{"type":"response.done","response":{"id":"r1","output":[{"type":"web_search_call"},{"type":"file_search_call"},{"type":"shell_call"}],"tool_usage":{"web_search_calls":2e2,"image_gen":{"output_tokens":8}}}}`)) + if err != nil { + t.Fatal(err) + } + if gjson.GetBytes(got, "response.usage").Exists() { + t.Fatalf("invented measured usage: %s", got) + } + for path, want := range map[string]string{"response.tool_usage.web_search_calls": "200", "response.tool_usage.file_search_calls": "1", "response.tool_usage.unpriced_server_tools": "true", "response.tool_usage.image_gen.output_tokens": "8"} { + if actual := gjson.GetBytes(got, path).String(); actual != want { + t.Errorf("%s = %q, want %q", path, actual, want) + } + } + if raw := gjson.GetBytes(got, "response.tool_usage.web_search_calls").Raw; raw != "2e2" { + t.Errorf("provider numeric representation changed: %s", got) + } +} + +func TestProjectUsageFrameAllocationDoesNotScaleWithTranscript(t *testing.T) { + measure := func(size int64) float64 { + return testing.AllocsPerRun(2, func() { + reader := io.MultiReader(strings.NewReader(`{"ignored":"`), io.LimitReader(repeatedProjectionByte('a'), size), strings.NewReader(`","type":"response.done","response":{"usage":{"total_tokens":12}}}`)) + if _, err := projectUsageFrame(reader); err != nil { + t.Fatal(err) + } + }) + } + small, large := measure(1<<20), measure(16<<20) + if large > small+10 { + t.Fatalf("allocations scale with skipped transcript: 1 MiB=%v, 16 MiB=%v", small, large) + } +} + +func TestProjectUsageFramePreservesTranscriptionSessionAndNull(t *testing.T) { + for _, payload := range []string{ + `{"type":"session.updated","session":{"model":"gpt-realtime","audio":{"input":{"transcription":{"model":"gpt-4o-transcribe"}}},"input_audio_transcription":{"model":"whisper-1"}}}`, + `{"type":"session.updated","session":{"input_audio_transcription":null}}`, + `{"type":"conversation.item.input_audio_transcription.completed","item_id":"item-1","content_index":2,"usage":{"input_tokens":10,"output_tokens":3},"service_tier":"default"}`, + } { + got, err := projectUsageFrame(strings.NewReader(payload)) + if err != nil { + t.Fatal(err) + } + var original, projected any + _ = json.Unmarshal([]byte(payload), &original) + _ = json.Unmarshal(got, &projected) + a, _ := json.Marshal(original) + b, _ := json.Marshal(projected) + if !bytes.Equal(a, b) { + t.Errorf("projection = %s, want %s", got, a) + } + } +} + +func TestProjectUsageFrameRejectsMalformedAndOversizedMetadata(t *testing.T) { + for _, payload := range []string{ + `{"type":"response.done","response":{"usage":{"input_tokens":1}}`, + `{"ignored":"bad\q"}`, `{"ignored":[1,]}`, `{"ignored":01}`, `{"ignored":1.}`, + `{"ignored":true false}`, `{"ignored":"line` + "\n" + `break"}`, `{"ignored":{foo:1}}`, + `{"response":{"usage":{"input_tokens":1,}}}`, `{} false`, + `{"ignored":` + strings.Repeat("[", 300) + strings.Repeat("]", 300) + `}`, + `{"response":{"usage":{"dimension":"` + strings.Repeat("x", 300<<10) + `"}}}`, + } { + if got, err := projectUsageFrame(strings.NewReader(payload)); err == nil || got != nil { + t.Errorf("accepted malformed/oversized JSON (%d bytes): projection %q, error %v", len(payload), got, err) + } + } +} + +func TestProjectUsageFrameValidatesSkippedEscapesAndLongKeys(t *testing.T) { + payload := `{"` + strings.Repeat("x", 1<<20) + `":"escaped \" \\ \u0041","ignored":[true,false,null,-1.23e+45,{"x":"y"}],"ty\u0070e":"response.done","response":{"id":"r1","usage":{"total_tokens":12}}}` + got, err := projectUsageFrame(strings.NewReader(payload)) + if err != nil || gjson.GetBytes(got, "type").String() != "response.done" { + t.Fatalf("projection = %s, error = %v", got, err) + } +} + +type repeatedProjectionByte byte + +func (b repeatedProjectionByte) Read(p []byte) (int, error) { + for i := range p { + p[i] = byte(b) + } + return len(p), nil +} diff --git a/internal/client/codex/live/usage_frame_projection_validation_test.go b/internal/client/codex/live/usage_frame_projection_validation_test.go new file mode 100644 index 00000000000..4884428aa5b --- /dev/null +++ b/internal/client/codex/live/usage_frame_projection_validation_test.go @@ -0,0 +1,165 @@ +package live + +import ( + "bytes" + "encoding/json" + "errors" + "io" + "reflect" + "strings" + "testing" + "testing/iotest" + + "github.com/tidwall/gjson" +) + +func TestUsageProjectionValidationGrammarMatchesJSON(t *testing.T) { + values := []string{ + `0`, `-0`, `12`, `-123456789012345678901234567890`, `0.0`, `-0.01`, `1e0`, `1E+10`, `1e-10`, `1e9999`, + `true`, `false`, `null`, `""`, `"\"\\\/\b\f\n\r\t"`, `"\u0041\u00e9\uD83D\uDE00"`, `"\ud800"`, `"Русский 日本語 😀"`, + `[]`, `{}`, `[0,true,null,{"nested":["\u0041",-1.2e+3]}]`, `{"same":1,"same":2}`, + `01`, `-01`, `+1`, `.1`, `1.`, `1e`, `1e+`, `1e-`, `--1`, `NaN`, `Infinity`, `0x10`, `1e1.2`, + `tru`, `True`, `nul`, `false true`, `"\q"`, `"\u000"`, `"\u00xz"`, `"\x20"`, `"unterminated`, + `"raw` + "\x00" + `control"`, `"raw` + "\n" + `newline"`, `[1,]`, `[,1]`, `{"a":1,}`, `{"a" 1}`, `{a:1}`, `{"a":}`, + } + for _, value := range values { + for _, key := range []string{"ignored", "usage"} { + payload := ` {"type":"response.done","` + key + `":` + value + `,"response":{"id":"r1"}} ` + wantValid := json.Valid([]byte(payload)) + for _, fragmented := range []bool{false, true} { + var reader io.Reader = strings.NewReader(payload) + if fragmented { + reader = iotest.OneByteReader(reader) + } + got, err := projectUsageFrame(reader) + if (err == nil) != wantValid { + t.Fatalf("JSON agreement mismatch key=%s value=%q fragmented=%v valid=%v error=%v", key, value, fragmented, wantValid, err) + } + if err != nil { + if got != nil { + t.Fatalf("invalid frame returned partial projection: %s", got) + } + continue + } + if !json.Valid(got) || gjson.GetBytes(got, "type").String() != "response.done" || gjson.GetBytes(got, "response.id").String() != "r1" { + t.Fatalf("invalid metadata projection for value %q: %s", value, got) + } + } + } + } +} + +func TestUsageProjectionValidationEveryTruncatedPrefix(t *testing.T) { + payload := `{"ignored":[null,-1.25e-10,{"key":"\uD83D\uDE00\\quoted\""}],"type":"response.done","response":{"usage":{"input_tokens":9007199254740993,"output_tokens":0}}}` + for end := 0; end <= len(payload); end++ { + prefix := payload[:end] + got, err := projectUsageFrame(iotest.OneByteReader(strings.NewReader(prefix))) + if (err == nil) != json.Valid([]byte(prefix)) { + t.Fatalf("prefix %d validity mismatch: %q, projection=%s error=%v", end, prefix, got, err) + } + if err != nil && got != nil { + t.Fatalf("prefix %d emitted partial metadata", end) + } + } + for _, suffix := range []string{"\n\t\r ", "{}", "false", "0", "\x00"} { + document := payload + suffix + _, err := projectUsageFrame(iotest.OneByteReader(strings.NewReader(document))) + if (err == nil) != json.Valid([]byte(document)) { + t.Fatalf("suffix %q: error=%v", suffix, err) + } + } +} + +func TestUsageProjectionValidationSelectedFieldsExact(t *testing.T) { + payload := `{"ty\u0070e":"response.done","item_id":"item-\u2603","content_index":2,"usage":{"input_tokens":9007199254740993,"cost_in_usd_ticks":9223372036854775807,"nested":{"unknown":true}},"tool_usage":{"image_gen":{"input_tokens":7}},"service_tier":"priority","response":{"id":"\ud83d\ude00","model":"model","status":"completed","service_tier":"flex","usage":{"input_tokens":1e3,"output_tokens":-0,"modality":[1,2,3]},"tool_usage":{"custom":"\u0061"}},"session":{"model":"realtime","audio":{"input":{"transcription":{"model":"transcribe"}}},"input_audio_transcription":{"model":"legacy"}}}` + got, err := projectUsageFrame(iotest.OneByteReader(strings.NewReader(payload))) + if err != nil { + t.Fatal(err) + } + decode := func(data []byte) any { + t.Helper() + decoder := json.NewDecoder(bytes.NewReader(data)) + decoder.UseNumber() + var value any + if err := decoder.Decode(&value); err != nil { + t.Fatal(err) + } + return value + } + if !reflect.DeepEqual(decode(got), decode([]byte(payload))) { + t.Fatalf("selected metadata changed:\ngot %s\nwant %s", got, payload) + } +} + +func TestUsageProjectionValidationDuplicateKeys(t *testing.T) { + tests := []struct { + name, payload string + web, file int64 + }{ + {"last output empty", `{"response":{"output":[{"type":"web_search_call"}],"output":[],"usage":{"input_tokens":1}}}`, 0, 0}, + {"last output replacement", `{"response":{"output":[{"type":"web_search_call"}],"output":[{"type":"file_search_call"}],"usage":{"input_tokens":1}}}`, 0, 1}, + {"last output null", `{"response":{"output":[{"type":"web_search_call"}],"output":null,"usage":{"input_tokens":1}}}`, 0, 0}, + {"last type string", `{"response":{"output":[{"type":"web_search_call","type":"file_search_call"}],"usage":{"input_tokens":1}}}`, 0, 1}, + {"last type null", `{"response":{"output":[{"type":"web_search_call","type":null}],"usage":{"input_tokens":1}}}`, 0, 0}, + {"last type long", `{"response":{"output":[{"type":"web_search_call","type":"` + strings.Repeat("x", 300) + `"}],"usage":{"input_tokens":1}}}`, 0, 0}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + if !json.Valid([]byte(tc.payload)) { + t.Fatal("invalid fixture") + } + got, err := projectUsageFrame(iotest.OneByteReader(strings.NewReader(tc.payload))) + if err != nil { + t.Fatal(err) + } + if web, file := gjson.GetBytes(got, "response.usage.web_search_calls").Int(), gjson.GetBytes(got, "response.usage.file_search_calls").Int(); web != tc.web || file != tc.file { + t.Fatalf("duplicate keys changed billable counts: web=%d file=%d projection=%s", web, file, got) + } + }) + } +} + +func TestUsageProjectionValidationLargeUnknownData(t *testing.T) { + for _, kind := range []string{"string", "number", "key"} { + t.Run(kind, func(t *testing.T) { + prefix, suffix := `{"ignored":"`, `","type":"response.done","response":{"id":"after-large","usage":{"input_tokens":1}}}` + fill := byte('a') + if kind == "number" { + prefix = `{"ignored":1` + suffix = `,"type":"response.done","response":{"id":"after-large","usage":{"input_tokens":1}}}` + fill = '0' + } + if kind == "key" { + prefix = `{"` + suffix = `":null,"type":"response.done","response":{"id":"after-large","usage":{"input_tokens":1}}}` + } + reader := io.MultiReader(strings.NewReader(prefix), io.LimitReader(validationRepeatedByte(fill), 2<<20), strings.NewReader(suffix)) + got, err := projectUsageFrame(iotest.OneByteReader(reader)) + if err != nil { + t.Fatal(err) + } + if len(got) > 1024 || !json.Valid(got) || gjson.GetBytes(got, "response.id").String() != "after-large" || gjson.GetBytes(got, "response.usage.input_tokens").Int() != 1 { + t.Fatalf("large unknown %s affected projection: %s", kind, got) + } + }) + } +} + +type validationRepeatedByte byte + +func (b validationRepeatedByte) Read(p []byte) (int, error) { + for i := range p { + p[i] = byte(b) + } + return len(p), nil +} + +func TestUsageProjectionValidationPropagatesReaderError(t *testing.T) { + failure := errors.New("validation transport failure") + for _, prefix := range []string{``, `{"ignored":"abc`, `{"ignored":1`, `{"type":"response.done"}`} { + got, err := projectUsageFrame(io.MultiReader(strings.NewReader(prefix), iotest.ErrReader(failure))) + if err == nil || got != nil { + t.Fatalf("reader error emitted accounting: prefix=%q got=%s error=%v", prefix, got, err) + } + } +} diff --git a/internal/client/codex/live/usage_test.go b/internal/client/codex/live/usage_test.go index 0c71a7170c3..e074d4a969c 100644 --- a/internal/client/codex/live/usage_test.go +++ b/internal/client/codex/live/usage_test.go @@ -1,6 +1,7 @@ package live import ( + "bytes" "context" "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/usage" @@ -77,3 +78,99 @@ func TestTranscriptionUsesItsOwnModelAndContentIdentity(t *testing.T) { } } } + +func TestOversizedTerminalUsageFramePreservesAccounting(t *testing.T) { + for _, pending := range []bool{false, true} { + for _, disconnected := range []bool{false, true} { + name := "terminal_only" + if pending { + name = "pending" + } + if disconnected { + name += "_disconnected" + } + t.Run(name, func(t *testing.T) { + sink := &liveReviewSink{authID: t.Name(), records: make(chan usage.Record, 10)} + usage.RegisterPlugin(sink) + o := newLiveUsageObserver(context.Background(), &auth.Auth{ID: t.Name()}, "gpt-realtime") + if pending { + o.observe([]byte(`{"type":"response.created","response":{"id":"large-response","model":"gpt-realtime"}}`)) + } + payload := `{"response":{"output":[{"type":"message","content":[{"text":"` + strings.Repeat(`large escaped \" text `, 100000) + `"}]},{"type":"web_search_call"},{"type":"web_search_call"}],"usage":{"input_tokens":100,"output_tokens":20,"total_tokens":120,"input_tokens_details":{"audio_tokens":60,"text_tokens":40}},"id":"large-response","model":"gpt-realtime","status":"completed"},"type":"response.done"}` + var forwarded bytes.Buffer + var writer io.Writer = &forwarded + if disconnected { + writer = disconnectedUsageWriter{} + } + err := forwardUsageFrame(writer, strings.NewReader(payload), o.observe) + if disconnected && err != io.ErrClosedPipe || !disconnected && err != nil { + t.Fatalf("forward error: %v", err) + } + if !disconnected && forwarded.String() != payload { + t.Fatal("forwarded frame was changed") + } + o.close() + if len(sink.records) != 1 { + t.Fatalf("want one record, got %d", len(sink.records)) + } + r := <-sink.records + if r.Failed || !r.Detail.UsageObserved || r.Detail.InputTokens != 100 || r.Detail.OutputTokens != 20 || r.Model != "gpt-realtime" { + t.Fatalf("lost terminal accounting: %+v", r) + } + if !strings.Contains(r.Detail.RawUsage, `"web_search_calls":2`) || !strings.Contains(r.Detail.RawUsage, `"audio_tokens":60`) { + t.Fatalf("lost billing dimensions: %s", r.Detail.RawUsage) + } + }) + } + } +} + +func TestOversizedTranscriptionFramePreservesAccounting(t *testing.T) { + sink := &liveReviewSink{authID: t.Name(), records: make(chan usage.Record, 10)} + usage.RegisterPlugin(sink) + o := newLiveUsageObserver(context.Background(), &auth.Auth{ID: t.Name()}, "gpt-realtime") + o.observe([]byte(`{"type":"session.updated","session":{"audio":{"input":{"transcription":{"model":"gpt-4o-transcribe"}}}}}`)) + payload := `{"transcript":"` + strings.Repeat("text ", 300000) + `","type":"conversation.item.input_audio_transcription.completed","item_id":"item1","content_index":2,"usage":{"input_tokens":40,"output_tokens":5,"total_tokens":45}}` + if err := forwardUsageFrame(io.Discard, strings.NewReader(payload), o.observe); err != nil { + t.Fatal(err) + } + if len(sink.records) != 1 { + t.Fatalf("want transcription record, got %d", len(sink.records)) + } + r := <-sink.records + if r.Failed || r.Model != "gpt-4o-transcribe" || r.Kind != "tool" || r.Detail.TotalTokens != 45 { + t.Fatalf("lost transcription accounting: %+v", r) + } +} + +func TestLiveBillingWithoutTokensRemainsUnobserved(t *testing.T) { + d, ok := liveUsageDetail([]byte(`{"type":"response.done","response":{"id":"tool-only","tool_usage":{"web_search_calls":2}}}`)) + if !ok || d.UsageObserved || !strings.Contains(d.RawUsage, `"web_search_calls":2`) { + t.Fatalf("lost unobserved billing metadata: %+v ok=%v", d, ok) + } +} + +func TestUsageFrameMalformedContentDoesNotInterruptForwarding(t *testing.T) { + payload := `{"type":"response.done","response":{"usage":{"input_tokens":1}}} trailing` + var forwarded bytes.Buffer + observed := false + if err := forwardUsageFrame(&forwarded, strings.NewReader(payload), func([]byte) { observed = true }); err != nil { + t.Fatal(err) + } + if forwarded.String() != payload || observed { + t.Fatalf("forwarded changed or malformed event observed: observed=%v", observed) + } +} + +type usageReadFailure struct{} + +func (usageReadFailure) Read([]byte) (int, error) { return 0, io.ErrUnexpectedEOF } + +func TestUsageFrameReadFailureDoesNotPublishTerminal(t *testing.T) { + payload := `{"type":"response.done","response":{"id":"partial","usage":{"input_tokens":1}}}` + observed := false + err := forwardUsageFrame(io.Discard, io.MultiReader(strings.NewReader(payload), usageReadFailure{}), func([]byte) { observed = true }) + if err != io.ErrUnexpectedEOF || observed { + t.Fatalf("incomplete frame error=%v observed=%v", err, observed) + } +}