From 5ef2f45ad34025b0e357b861602bea5ce6bb6ea3 Mon Sep 17 00:00:00 2001 From: zap Date: Sun, 27 Sep 2026 19:46:06 +0800 Subject: [PATCH] fix(usage): record real backend token counts and mark estimated rows Issue #188: usage statistics recorded estimated token counts instead of real backend-reported values. Proxy streaming paths did not request include_usage; local paths ignored engine-reported tokens; users could not distinguish real from estimated. - Add OnUsage callback in Options to thread real prompt/completion tokens from llama-server and OpenAI-compatible engines to handlers - Add stream_options.include_usage to proxy request bodies so backends report real usage in stream tails - Add resolveUsageTokens with per-side sanity validation: use real values when available, fall back to estimation otherwise, mark mixed rows as estimated - Improve CJK token estimation (0.6 tokens/char vs flat 4 chars/token) - Capture partial stream content for fallback estimation when upstream reports no usage; record even on mid-stream disconnect - Add estimated_requests column (PRAGMA-checked migration), expose in API/OpenAPI/frontend with amber badge - Move Anthropic native stream recording outside err==nil guard - Add tests for proxy stream usage, fallback estimation, and mid-stream error recording 22 files, +1042/-206 --- internal/config/api_auth.go | 139 ++++++------ internal/config/api_usage_store.go | 101 ++++++--- internal/inference/llama.go | 16 ++ internal/inference/openai.go | 27 ++- internal/inference/options.go | 6 + internal/server/handlers.go | 30 ++- internal/server/handlers_anthropic.go | 125 +++++++++-- internal/server/handlers_anthropic_test.go | 172 +++++++++++++++ internal/server/handlers_api_keys.go | 43 ++-- internal/server/handlers_chat_tools.go | 5 + internal/server/handlers_openai.go | 50 ++++- internal/server/handlers_openai_test.go | 198 +++++++++++++++++- internal/server/handlers_responses.go | 27 ++- internal/server/observability.go | 28 +-- internal/server/server_test.go | 9 + internal/server/static/openapi/local-api.json | 4 + internal/server/usage.go | 158 +++++++++++++- openapi/local-api.json | 4 + pkg/api/types.go | 89 ++++---- web/src/api/client.ts | 1 + web/src/i18n.ts | 2 + web/src/pages/AIGateway.tsx | 14 +- 22 files changed, 1042 insertions(+), 206 deletions(-) diff --git a/internal/config/api_auth.go b/internal/config/api_auth.go index 9d8d5cfe..52c15a79 100644 --- a/internal/config/api_auth.go +++ b/internal/config/api_auth.go @@ -208,55 +208,58 @@ type APIUsageEvent struct { LimitedCount int64 InputTokens int64 OutputTokens int64 + Estimated bool CreatedAt time.Time } type APIUsageRecord struct { - APIKeyID string `json:"api_key_id"` - APIKeyName string `json:"api_key_name"` - Model string `json:"model"` - Source string `json:"source,omitempty"` - SourceType string `json:"source_type,omitempty"` - SourceName string `json:"source_name,omitempty"` - PoolID string `json:"pool_id,omitempty"` - PoolName string `json:"pool_name,omitempty"` - PoolModel string `json:"pool_model,omitempty"` - ActualMemberID string `json:"actual_member_id,omitempty"` - MemberModel string `json:"member_model,omitempty"` - EstimatedCost float64 `json:"estimated_cost,omitempty"` - CostCurrency string `json:"cost_currency,omitempty"` - CostKnown bool `json:"cost_known"` - FallbackCount int64 `json:"fallback_count,omitempty"` - LimitedCount int64 `json:"limited_count,omitempty"` - Requests int64 `json:"requests"` - InputTokens int64 `json:"input_tokens"` - OutputTokens int64 `json:"output_tokens"` - TotalTokens int64 `json:"total_tokens"` - LastUsedAt time.Time `json:"last_used_at"` + APIKeyID string `json:"api_key_id"` + APIKeyName string `json:"api_key_name"` + Model string `json:"model"` + Source string `json:"source,omitempty"` + SourceType string `json:"source_type,omitempty"` + SourceName string `json:"source_name,omitempty"` + PoolID string `json:"pool_id,omitempty"` + PoolName string `json:"pool_name,omitempty"` + PoolModel string `json:"pool_model,omitempty"` + ActualMemberID string `json:"actual_member_id,omitempty"` + MemberModel string `json:"member_model,omitempty"` + EstimatedCost float64 `json:"estimated_cost,omitempty"` + CostCurrency string `json:"cost_currency,omitempty"` + CostKnown bool `json:"cost_known"` + FallbackCount int64 `json:"fallback_count,omitempty"` + LimitedCount int64 `json:"limited_count,omitempty"` + Requests int64 `json:"requests"` + EstimatedRequests int64 `json:"estimated_requests,omitempty"` + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + TotalTokens int64 `json:"total_tokens"` + LastUsedAt time.Time `json:"last_used_at"` } type APIUsageEventRecord struct { - APIKeyID string `json:"api_key_id"` - APIKeyName string `json:"api_key_name"` - Model string `json:"model"` - Source string `json:"source,omitempty"` - SourceType string `json:"source_type,omitempty"` - SourceName string `json:"source_name,omitempty"` - PoolID string `json:"pool_id,omitempty"` - PoolName string `json:"pool_name,omitempty"` - PoolModel string `json:"pool_model,omitempty"` - ActualMemberID string `json:"actual_member_id,omitempty"` - MemberModel string `json:"member_model,omitempty"` - EstimatedCost float64 `json:"estimated_cost,omitempty"` - CostCurrency string `json:"cost_currency,omitempty"` - CostKnown bool `json:"cost_known"` - FallbackCount int64 `json:"fallback_count,omitempty"` - LimitedCount int64 `json:"limited_count,omitempty"` - Requests int64 `json:"requests,omitempty"` - InputTokens int64 `json:"input_tokens"` - OutputTokens int64 `json:"output_tokens"` - TotalTokens int64 `json:"total_tokens"` - CreatedAt time.Time `json:"created_at"` + APIKeyID string `json:"api_key_id"` + APIKeyName string `json:"api_key_name"` + Model string `json:"model"` + Source string `json:"source,omitempty"` + SourceType string `json:"source_type,omitempty"` + SourceName string `json:"source_name,omitempty"` + PoolID string `json:"pool_id,omitempty"` + PoolName string `json:"pool_name,omitempty"` + PoolModel string `json:"pool_model,omitempty"` + ActualMemberID string `json:"actual_member_id,omitempty"` + MemberModel string `json:"member_model,omitempty"` + EstimatedCost float64 `json:"estimated_cost,omitempty"` + CostCurrency string `json:"cost_currency,omitempty"` + CostKnown bool `json:"cost_known"` + FallbackCount int64 `json:"fallback_count,omitempty"` + LimitedCount int64 `json:"limited_count,omitempty"` + Requests int64 `json:"requests,omitempty"` + EstimatedRequests int64 `json:"estimated_requests,omitempty"` + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + TotalTokens int64 `json:"total_tokens"` + CreatedAt time.Time `json:"created_at"` } type APIUsageState struct { @@ -349,6 +352,7 @@ func compactAPIUsageEvents(events []APIUsageEventRecord) []APIUsageEventRecord { out[i].SourceName = latestNonEmpty(out[i].SourceName, event.SourceName) out[i].PoolName = latestNonEmpty(out[i].PoolName, event.PoolName) out[i].Requests += event.Requests + out[i].EstimatedRequests += event.EstimatedRequests out[i].FallbackCount += event.FallbackCount out[i].LimitedCount += event.LimitedCount out[i].EstimatedCost += event.EstimatedCost @@ -384,6 +388,7 @@ func upsertAPIUsageRecord(state *APIUsageState, event APIUsageEventRecord) { state.Records[i].SourceName = event.SourceName state.Records[i].PoolName = latestNonEmpty(state.Records[i].PoolName, event.PoolName) state.Records[i].Requests += requests + state.Records[i].EstimatedRequests += event.EstimatedRequests state.Records[i].FallbackCount += event.FallbackCount state.Records[i].LimitedCount += event.LimitedCount state.Records[i].EstimatedCost += event.EstimatedCost @@ -395,27 +400,28 @@ func upsertAPIUsageRecord(state *APIUsageState, event APIUsageEventRecord) { } } state.Records = append(state.Records, APIUsageRecord{ - APIKeyID: event.APIKeyID, - APIKeyName: event.APIKeyName, - Model: event.Model, - Source: event.Source, - SourceType: event.SourceType, - SourceName: event.SourceName, - PoolID: event.PoolID, - PoolName: event.PoolName, - PoolModel: event.PoolModel, - ActualMemberID: event.ActualMemberID, - MemberModel: event.MemberModel, - EstimatedCost: event.EstimatedCost, - CostCurrency: event.CostCurrency, - CostKnown: event.CostKnown, - FallbackCount: event.FallbackCount, - LimitedCount: event.LimitedCount, - Requests: requests, - InputTokens: event.InputTokens, - OutputTokens: event.OutputTokens, - TotalTokens: apiUsageEventTotalTokens(event), - LastUsedAt: event.CreatedAt, + APIKeyID: event.APIKeyID, + APIKeyName: event.APIKeyName, + Model: event.Model, + Source: event.Source, + SourceType: event.SourceType, + SourceName: event.SourceName, + PoolID: event.PoolID, + PoolName: event.PoolName, + PoolModel: event.PoolModel, + ActualMemberID: event.ActualMemberID, + MemberModel: event.MemberModel, + EstimatedCost: event.EstimatedCost, + CostCurrency: event.CostCurrency, + CostKnown: event.CostKnown, + FallbackCount: event.FallbackCount, + LimitedCount: event.LimitedCount, + Requests: requests, + EstimatedRequests: event.EstimatedRequests, + InputTokens: event.InputTokens, + OutputTokens: event.OutputTokens, + TotalTokens: apiUsageEventTotalTokens(event), + LastUsedAt: event.CreatedAt, }) } @@ -463,6 +469,13 @@ func apiUsageEventRequests(event APIUsageEventRecord) int64 { return 1 } +func apiUsageEventEstimatedRequests(event APIUsageEventRecord) int64 { + if event.EstimatedRequests > 0 { + return event.EstimatedRequests + } + return 0 +} + func apiUsageEventTotalTokens(event APIUsageEventRecord) int64 { if event.TotalTokens != 0 { return event.TotalTokens diff --git a/internal/config/api_usage_store.go b/internal/config/api_usage_store.go index 2d164e9b..5cf33f2f 100644 --- a/internal/config/api_usage_store.go +++ b/internal/config/api_usage_store.go @@ -45,6 +45,7 @@ CREATE TABLE IF NOT EXISTS api_usage_events ( fallback_count INTEGER NOT NULL DEFAULT 0, limited_count INTEGER NOT NULL DEFAULT 0, requests INTEGER NOT NULL DEFAULT 0, + estimated_requests INTEGER NOT NULL DEFAULT 0, input_tokens INTEGER NOT NULL DEFAULT 0, output_tokens INTEGER NOT NULL DEFAULT 0, total_tokens INTEGER NOT NULL DEFAULT 0, @@ -65,15 +66,15 @@ CREATE TABLE IF NOT EXISTS api_usage_meta ( const apiUsageColumns = `day, api_key_id, model, source, source_type, pool_id, pool_model, actual_member_id, member_model, cost_currency, cost_known, api_key_name, source_name, - pool_name, estimated_cost, fallback_count, limited_count, requests, input_tokens, - output_tokens, total_tokens, created_at` + pool_name, estimated_cost, fallback_count, limited_count, requests, estimated_requests, + input_tokens, output_tokens, total_tokens, created_at` // apiUsageUpsertStatement folds a new event into its day bucket in a single // write. Names fall back to the stored value when the new event omits them, // matching latestNonEmpty. const apiUsageUpsertStatement = ` INSERT INTO api_usage_events (` + apiUsageColumns + `) -VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT ( day, api_key_id, model, source, source_type, pool_id, pool_model, actual_member_id, member_model, @@ -86,6 +87,7 @@ ON CONFLICT ( fallback_count = api_usage_events.fallback_count + excluded.fallback_count, limited_count = api_usage_events.limited_count + excluded.limited_count, requests = api_usage_events.requests + excluded.requests, + estimated_requests = api_usage_events.estimated_requests + excluded.estimated_requests, input_tokens = api_usage_events.input_tokens + excluded.input_tokens, output_tokens = api_usage_events.output_tokens + excluded.output_tokens, total_tokens = api_usage_events.total_tokens + excluded.total_tokens, @@ -186,6 +188,21 @@ func (s *APIUsageStore) openLocked() (*sql.DB, error) { _ = db.Close() return nil, fmt.Errorf("initializing API usage database: %w", err) } + // Older databases created before the estimated_requests column existed + // need an ALTER TABLE; new databases already have it from the schema. + // Check PRAGMA table_info so we only ALTER when the column is missing, + // and any real failure surfaces instead of being silently swallowed. + needsMigration, err := needsEstimatedRequestsColumn(db) + if err != nil { + _ = db.Close() + return nil, fmt.Errorf("checking API usage database schema: %w", err) + } + if needsMigration { + if _, err := db.Exec("ALTER TABLE api_usage_events ADD COLUMN estimated_requests INTEGER NOT NULL DEFAULT 0"); err != nil { + _ = db.Close() + return nil, fmt.Errorf("migrating API usage database: %w", err) + } + } s.db = db // A failed import must not take usage recording down: the JSON file stays // in place so the next start can retry it. @@ -299,7 +316,7 @@ func apiUsageSelectEvents(db *sql.DB, options APIUsageListOptions) ([]APIUsageEv &event.PoolID, &event.PoolModel, &event.ActualMemberID, &event.MemberModel, &event.CostCurrency, &costKnown, &event.APIKeyName, &event.SourceName, &event.PoolName, &event.EstimatedCost, &event.FallbackCount, &event.LimitedCount, - &event.Requests, &event.InputTokens, &event.OutputTokens, &event.TotalTokens, + &event.Requests, &event.EstimatedRequests, &event.InputTokens, &event.OutputTokens, &event.TotalTokens, &createdAt, ); err != nil { return nil, fmt.Errorf("reading API usage: %w", err) @@ -347,27 +364,28 @@ func apiUsageEventRecord(event APIUsageEvent) (APIUsageEventRecord, bool) { event.CostKnown, event.CostCurrency, event.EstimatedCost, ) compacted := compactAPIUsageEvents([]APIUsageEventRecord{{ - APIKeyID: event.APIKeyID, - APIKeyName: event.APIKeyName, - Model: event.Model, - Source: event.Source, - SourceType: event.SourceType, - SourceName: event.SourceName, - PoolID: event.PoolID, - PoolName: event.PoolName, - PoolModel: event.PoolModel, - ActualMemberID: event.ActualMemberID, - MemberModel: event.MemberModel, - EstimatedCost: estimatedCost, - CostCurrency: costCurrency, - CostKnown: costKnown, - FallbackCount: event.FallbackCount, - LimitedCount: event.LimitedCount, - Requests: 1, - InputTokens: event.InputTokens, - OutputTokens: event.OutputTokens, - TotalTokens: event.InputTokens + event.OutputTokens, - CreatedAt: createdAt, + APIKeyID: event.APIKeyID, + APIKeyName: event.APIKeyName, + Model: event.Model, + Source: event.Source, + SourceType: event.SourceType, + SourceName: event.SourceName, + PoolID: event.PoolID, + PoolName: event.PoolName, + PoolModel: event.PoolModel, + ActualMemberID: event.ActualMemberID, + MemberModel: event.MemberModel, + EstimatedCost: estimatedCost, + CostCurrency: costCurrency, + CostKnown: costKnown, + FallbackCount: event.FallbackCount, + LimitedCount: event.LimitedCount, + Requests: 1, + EstimatedRequests: boolToInt64(event.Estimated), + InputTokens: event.InputTokens, + OutputTokens: event.OutputTokens, + TotalTokens: event.InputTokens + event.OutputTokens, + CreatedAt: createdAt, }}) if len(compacted) == 0 { return APIUsageEventRecord{}, false @@ -399,6 +417,7 @@ func apiUsageInsertArgs(event APIUsageEventRecord) []any { event.FallbackCount, event.LimitedCount, apiUsageEventRequests(event), + apiUsageEventEstimatedRequests(event), event.InputTokens, event.OutputTokens, apiUsageEventTotalTokens(event), @@ -413,9 +432,41 @@ func apiUsageTimeToStorage(value time.Time) int64 { return value.UTC().UnixNano() } +func boolToInt64(v bool) int64 { + if v { + return 1 + } + return 0 +} + func apiUsageTimeFromStorage(value int64) time.Time { if value == 0 { return time.Time{} } return time.Unix(0, value).UTC() } + +func needsEstimatedRequestsColumn(db *sql.DB) (bool, error) { + rows, err := db.Query("PRAGMA table_info(api_usage_events)") + if err != nil { + return false, fmt.Errorf("querying table_info: %w", err) + } + defer rows.Close() + for rows.Next() { + var cid int + var name, ctype string + var notnull int + var dfltValue any + var pk int + if err := rows.Scan(&cid, &name, &ctype, ¬null, &dfltValue, &pk); err != nil { + return false, fmt.Errorf("scanning table_info: %w", err) + } + if name == "estimated_requests" { + return false, nil + } + } + if err := rows.Err(); err != nil { + return false, fmt.Errorf("reading table_info: %w", err) + } + return true, nil +} diff --git a/internal/inference/llama.go b/internal/inference/llama.go index 360414c3..332edce9 100644 --- a/internal/inference/llama.go +++ b/internal/inference/llama.go @@ -1026,6 +1026,7 @@ func trimOldestNonSystemMessage(messages []Message) ([]Message, bool) { type llamaChatUsage struct { PromptTokens int64 `json:"prompt_tokens"` + CompletionTokens int64 `json:"completion_tokens"` PromptTokensDetails struct { CachedTokens int64 `json:"cached_tokens"` } `json:"prompt_tokens_details"` @@ -1057,6 +1058,19 @@ func recordLlamaCacheUsage(collector *CacheUsageCollector, usage llamaChatUsage, } } +// reportLlamaGenerationUsage forwards the prompt/completion token counts a +// backend advertised to the caller's OnUsage callback. It is a no-op when the +// caller did not request usage reporting or the backend reported nothing. +func reportLlamaGenerationUsage(onUsage func(int64, int64), usage llamaChatUsage) { + if onUsage == nil { + return + } + if usage.PromptTokens <= 0 && usage.CompletionTokens <= 0 { + return + } + onUsage(usage.PromptTokens, usage.CompletionTokens) +} + func (e *llamaEngine) handleStream(body io.Reader, onToken TokenCallback, opts Options) (string, error) { return e.handleStreamWithCacheUsage(body, onToken, opts, nil) } @@ -1092,6 +1106,7 @@ func (e *llamaEngine) handleStreamWithCacheUsage(body io.Reader, onToken TokenCa continue } recordLlamaCacheUsage(collector, chunk.Usage, chunk.Timings) + reportLlamaGenerationUsage(opts.OnUsage, chunk.Usage) if len(chunk.Choices) > 0 { d := chunk.Choices[0].Delta // Use at most one delta text per chunk. Some llama-server builds populate both @@ -1146,6 +1161,7 @@ func (e *llamaEngine) handleNonStreamWithCacheUsage(body io.Reader, opts Options return "", fmt.Errorf("decoding response: %w", err) } recordLlamaCacheUsage(collector, resp.Usage, resp.Timings) + reportLlamaGenerationUsage(opts.OnUsage, resp.Usage) if len(resp.Choices) == 0 { return "", fmt.Errorf("no choices in response") } diff --git a/internal/inference/openai.go b/internal/inference/openai.go index 5f80316c..b578274c 100644 --- a/internal/inference/openai.go +++ b/internal/inference/openai.go @@ -330,16 +330,16 @@ func (e *openAIEngine) Chat(ctx context.Context, messages []Message, opts Option collector := cacheUsageCollectorFromContext(ctx) if stream { - return e.handleStreamWithCacheUsage(resp.Body, onToken, collector) + return e.handleStreamWithCacheUsage(resp.Body, onToken, collector, opts.OnUsage) } - return e.handleJSONResponseWithCacheUsage(resp.Body, collector) + return e.handleJSONResponseWithCacheUsage(resp.Body, collector, opts.OnUsage) } func (e *openAIEngine) handleStream(body io.Reader, onToken TokenCallback) (string, error) { - return e.handleStreamWithCacheUsage(body, onToken, nil) + return e.handleStreamWithCacheUsage(body, onToken, nil, nil) } -func (e *openAIEngine) handleStreamWithCacheUsage(body io.Reader, onToken TokenCallback, collector *CacheUsageCollector) (string, error) { +func (e *openAIEngine) handleStreamWithCacheUsage(body io.Reader, onToken TokenCallback, collector *CacheUsageCollector, onUsage func(int64, int64)) (string, error) { scanner := bufio.NewScanner(body) var full strings.Builder reasoningOpen := false @@ -366,6 +366,7 @@ func (e *openAIEngine) handleStreamWithCacheUsage(body io.Reader, onToken TokenC continue } recordOpenAICacheUsage(collector, chatResp.Usage) + reportOpenAIGenerationUsage(onUsage, chatResp.Usage) if len(chatResp.Choices) == 0 || chatResp.Choices[0].Delta == nil { continue } @@ -396,15 +397,16 @@ func (e *openAIEngine) handleStreamWithCacheUsage(body io.Reader, onToken TokenC } func (e *openAIEngine) handleJSONResponse(body io.Reader) (string, error) { - return e.handleJSONResponseWithCacheUsage(body, nil) + return e.handleJSONResponseWithCacheUsage(body, nil, nil) } -func (e *openAIEngine) handleJSONResponseWithCacheUsage(body io.Reader, collector *CacheUsageCollector) (string, error) { +func (e *openAIEngine) handleJSONResponseWithCacheUsage(body io.Reader, collector *CacheUsageCollector, onUsage func(int64, int64)) (string, error) { var chatResp api.OpenAIChatResponse if err := json.NewDecoder(body).Decode(&chatResp); err != nil { return "", fmt.Errorf("decoding response: %w", err) } recordOpenAICacheUsage(collector, chatResp.Usage) + reportOpenAIGenerationUsage(onUsage, chatResp.Usage) if len(chatResp.Choices) == 0 || chatResp.Choices[0].Message == nil { return "", fmt.Errorf("no message in response") } @@ -435,6 +437,19 @@ func recordOpenAICacheUsage(collector *CacheUsageCollector, usage api.OpenAIUsag collector.add(read, write, prompt) } +// reportOpenAIGenerationUsage forwards the prompt/completion token counts a +// backend advertised to the caller's OnUsage callback. It is a no-op when the +// caller did not request usage reporting or the backend reported nothing. +func reportOpenAIGenerationUsage(onUsage func(int64, int64), usage api.OpenAIUsage) { + if onUsage == nil { + return + } + if usage.PromptTokens <= 0 && usage.CompletionTokens <= 0 { + return + } + onUsage(int64(usage.PromptTokens), int64(usage.CompletionTokens)) +} + // SupportsNativeToolStreaming reports that cloud and third-party // OpenAI-compatible backends return standard streaming tool-call deltas, // so tool requests do not need local aggregation and normalization. diff --git a/internal/inference/options.go b/internal/inference/options.go index 3140691e..7cff2814 100644 --- a/internal/inference/options.go +++ b/internal/inference/options.go @@ -12,6 +12,12 @@ type Options struct { // DisableThinking forces routing-style requests to skip provider thinking // modes (Qwen enable_thinking=false, GLM/Kimi/DeepSeek thinking.type=disabled). DisableThinking bool + // OnUsage reports the token usage a backend advertised for one generation + // call, when available. Backends that cannot report usage leave it uncalled; + // it may fire more than once (e.g. once per streamed chunk carrying usage), + // so callers should keep the last non-zero values rather than accumulate. + // Callers must keep their own estimate as a fallback for the no-usage case. + OnUsage func(promptTokens, completionTokens int64) } // DefaultOptions returns sensible defaults. MaxTokens follows Ollama and diff --git a/internal/server/handlers.go b/internal/server/handlers.go index 022c87cd..f757d1fc 100644 --- a/internal/server/handlers.go +++ b/internal/server/handlers.go @@ -866,6 +866,16 @@ func (s *Server) handleGenerate(w http.ResponseWriter, r *http.Request) { inputTokens = 1 } + var reportedIn, reportedOut int64 + opts.OnUsage = func(promptTokens, completionTokens int64) { + if promptTokens > 0 { + reportedIn = promptTokens + } + if completionTokens > 0 { + reportedOut = completionTokens + } + } + if stream { w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("Cache-Control", "no-cache") @@ -892,7 +902,7 @@ func (s *Server) handleGenerate(w http.ResponseWriter, r *http.Request) { }) return } - s.recordAPIUsage(r, req.Model, "", inputTokens, estimateAnthropicTokens(full.String())) + s.recordResolvedUsage(r, req.Model, "", reportedIn, reportedOut, inputTokens, estimateAnthropicTokens(full.String())) writeSSE(w, api.GenerateResponse{ Model: req.Model, Done: true, @@ -911,7 +921,7 @@ func (s *Server) handleGenerate(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusInternalServerError, err.Error()) return } - s.recordAPIUsage(r, req.Model, "", inputTokens, estimateAnthropicTokens(response)) + s.recordResolvedUsage(r, req.Model, "", reportedIn, reportedOut, inputTokens, estimateAnthropicTokens(response)) writeJSON(w, http.StatusOK, api.GenerateResponse{ Model: req.Model, Response: response, @@ -1021,6 +1031,16 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { return } + var reportedIn, reportedOut int64 + opts.OnUsage = func(promptTokens, completionTokens int64) { + if promptTokens > 0 { + reportedIn = promptTokens + } + if completionTokens > 0 { + reportedOut = completionTokens + } + } + if stream { if requestWantsSSE(r) { w.Header().Set("Content-Type", "text/event-stream") @@ -1076,7 +1096,7 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { return } _ = fullResp - s.recordAPIUsage(r, req.Model, req.Source, inputTokens, estimateAnthropicTokens(full.String())) + s.recordResolvedUsage(r, req.Model, req.Source, reportedIn, reportedOut, inputTokens, estimateAnthropicTokens(full.String())) writeSSE(w, api.ChatResponse{ Model: req.Model, Done: true, @@ -1137,7 +1157,7 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { }) return } - s.recordAPIUsage(r, req.Model, req.Source, inputTokens, estimateAnthropicTokens(full.String())) + s.recordResolvedUsage(r, req.Model, req.Source, reportedIn, reportedOut, inputTokens, estimateAnthropicTokens(full.String())) writeNDJSON(w, api.ChatResponse{ Model: req.Model, Message: &api.Message{ @@ -1162,7 +1182,7 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { writeInferenceError(w, err) return } - s.recordAPIUsage(r, req.Model, req.Source, inputTokens, estimateAnthropicTokens(response)) + s.recordResolvedUsage(r, req.Model, req.Source, reportedIn, reportedOut, inputTokens, estimateAnthropicTokens(response)) writeJSON(w, http.StatusOK, api.ChatResponse{ Model: req.Model, Message: &api.Message{ diff --git a/internal/server/handlers_anthropic.go b/internal/server/handlers_anthropic.go index 3a4a2275..a3535913 100644 --- a/internal/server/handlers_anthropic.go +++ b/internal/server/handlers_anthropic.go @@ -9,7 +9,6 @@ import ( "net/http" "strings" "time" - "unicode/utf8" "github.com/opencsgs/csglite/internal/inference" "github.com/opencsgs/csglite/pkg/api" @@ -90,6 +89,19 @@ func (s *Server) handleAnthropicMessages(w http.ResponseWriter, r *http.Request) return } + // Fallback for engines that do not implement ChatCompletionProxier. + // All real Engines currently implement it, so this branch is only + // reached by non-proxier fakes. + var reportedIn, reportedOut int64 + opts.OnUsage = func(promptTokens, completionTokens int64) { + if promptTokens > 0 { + reportedIn = promptTokens + } + if completionTokens > 0 { + reportedOut = completionTokens + } + } + messages := anthropicMessagesToInference(req) if req.Stream { @@ -143,7 +155,7 @@ func (s *Server) handleAnthropicMessages(w http.ResponseWriter, r *http.Request) writeAnthropicSSE(w, "message_stop", map[string]interface{}{ "type": "message_stop", }) - s.recordAPIUsage(r, req.Model, req.Source, inputTokens, outputTokens) + s.recordResolvedUsage(r, req.Model, req.Source, reportedIn, reportedOut, inputTokens, outputTokens) return } @@ -161,7 +173,7 @@ func (s *Server) handleAnthropicMessages(w http.ResponseWriter, r *http.Request) } anthropicResp := buildAnthropicMessageResponse(id, req.Model, response, inputTokens) - s.recordAPIUsage(r, req.Model, req.Source, anthropicResp.Usage.InputTokens, anthropicResp.Usage.OutputTokens) + s.recordResolvedUsage(r, req.Model, req.Source, reportedIn, reportedOut, inputTokens, estimateAnthropicTokens(response)) writeJSON(w, http.StatusOK, anthropicResp) } @@ -189,10 +201,26 @@ func (s *Server) tryNativeAnthropicMessages( copyAnthropicUpstreamHeaders(w.Header(), resp.Header) if req.Stream && resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices { w.WriteHeader(resp.StatusCode) - _, err = io.Copy(openAIStreamWriter{ResponseWriter: w}, resp.Body) - if err == nil { - s.recordAPIUsage(r, req.Model, req.Source, inputTokens, 0) + capture := &streamUsageCapture{} + _, err = io.Copy(openAIStreamWriter{ResponseWriter: w}, io.TeeReader(resp.Body, capture)) + recordInput, recordOutput := inputTokens, 0 + if capturedInput, capturedOutput, ok := capture.usage(); ok { + if capturedInput > 0 { + recordInput = capturedInput + } + if capturedOutput > 0 { + recordOutput = capturedOutput + } else { + recordOutput = estimateAnthropicTokens(extractAnthropicStreamContent(capture.tail)) + } + if capturedInput <= 0 || capturedOutput <= 0 { + r = markUsageEstimated(r) + } + } else { + recordOutput = estimateAnthropicTokens(extractAnthropicStreamContent(capture.tail)) + r = markUsageEstimated(r) } + s.recordAPIUsage(r, req.Model, req.Source, recordInput, recordOutput) return true, err } @@ -205,15 +233,31 @@ func (s *Server) tryNativeAnthropicMessages( return true, fmt.Errorf("writing Anthropic messages response: %w", err) } if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices { - var usageEnvelope struct { - Usage api.AnthropicUsage `json:"usage"` + var envelope struct { + Usage api.AnthropicUsage `json:"usage"` + Content []struct { + Type string `json:"type"` + Text string `json:"text"` + } `json:"content"` } - if json.Unmarshal(body, &usageEnvelope) == nil { - recordInput := usageEnvelope.Usage.InputTokens + if json.Unmarshal(body, &envelope) == nil { + recordInput := envelope.Usage.InputTokens if recordInput == 0 { recordInput = inputTokens + r = markUsageEstimated(r) + } + recordOutput := envelope.Usage.OutputTokens + if recordOutput == 0 { + var sb strings.Builder + for _, block := range envelope.Content { + if block.Type == "text" { + sb.WriteString(block.Text) + } + } + recordOutput = estimateAnthropicTokens(sb.String()) + r = markUsageEstimated(r) } - s.recordAPIUsage(r, req.Model, req.Source, recordInput, usageEnvelope.Usage.OutputTokens) + s.recordAPIUsage(r, req.Model, req.Source, recordInput, recordOutput) } } return true, nil @@ -318,6 +362,9 @@ func (s *Server) handleAnthropicMessagesProxy( } if !req.Stream { + if openAIResp.Usage.PromptTokens == 0 || openAIResp.Usage.CompletionTokens == 0 { + r = markUsageEstimated(r) + } s.recordAPIUsage(r, req.Model, req.Source, anthropicResp.Usage.InputTokens, anthropicResp.Usage.OutputTokens) writeJSON(w, http.StatusOK, anthropicResp) return @@ -326,6 +373,9 @@ func (s *Server) handleAnthropicMessagesProxy( w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("Cache-Control", "no-cache") w.Header().Set("Connection", "keep-alive") + if openAIResp.Usage.PromptTokens == 0 || openAIResp.Usage.CompletionTokens == 0 { + r = markUsageEstimated(r) + } s.recordAPIUsage(r, req.Model, req.Source, anthropicResp.Usage.InputTokens, anthropicResp.Usage.OutputTokens) writeAnthropicStreamedMessage(w, anthropicResp) } @@ -372,7 +422,8 @@ func (s *Server) handleAnthropicMessagesProxyStream( w.Header().Set("X-Accel-Buffering", "no") writeAnthropicMessageStart(w, id, req.Model, inputTokens) - outputTokens, err := streamOpenAIChatAsAnthropic(w, resp.Body) + reportedIn, reportedOut, estimatedOutput, err := streamOpenAIChatAsAnthropic(w, resp.Body) + s.recordResolvedUsage(r, req.Model, req.Source, int64(reportedIn), int64(reportedOut), inputTokens, estimatedOutput) if err != nil { writeAnthropicSSE(w, "error", anthropicErrorPayloadWithStatus( http.StatusBadGateway, @@ -381,7 +432,6 @@ func (s *Server) handleAnthropicMessagesProxyStream( )) return } - s.recordAPIUsage(r, req.Model, req.Source, inputTokens, outputTokens) } type anthropicStreamToolCall struct { @@ -390,7 +440,7 @@ type anthropicStreamToolCall struct { arguments strings.Builder } -func streamOpenAIChatAsAnthropic(w http.ResponseWriter, body io.Reader) (int, error) { +func streamOpenAIChatAsAnthropic(w http.ResponseWriter, body io.Reader) (promptTokens, completionTokens, estimatedOutput int, err error) { blockIndex := 0 openBlock := "" finishReason := "" @@ -447,8 +497,12 @@ func streamOpenAIChatAsAnthropic(w http.ResponseWriter, body io.Reader) (int, er outputText.WriteString(value) } - err := scanOpenAIChatStream(body, func(chunk api.OpenAIChatResponse) error { + err = scanOpenAIChatStream(body, func(chunk api.OpenAIChatResponse) error { + if chunk.Usage.PromptTokens > 0 { + promptTokens = chunk.Usage.PromptTokens + } if chunk.Usage.CompletionTokens > 0 { + completionTokens = chunk.Usage.CompletionTokens outputTokens = chunk.Usage.CompletionTokens } if len(chunk.Choices) == 0 { @@ -491,7 +545,8 @@ func streamOpenAIChatAsAnthropic(w http.ResponseWriter, body io.Reader) (int, er return nil }) if err != nil { - return outputTokens, err + estimatedOutput = estimateAnthropicTokens(outputText.String()) + return promptTokens, completionTokens, estimatedOutput, err } closeBlock() @@ -537,7 +592,8 @@ func streamOpenAIChatAsAnthropic(w http.ResponseWriter, body io.Reader) (int, er } } if outputTokens == 0 { - outputTokens = estimateAnthropicTokens(outputText.String()) + estimatedOutput = estimateAnthropicTokens(outputText.String()) + outputTokens = estimatedOutput } writeAnthropicSSE(w, "message_delta", map[string]interface{}{ "type": "message_delta", @@ -552,7 +608,7 @@ func streamOpenAIChatAsAnthropic(w http.ResponseWriter, body io.Reader) (int, er writeAnthropicSSE(w, "message_stop", map[string]interface{}{ "type": "message_stop", }) - return outputTokens, nil + return promptTokens, completionTokens, estimatedOutput, nil } func scanOpenAIChatStream(body io.Reader, onChunk func(api.OpenAIChatResponse) error) error { @@ -703,6 +759,9 @@ func (s *Server) handleAnthropicMessagesWithTools( } if !req.Stream { + if openAIResp.Usage.PromptTokens == 0 || openAIResp.Usage.CompletionTokens == 0 { + r = markUsageEstimated(r) + } s.recordAPIUsage(r, req.Model, req.Source, anthropicResp.Usage.InputTokens, anthropicResp.Usage.OutputTokens) writeJSON(w, http.StatusOK, anthropicResp) return @@ -711,6 +770,9 @@ func (s *Server) handleAnthropicMessagesWithTools( w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("Cache-Control", "no-cache") w.Header().Set("Connection", "keep-alive") + if openAIResp.Usage.PromptTokens == 0 || openAIResp.Usage.CompletionTokens == 0 { + r = markUsageEstimated(r) + } s.recordAPIUsage(r, req.Model, req.Source, anthropicResp.Usage.InputTokens, anthropicResp.Usage.OutputTokens) writeAnthropicStreamedMessage(w, anthropicResp) } @@ -840,6 +902,9 @@ func anthropicRequestToProxyBody(req api.AnthropicMessageRequest, opts inference "top_p": opts.TopP, "stream": stream, } + if stream { + body["stream_options"] = map[string]interface{}{"include_usage": true} + } if opts.MaxTokens > 0 { body["max_tokens"] = opts.MaxTokens } @@ -1082,13 +1147,35 @@ func estimateAnthropicTokens(text string) int { if text == "" { return 0 } - count := utf8.RuneCountInString(text) / 4 + cjk, other := 0, 0 + for _, r := range text { + if isCJKRune(r) { + cjk++ + } else { + other++ + } + } + // CJK-adjacent scripts (Han, Kana, Hangul, CJK punctuation) tokenize near + // 0.6 tokens/char on modern BPE tokenizers (Qwen/GLM/DeepSeek); Latin and + // other scripts stay near 4 chars/token. This is a rough fallback, not a + // tokenizer model — precision is intentionally limited. + count := cjk*3/5 + other/4 if count < 1 { count = 1 } return count } +func isCJKRune(r rune) bool { + return (r >= 0x3000 && r <= 0x303f) || // CJK Symbols and Punctuation + (r >= 0x3040 && r <= 0x30ff) || // Hiragana and Katakana + (r >= 0x3400 && r <= 0x4dbf) || // CJK Extension A + (r >= 0x4e00 && r <= 0x9fff) || // CJK Unified Ideographs + (r >= 0xac00 && r <= 0xd7af) || // Hangul Syllables + (r >= 0xf900 && r <= 0xfaff) || // CJK Compatibility Ideographs + (r >= 0xfe30 && r <= 0xfe4f) // CJK Compatibility Forms +} + func anthropicContentText(content interface{}) string { switch value := content.(type) { case nil: diff --git a/internal/server/handlers_anthropic_test.go b/internal/server/handlers_anthropic_test.go index e46d95f2..f5f2d54f 100644 --- a/internal/server/handlers_anthropic_test.go +++ b/internal/server/handlers_anthropic_test.go @@ -765,6 +765,58 @@ func TestTryNativeAnthropicMessagesRelaysSuccessfulResponse(t *testing.T) { } } +func TestTryNativeAnthropicMessagesStreamRecordsUpstreamUsage(t *testing.T) { + s := newTestServer(t) + stream := strings.Join([]string{ + `event: message_start`, + `data: {"type":"message_start","message":{"usage":{"input_tokens":120,"output_tokens":0}}}`, + ``, + `event: content_block_delta`, + `data: {"type":"content_block_delta","delta":{"type":"text_delta","text":"hi"}}`, + ``, + `event: message_delta`, + `data: {"type":"message_delta","usage":{"output_tokens":9}}`, + ``, + `event: message_stop`, + `data: {"type":"message_stop"}`, + ``, + }, "\n") + proxy := &nativeAnthropicTestProxy{ + status: http.StatusOK, + body: stream, + headers: http.Header{"Content-Type": {"text/event-stream"}}, + } + req := httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + w := httptest.NewRecorder() + + handled, err := s.tryNativeAnthropicMessages( + w, + req, + api.AnthropicMessageRequest{Model: "test/model", Source: "provider:test", Stream: true}, + map[string]interface{}{"model": "test/model"}, + proxy, + 4, + ) + if err != nil || !handled { + t.Fatalf("handled=%v err=%v, want native stream relayed", handled, err) + } + if w.Code != http.StatusOK || !strings.Contains(w.Body.String(), `"text_delta"`) { + t.Fatalf("status=%d body=%s", w.Code, w.Body.String()) + } + state, err := s.apiUsage.List(config.APIUsageListOptions{}) + if err != nil { + t.Fatal(err) + } + if len(state.Records) != 1 { + t.Fatalf("usage records = %#v", state.Records) + } + record := state.Records[0] + if record.InputTokens != 120 || record.OutputTokens != 9 || record.TotalTokens != 129 { + t.Fatalf("usage tokens = input %d output %d total %d, want 120/9/129", + record.InputTokens, record.OutputTokens, record.TotalTokens) + } +} + func TestTryNativeAnthropicMessagesFallsBackOnlyForUnsupportedEndpoint(t *testing.T) { for _, status := range []int{http.StatusNotFound, http.StatusMethodNotAllowed, http.StatusNotImplemented} { t.Run(http.StatusText(status), func(t *testing.T) { @@ -807,3 +859,123 @@ func TestTryNativeAnthropicMessagesFallsBackOnlyForUnsupportedEndpoint(t *testin t.Fatalf("handled=%v status=%d err=%v, want native authentication error relayed", handled, w.Code, err) } } + +func TestHandleAnthropicMessagesProxyStreamRecordsUpstreamUsage(t *testing.T) { + useIsolatedStorageHome(t) + engine := &fakeChatCompletionEngine{ + resp: api.OpenAIChatResponse{ + ID: "chatcmpl-anthropic-usage", + Object: "chat.completion", + Created: 123, + Model: "test/model", + Choices: []api.OpenAIChoice{{ + Index: 0, + Message: &api.Message{Role: "assistant", Content: "the answer is 42"}, + }}, + Usage: api.OpenAIUsage{PromptTokens: 120, CompletionTokens: 8, TotalTokens: 128}, + }, + } + s := newAnthropicProxyTestServer(t, engine) + + body := `{"model":"test/model","messages":[{"role":"user","content":"what is the answer to everything"}],"stream":true}` + req := httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body)) + req.Header.Set("Anthropic-Version", "2023-06-01") + w := httptest.NewRecorder() + + s.handleAnthropicMessages(w, req) + + if w.Code != http.StatusOK { + t.Fatalf("status = %d body=%s", w.Code, w.Body.String()) + } + if engine.lastReq == nil { + t.Fatal("proxy request was not made") + } + if opts, ok := engine.lastReq["stream_options"].(map[string]interface{}); !ok || opts["include_usage"] != true { + t.Fatalf("stream_options.include_usage not set in proxy request: %#v", engine.lastReq["stream_options"]) + } + state, err := s.apiUsage.List(config.APIUsageListOptions{}) + if err != nil { + t.Fatal(err) + } + if len(state.Records) != 1 { + t.Fatalf("usage records = %#v", state.Records) + } + record := state.Records[0] + if record.InputTokens != 120 || record.OutputTokens != 8 { + t.Fatalf("usage tokens = input %d output %d, want 120/8", record.InputTokens, record.OutputTokens) + } + if record.EstimatedRequests != 0 { + t.Fatalf("estimated requests = %d, want 0 (real usage should not be marked estimated)", record.EstimatedRequests) + } +} + +func TestHandleAnthropicMessagesProxyStreamFallsBackWithoutUpstreamUsage(t *testing.T) { + useIsolatedStorageHome(t) + engine := &fakeChatCompletionEngine{ + resp: api.OpenAIChatResponse{ + ID: "chatcmpl-anthropic-no-usage", + Object: "chat.completion", + Created: 123, + Model: "test/model", + Choices: []api.OpenAIChoice{{ + Index: 0, + Message: &api.Message{Role: "assistant", Content: "answer"}, + }}, + }, + } + s := newAnthropicProxyTestServer(t, engine) + + body := `{"model":"test/model","messages":[{"role":"user","content":"what is the answer to everything"}],"stream":true}` + req := httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body)) + req.Header.Set("Anthropic-Version", "2023-06-01") + w := httptest.NewRecorder() + + s.handleAnthropicMessages(w, req) + + if w.Code != http.StatusOK { + t.Fatalf("status = %d body=%s", w.Code, w.Body.String()) + } + state, err := s.apiUsage.List(config.APIUsageListOptions{}) + if err != nil { + t.Fatal(err) + } + if len(state.Records) != 1 { + t.Fatalf("usage records = %#v", state.Records) + } + record := state.Records[0] + if record.EstimatedRequests != 1 { + t.Fatalf("estimated requests = %d, want 1 (no upstream usage should be marked estimated)", record.EstimatedRequests) + } +} + +func TestHandleAnthropicMessagesProxyStreamRecordsOnMidStreamError(t *testing.T) { + useIsolatedStorageHome(t) + // Simulate a stream that sends one content chunk, then breaks. + streamBody := "data: {\"id\":\"x\",\"object\":\"chat.completion.chunk\",\"model\":\"test/model\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"partial answer here\"}}]}\n\n" + engine := &fakeChatCompletionEngine{ + streamBody: streamBody + "data: [BROKEN", + } + s := newAnthropicProxyTestServer(t, engine) + + body := `{"model":"test/model","messages":[{"role":"user","content":"hi"}],"stream":true}` + req := httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body)) + req.Header.Set("Anthropic-Version", "2023-06-01") + w := httptest.NewRecorder() + + s.handleAnthropicMessages(w, req) + + state, err := s.apiUsage.List(config.APIUsageListOptions{}) + if err != nil { + t.Fatal(err) + } + if len(state.Records) != 1 { + t.Fatalf("usage records = %#v (want 1 record even on stream error)", state.Records) + } + record := state.Records[0] + if record.EstimatedRequests != 1 { + t.Fatalf("estimated requests = %d, want 1 (mid-stream error should be marked estimated)", record.EstimatedRequests) + } + if record.OutputTokens == 0 { + t.Fatalf("output tokens = 0, want >0 (should estimate from partial stream content)") + } +} diff --git a/internal/server/handlers_api_keys.go b/internal/server/handlers_api_keys.go index 55c6a374..853b6ba1 100644 --- a/internal/server/handlers_api_keys.go +++ b/internal/server/handlers_api_keys.go @@ -307,27 +307,28 @@ func (s *Server) apiUsageRow(ctx context.Context, record config.APIUsageRecord, } } return api.APIUsageRow{ - APIKeyID: record.APIKeyID, - APIKeyName: record.APIKeyName, - Model: record.Model, - Source: source, - SourceType: sourceType, - SourceName: sourceName, - PoolID: record.PoolID, - PoolName: record.PoolName, - PoolModel: record.PoolModel, - ActualMemberID: record.ActualMemberID, - MemberModel: record.MemberModel, - EstimatedCost: record.EstimatedCost, - CostCurrency: record.CostCurrency, - CostKnown: record.CostKnown, - FallbackCount: record.FallbackCount, - LimitedCount: record.LimitedCount, - Requests: record.Requests, - InputTokens: record.InputTokens, - OutputTokens: record.OutputTokens, - TotalTokens: record.TotalTokens, - LastUsedAt: record.LastUsedAt, + APIKeyID: record.APIKeyID, + APIKeyName: record.APIKeyName, + Model: record.Model, + Source: source, + SourceType: sourceType, + SourceName: sourceName, + PoolID: record.PoolID, + PoolName: record.PoolName, + PoolModel: record.PoolModel, + ActualMemberID: record.ActualMemberID, + MemberModel: record.MemberModel, + EstimatedCost: record.EstimatedCost, + CostCurrency: record.CostCurrency, + CostKnown: record.CostKnown, + FallbackCount: record.FallbackCount, + LimitedCount: record.LimitedCount, + Requests: record.Requests, + EstimatedRequests: record.EstimatedRequests, + InputTokens: record.InputTokens, + OutputTokens: record.OutputTokens, + TotalTokens: record.TotalTokens, + LastUsedAt: record.LastUsedAt, } } diff --git a/internal/server/handlers_chat_tools.go b/internal/server/handlers_chat_tools.go index 5ac5c6e4..509ebd48 100644 --- a/internal/server/handlers_chat_tools.go +++ b/internal/server/handlers_chat_tools.go @@ -56,6 +56,11 @@ func (s *Server) handleChatWithTools(w http.ResponseWriter, r *http.Request, req inputTokens, outputTokens := openAIUsageTokens(openAIResp) if inputTokens == 0 { inputTokens = countMessageTokens(req.Messages) + r = markUsageEstimated(r) + } + if outputTokens == 0 && openAIResp.Usage.CompletionTokens == 0 { + outputTokens = estimateOpenAIOutputTokens(openAIResp) + r = markUsageEstimated(r) } s.recordAPIUsage(r, req.Model, req.Source, inputTokens, outputTokens) diff --git a/internal/server/handlers_openai.go b/internal/server/handlers_openai.go index d46c506a..b7881e75 100644 --- a/internal/server/handlers_openai.go +++ b/internal/server/handlers_openai.go @@ -97,6 +97,19 @@ func (s *Server) handleOpenAIChatCompletions(w http.ResponseWriter, r *http.Requ return } + // Fallback for engines that do not implement ChatCompletionProxier. + // All real engines (llama, openai, remote, providerPool) currently + // implement it, so this branch is only reached by non-proxier fakes. + var reportedIn, reportedOut int64 + opts.OnUsage = func(promptTokens, completionTokens int64) { + if promptTokens > 0 { + reportedIn = promptTokens + } + if completionTokens > 0 { + reportedOut = completionTokens + } + } + var messages []inference.Message for _, m := range req.Messages { messages = append(messages, inference.Message{Role: m.Role, Content: m.Content, ReasoningContent: m.ReasoningContent}) @@ -132,7 +145,7 @@ func (s *Server) handleOpenAIChatCompletions(w http.ResponseWriter, r *http.Requ writeSSE(w, apiErrorResponse{Error: err.Error(), ErrorCode: openAIInferenceStatus(err)}) return } - s.recordAPIUsage(r, req.Model, req.Source, inputTokens, estimateAnthropicTokens(full.String())) + s.recordResolvedUsage(r, req.Model, req.Source, reportedIn, reportedOut, inputTokens, estimateAnthropicTokens(full.String())) stop := "stop" writeSSE(w, api.OpenAIChatResponse{ @@ -163,7 +176,7 @@ func (s *Server) handleOpenAIChatCompletions(w http.ResponseWriter, r *http.Requ writeOpenAIInferenceError(w, err) return } - s.recordAPIUsage(r, req.Model, req.Source, inputTokens, estimateAnthropicTokens(response)) + s.recordResolvedUsage(r, req.Model, req.Source, reportedIn, reportedOut, inputTokens, estimateAnthropicTokens(response)) stop := "stop" writeJSON(w, http.StatusOK, api.OpenAIChatResponse{ @@ -244,8 +257,26 @@ func (s *Server) handleOpenAIChatCompletionsProxy( } w.WriteHeader(http.StatusOK) if stream { - _, _ = io.Copy(openAIStreamWriter{ResponseWriter: w}, resp.Body) - s.recordAPIUsageWithPool(r, req.Model, usageSource, countMessageTokens(req.Messages), 0, usagePool) + capture := &streamUsageCapture{} + _, _ = io.Copy(openAIStreamWriter{ResponseWriter: w}, io.TeeReader(resp.Body, capture)) + inputTokens, outputTokens := countMessageTokens(req.Messages), 0 + if capturedInput, capturedOutput, ok := capture.usage(); ok { + if capturedInput > 0 { + inputTokens = capturedInput + } + if capturedOutput > 0 { + outputTokens = capturedOutput + } else { + outputTokens = estimateAnthropicTokens(extractOpenAIStreamContent(capture.tail)) + } + if capturedInput <= 0 || capturedOutput <= 0 { + r = markUsageEstimated(r) + } + } else { + outputTokens = estimateAnthropicTokens(extractOpenAIStreamContent(capture.tail)) + r = markUsageEstimated(r) + } + s.recordAPIUsageWithPool(r, req.Model, usageSource, inputTokens, outputTokens, usagePool) } else { body, err := io.ReadAll(resp.Body) if err != nil { @@ -256,9 +287,15 @@ func (s *Server) handleOpenAIChatCompletionsProxy( inputTokens, outputTokens := openAIUsageTokens(openAIResp) if inputTokens == 0 { inputTokens = countMessageTokens(req.Messages) + r = markUsageEstimated(r) + } + if outputTokens == 0 && openAIResp.Usage.CompletionTokens == 0 { + outputTokens = estimateOpenAIOutputTokens(openAIResp) + r = markUsageEstimated(r) } s.recordAPIUsageWithPool(r, req.Model, usageSource, inputTokens, outputTokens, usagePool) } else { + r = markUsageEstimated(r) s.recordAPIUsageWithPool(r, req.Model, usageSource, countMessageTokens(req.Messages), 0, usagePool) } _, _ = w.Write(body) @@ -380,6 +417,11 @@ func (s *Server) handleOpenAIChatCompletionsWithTools( inputTokens, outputTokens := openAIUsageTokens(openAIResp) if inputTokens == 0 { inputTokens = countMessageTokens(req.Messages) + r = markUsageEstimated(r) + } + if outputTokens == 0 && openAIResp.Usage.CompletionTokens == 0 { + outputTokens = estimateOpenAIOutputTokens(openAIResp) + r = markUsageEstimated(r) } s.recordAPIUsage(r, req.Model, req.Source, inputTokens, outputTokens) diff --git a/internal/server/handlers_openai_test.go b/internal/server/handlers_openai_test.go index 563d05ff..7b174b95 100644 --- a/internal/server/handlers_openai_test.go +++ b/internal/server/handlers_openai_test.go @@ -23,8 +23,9 @@ import ( ) type fakeChatCompletionEngine struct { - resp api.OpenAIChatResponse - lastReq map[string]interface{} + resp api.OpenAIChatResponse + lastReq map[string]interface{} + streamBody string } type fakeNativeToolStreamingEngine struct { @@ -73,6 +74,13 @@ func (e *fakeChatCompletionEngine) ModelName() string { return "test/model" } func (e *fakeChatCompletionEngine) ChatCompletion(_ context.Context, reqBody map[string]interface{}) (*http.Response, error) { e.lastReq = reqBody if stream, _ := reqBody["stream"].(bool); stream { + if e.streamBody != "" { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(e.streamBody)), + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + }, nil + } var body strings.Builder for _, choice := range e.resp.Choices { delta := choice.Message @@ -1501,3 +1509,189 @@ func TestProviderConfiguredHeadersAreAutomaticallyAddedToClientChatRequest(t *te assertProviderHeaders(t, <-received) } + +func TestHandleOpenAIChatCompletionsProxyStreamRecordsUpstreamUsage(t *testing.T) { + useIsolatedStorageHome(t) + cfg := &config.Config{ModelDir: t.TempDir()} + if err := model.SaveManifest(cfg.ModelDir, &model.LocalModel{ + Namespace: "test", + Name: "model", + Format: model.FormatGGUF, + Size: 1, + Files: []string{"model.gguf", "config.json"}, + DownloadedAt: time.Now(), + }); err != nil { + t.Fatalf("save model manifest: %v", err) + } + modelDir := filepath.Join(cfg.ModelDir, "test", "model") + if err := os.MkdirAll(modelDir, 0o755); err != nil { + t.Fatalf("mkdir model dir: %v", err) + } + if err := os.WriteFile(filepath.Join(modelDir, "config.json"), []byte(`{"max_position_embeddings":40960}`), 0o644); err != nil { + t.Fatalf("write config.json: %v", err) + } + + engine := &fakeChatCompletionEngine{ + resp: api.OpenAIChatResponse{ + ID: "chatcmpl-usage", + Object: "chat.completion", + Created: 123, + Model: "test/model", + Choices: []api.OpenAIChoice{{ + Index: 0, + Message: &api.Message{Role: "assistant", Content: "the answer is 42"}, + }}, + Usage: api.OpenAIUsage{PromptTokens: 120, CompletionTokens: 8, TotalTokens: 128}, + }, + } + s := newTestServerWithConfig(t, cfg) + s.engines["test/model"] = &managedEngine{engine: engine, numCtx: 16384, numParallel: 4} + + body := `{"model":"test/model","messages":[{"role":"user","content":"what is the answer to everything"}],"stream":true}` + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(body)) + w := httptest.NewRecorder() + + s.handleOpenAIChatCompletions(w, req) + + if w.Code != http.StatusOK { + t.Fatalf("status = %d body=%s", w.Code, w.Body.String()) + } + if engine.lastReq == nil { + t.Fatal("proxy request was not made") + } + if opts, ok := engine.lastReq["stream_options"].(map[string]interface{}); !ok || opts["include_usage"] != true { + t.Fatalf("stream_options.include_usage not set in proxy request: %#v", engine.lastReq["stream_options"]) + } + if !strings.Contains(w.Body.String(), `"prompt_tokens":120`) { + t.Fatalf("streamed usage chunk was not forwarded: %s", w.Body.String()) + } + state, err := s.apiUsage.List(config.APIUsageListOptions{}) + if err != nil { + t.Fatal(err) + } + if len(state.Records) != 1 { + t.Fatalf("usage records = %#v", state.Records) + } + record := state.Records[0] + if record.InputTokens != 120 || record.OutputTokens != 8 || record.TotalTokens != 128 { + t.Fatalf("usage tokens = input %d output %d total %d, want 120/8/128", + record.InputTokens, record.OutputTokens, record.TotalTokens) + } + if record.EstimatedRequests != 0 { + t.Fatalf("estimated requests = %d, want 0 (real usage should not be marked estimated)", record.EstimatedRequests) + } +} + +func TestHandleOpenAIChatCompletionsProxyStreamFallsBackWithoutUpstreamUsage(t *testing.T) { + useIsolatedStorageHome(t) + cfg := &config.Config{ModelDir: t.TempDir()} + if err := model.SaveManifest(cfg.ModelDir, &model.LocalModel{ + Namespace: "test", + Name: "model", + Format: model.FormatGGUF, + Size: 1, + Files: []string{"model.gguf", "config.json"}, + DownloadedAt: time.Now(), + }); err != nil { + t.Fatalf("save model manifest: %v", err) + } + modelDir := filepath.Join(cfg.ModelDir, "test", "model") + if err := os.MkdirAll(modelDir, 0o755); err != nil { + t.Fatalf("mkdir model dir: %v", err) + } + if err := os.WriteFile(filepath.Join(modelDir, "config.json"), []byte(`{"max_position_embeddings":40960}`), 0o644); err != nil { + t.Fatalf("write config.json: %v", err) + } + + engine := &fakeChatCompletionEngine{ + resp: api.OpenAIChatResponse{ + ID: "chatcmpl-no-usage", + Object: "chat.completion", + Created: 123, + Model: "test/model", + Choices: []api.OpenAIChoice{{ + Index: 0, + Message: &api.Message{Role: "assistant", Content: "answer"}, + }}, + }, + } + s := newTestServerWithConfig(t, cfg) + s.engines["test/model"] = &managedEngine{engine: engine, numCtx: 16384, numParallel: 4} + + body := `{"model":"test/model","messages":[{"role":"user","content":"what is the answer to everything"}],"stream":true}` + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(body)) + w := httptest.NewRecorder() + + s.handleOpenAIChatCompletions(w, req) + + if w.Code != http.StatusOK { + t.Fatalf("status = %d body=%s", w.Code, w.Body.String()) + } + state, err := s.apiUsage.List(config.APIUsageListOptions{}) + if err != nil { + t.Fatal(err) + } + if len(state.Records) != 1 { + t.Fatalf("usage records = %#v", state.Records) + } + record := state.Records[0] + if want := countMessageTokens([]api.Message{{Role: "user", Content: "what is the answer to everything"}}); record.InputTokens != int64(want) { + t.Fatalf("input tokens = %d, want fallback estimate %d", record.InputTokens, want) + } + if record.OutputTokens != 1 { + t.Fatalf("output tokens = %d, want 1 (estimated from stream content)", record.OutputTokens) + } + if record.EstimatedRequests != 1 { + t.Fatalf("estimated requests = %d, want 1", record.EstimatedRequests) + } +} + +func TestHandleOpenAIChatCompletionsProxyStreamRecordsOnMidStreamError(t *testing.T) { + useIsolatedStorageHome(t) + cfg := &config.Config{ModelDir: t.TempDir()} + if err := model.SaveManifest(cfg.ModelDir, &model.LocalModel{ + Namespace: "test", + Name: "model", + Format: model.FormatGGUF, + Size: 1, + Files: []string{"model.gguf", "config.json"}, + DownloadedAt: time.Now(), + }); err != nil { + t.Fatalf("save model manifest: %v", err) + } + modelDir := filepath.Join(cfg.ModelDir, "test", "model") + if err := os.MkdirAll(modelDir, 0o755); err != nil { + t.Fatalf("mkdir model dir: %v", err) + } + if err := os.WriteFile(filepath.Join(modelDir, "config.json"), []byte(`{"max_position_embeddings":40960}`), 0o644); err != nil { + t.Fatalf("write config.json: %v", err) + } + + streamBody := "data: {\"id\":\"x\",\"object\":\"chat.completion.chunk\",\"model\":\"test/model\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"partial answer here\"}}]}\n\n" + engine := &fakeChatCompletionEngine{ + streamBody: streamBody + "data: [BROKEN", + } + s := newTestServerWithConfig(t, cfg) + s.engines["test/model"] = &managedEngine{engine: engine, numCtx: 16384, numParallel: 4} + + body := `{"model":"test/model","messages":[{"role":"user","content":"hi"}],"stream":true}` + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(body)) + w := httptest.NewRecorder() + + s.handleOpenAIChatCompletions(w, req) + + state, err := s.apiUsage.List(config.APIUsageListOptions{}) + if err != nil { + t.Fatal(err) + } + if len(state.Records) != 1 { + t.Fatalf("usage records = %#v (want 1 record even on stream error)", state.Records) + } + record := state.Records[0] + if record.EstimatedRequests != 1 { + t.Fatalf("estimated requests = %d, want 1 (mid-stream error should be marked estimated)", record.EstimatedRequests) + } + if record.OutputTokens == 0 { + t.Fatalf("output tokens = 0, want >0 (should estimate from partial stream content)") + } +} diff --git a/internal/server/handlers_responses.go b/internal/server/handlers_responses.go index 7f9f4312..ca6afa55 100644 --- a/internal/server/handlers_responses.go +++ b/internal/server/handlers_responses.go @@ -74,6 +74,19 @@ func (s *Server) handleOpenAIResponses(w http.ResponseWriter, r *http.Request) { return } + // Fallback for engines that do not implement ChatCompletionProxier. + // All real Engines currently implement it, so this branch is only + // reached by non-proxier fakes. + var reportedIn, reportedOut int64 + opts.OnUsage = func(promptTokens, completionTokens int64) { + if promptTokens > 0 { + reportedIn = promptTokens + } + if completionTokens > 0 { + reportedOut = completionTokens + } + } + messages := responsesRequestMessages(req) inputTokens := countResponsesTokens(req) id := fmt.Sprintf("resp_%d", time.Now().UnixNano()) @@ -172,7 +185,7 @@ func (s *Server) handleOpenAIResponses(w http.ResponseWriter, r *http.Request) { "type": "response.completed", "response": buildResponsesResponse(id, itemID, req.Model, text, created, "completed", inputTokens), }) - s.recordAPIUsage(r, req.Model, req.Source, inputTokens, estimateAnthropicTokens(text)) + s.recordResolvedUsage(r, req.Model, req.Source, reportedIn, reportedOut, inputTokens, estimateAnthropicTokens(text)) fmt.Fprintf(w, "event: done\ndata: [DONE]\n\n") if f, ok := w.(http.Flusher); ok { f.Flush() @@ -193,7 +206,7 @@ func (s *Server) handleOpenAIResponses(w http.ResponseWriter, r *http.Request) { return } text = normalizeResponsesVisibleText(text) - s.recordAPIUsage(r, req.Model, req.Source, inputTokens, estimateAnthropicTokens(text)) + s.recordResolvedUsage(r, req.Model, req.Source, reportedIn, reportedOut, inputTokens, estimateAnthropicTokens(text)) writeJSON(w, http.StatusOK, buildResponsesResponse(id, itemID, req.Model, text, created, "completed", inputTokens)) } @@ -247,6 +260,11 @@ func (s *Server) handleOpenAIResponsesProxy( recordInputTokens, outputTokens := openAIUsageTokens(openAIResp) if recordInputTokens == 0 { recordInputTokens = inputTokens + r = markUsageEstimated(r) + } + if outputTokens == 0 && openAIResp.Usage.CompletionTokens == 0 { + outputTokens = estimateOpenAIOutputTokens(openAIResp) + r = markUsageEstimated(r) } s.recordAPIUsage(r, req.Model, req.Source, recordInputTokens, outputTokens) if req.Stream { @@ -331,6 +349,11 @@ func (s *Server) handleOpenAIResponsesWithTools( recordInputTokens, outputTokens := openAIUsageTokens(openAIResp) if recordInputTokens == 0 { recordInputTokens = inputTokens + r = markUsageEstimated(r) + } + if outputTokens == 0 && openAIResp.Usage.CompletionTokens == 0 { + outputTokens = estimateOpenAIOutputTokens(openAIResp) + r = markUsageEstimated(r) } s.recordAPIUsage(r, req.Model, req.Source, recordInputTokens, outputTokens) diff --git a/internal/server/observability.go b/internal/server/observability.go index cc3a8308..1ad08876 100644 --- a/internal/server/observability.go +++ b/internal/server/observability.go @@ -141,16 +141,21 @@ func (w *observationResponseWriter) capture(p []byte) { } func (w *observationResponseWriter) captureUsageTail(p []byte) { + w.usageTail = appendUsageTail(w.usageTail, p) +} + +// appendUsageTail keeps only the newest observabilityUsageTailLimit bytes, the +// region where a streamed response advertises its token usage. +func appendUsageTail(tail []byte, p []byte) []byte { if len(p) >= observabilityUsageTailLimit { - w.usageTail = append(w.usageTail[:0], p[len(p)-observabilityUsageTailLimit:]...) - return + return append(tail[:0], p[len(p)-observabilityUsageTailLimit:]...) } - overflow := len(w.usageTail) + len(p) - observabilityUsageTailLimit + overflow := len(tail) + len(p) - observabilityUsageTailLimit if overflow > 0 { - copy(w.usageTail, w.usageTail[overflow:]) - w.usageTail = w.usageTail[:len(w.usageTail)-overflow] + copy(tail, tail[overflow:]) + tail = tail[:len(tail)-overflow] } - w.usageTail = append(w.usageTail, p...) + return append(tail, p...) } func (w *observationResponseWriter) Flush() { @@ -487,11 +492,11 @@ func updateObservationResponseUsage(value map[string]any, result *observationRes switch { case hasAnthropicRead: result.eligibleTokens = max(result.eligibleTokens, inputTokens+readTokens+creationTokens) -case hasTopLevelRead: - switch { - case readTokens > 0 && inputTokens < readTokens: - inputTokens += readTokens - } + case hasTopLevelRead: + switch { + case readTokens > 0 && inputTokens < readTokens: + inputTokens += readTokens + } result.eligibleTokens = max(result.eligibleTokens, inputTokens) case hasNestedRead || hasCreation: result.eligibleTokens = max(result.eligibleTokens, inputTokens) @@ -515,7 +520,6 @@ func observationJSONInt(value map[string]any, keys ...string) (int64, bool) { return 0, false } - func observationNestedJSONInt(value map[string]any, target string, parents ...string) (int64, bool) { for _, parent := range parents { nested, ok := value[parent].(map[string]any) diff --git a/internal/server/server_test.go b/internal/server/server_test.go index 4d0299db..47af15fd 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -107,6 +107,15 @@ func newTestServerWithConfig(t *testing.T, cfg *config.Config) *Server { return s } +// useIsolatedStorageHome points the storage root at a per-test directory so +// persisted state such as API usage does not leak between tests. +func useIsolatedStorageHome(t *testing.T) { + t.Helper() + home := t.TempDir() + t.Setenv("HOME", home) + t.Setenv("USERPROFILE", home) +} + func TestDisplayServerAddr(t *testing.T) { tests := map[string]string{ ":11435": "localhost:11435", diff --git a/internal/server/static/openapi/local-api.json b/internal/server/static/openapi/local-api.json index d0013282..bf7e0df3 100644 --- a/internal/server/static/openapi/local-api.json +++ b/internal/server/static/openapi/local-api.json @@ -8159,6 +8159,10 @@ "type": "integer", "minimum": 0 }, + "estimated_requests": { + "type": "integer", + "minimum": 0 + }, "input_tokens": { "type": "integer", "minimum": 0 diff --git a/internal/server/usage.go b/internal/server/usage.go index 86a4c953..a6f480e0 100644 --- a/internal/server/usage.go +++ b/internal/server/usage.go @@ -1,7 +1,9 @@ package server import ( + "bytes" "context" + "encoding/json" "math" "net/http" "strings" @@ -99,6 +101,151 @@ type apiUsagePoolMetadata struct { LimitedCount int64 } +// streamUsageCapture keeps the tail of a proxied upstream stream so the usage +// block a provider emits at the end of the stream can still be recorded after +// the bytes have been forwarded to the client. +type streamUsageCapture struct { + tail []byte +} + +func (c *streamUsageCapture) Write(p []byte) (int, error) { + c.tail = appendUsageTail(c.tail, p) + return len(p), nil +} + +// usage reports the token counts advertised by the streamed response. It +// returns false when the stream carried no usage at all, leaving the caller to +// fall back to its own estimate. +func (c *streamUsageCapture) usage() (int, int, bool) { + result := observationResponseUsageFromBodies(c.tail) + if result.inputTokens <= 0 && result.outputTokens <= 0 { + return 0, 0, false + } + return int(result.inputTokens), int(result.outputTokens), true +} + +// extractOpenAIStreamContent scans captured SSE bytes from an OpenAI streaming +// response and concatenates assistant content deltas into a single string. +func extractOpenAIStreamContent(tail []byte) string { + var sb strings.Builder + for _, line := range bytes.Split(tail, []byte("\n")) { + line = bytes.TrimSpace(line) + if !bytes.HasPrefix(line, []byte("data: ")) { + continue + } + data := bytes.TrimPrefix(line, []byte("data: ")) + if bytes.Equal(bytes.TrimSpace(data), []byte("[DONE]")) { + continue + } + var chunk struct { + Choices []struct { + Delta struct { + Content string `json:"content"` + } `json:"delta"` + } `json:"choices"` + } + if json.Unmarshal(data, &chunk) == nil { + for _, ch := range chunk.Choices { + sb.WriteString(ch.Delta.Content) + } + } + } + return sb.String() +} + +// extractAnthropicStreamContent scans captured SSE bytes from an Anthropic +// streaming response and concatenates text deltas into a single string. +func extractAnthropicStreamContent(tail []byte) string { + var sb strings.Builder + for _, line := range bytes.Split(tail, []byte("\n")) { + line = bytes.TrimSpace(line) + if !bytes.HasPrefix(line, []byte("data: ")) { + continue + } + data := bytes.TrimPrefix(line, []byte("data: ")) + var chunk struct { + Type string `json:"type"` + Delta struct { + Type string `json:"type"` + Text string `json:"text"` + } `json:"delta"` + } + if json.Unmarshal(data, &chunk) == nil { + if chunk.Type == "content_block_delta" && chunk.Delta.Text != "" { + sb.WriteString(chunk.Delta.Text) + } + } + } + return sb.String() +} + +// sanePrompt returns true when a prompt-token count is positive and below a +// generous single-request ceiling. A zero or negative prompt means the backend +// did not report it, so the caller should fall back to its own estimate. +func sanePrompt(prompt int64) bool { + const maxReasonableTokens = 10_000_000 + return prompt > 0 && prompt <= maxReasonableTokens +} + +// saneCompletion returns true when a completion-token count is positive and +// below a generous single-request ceiling. A zero completion is treated as +// "not reported" rather than "real zero output": the OnUsage callback only +// stores values > 0, so a zero reaching resolveUsageTokens means the backend +// omitted completion_tokens. Falling back to the caller's text-based estimate +// is safer than trusting the missing value as real. When the model truly +// produced no output the estimate is also 0, so the recorded value is the same +// — only the "estimated" tag differs, which is an acceptable trade-off. +func saneCompletion(completion int64) bool { + const maxReasonableTokens = 10_000_000 + return completion > 0 && completion <= maxReasonableTokens +} + +// resolveUsageTokens prefers the token counts a backend actually reported and +// falls back to the caller's estimate per side: when the backend reports only +// one of prompt/completion, the other side uses the estimate and the result is +// marked as not fully real. The returned bool is true only when both sides +// came from the backend. +func resolveUsageTokens(realIn, realOut int64, estIn, estOut int) (in, out int, real bool) { + inSane := sanePrompt(realIn) + outSane := saneCompletion(realOut) + if inSane && outSane { + return int(realIn), int(realOut), true + } + in, out = estIn, estOut + if inSane { + in = int(realIn) + } + if outSane { + out = int(realOut) + } + return in, out, false +} + +type usageEstimatedContextKey struct{} + +// markUsageEstimated tags the request so recordAPIUsage counts the recorded +// values as estimated rather than real. Used at sites that fell back to a +// heuristic because the backend reported no usable usage. +func markUsageEstimated(r *http.Request) *http.Request { + return r.WithContext(context.WithValue(r.Context(), usageEstimatedContextKey{}, true)) +} + +func usageEstimatedFromContext(ctx context.Context) bool { + v, _ := ctx.Value(usageEstimatedContextKey{}).(bool) + return v +} + +// recordResolvedUsage records token usage, preferring backend-reported values +// over the caller's estimate and tagging the record as estimated when the +// backend reported nothing usable. +func (s *Server) recordResolvedUsage(r *http.Request, model, source string, realIn, realOut int64, estIn, estOut int) { + in, out, real := resolveUsageTokens(realIn, realOut, estIn, estOut) + if !real { + r = markUsageEstimated(r) + } + s.recordAPIUsage(r, model, source, in, out) +} + func (s *Server) recordAPIUsage(r *http.Request, model, source string, inputTokens, outputTokens int) { memberSource, pool := providerPoolUsageCaptureFromContext(r.Context()).get() if memberSource != "" { @@ -159,6 +306,7 @@ func (s *Server) recordAPIUsageWithPool(r *http.Request, model, source string, i LimitedCount: poolMetadataCount(pool, func(value *apiUsagePoolMetadata) int64 { return value.LimitedCount }), InputTokens: int64(inputTokens), OutputTokens: int64(outputTokens), + Estimated: usageEstimatedFromContext(r.Context()), }) } @@ -271,9 +419,17 @@ func openAIUsageTokens(resp api.OpenAIChatResponse) (int, int) { if resp.Usage.TotalTokens > 0 || resp.Usage.PromptTokens > 0 || resp.Usage.CompletionTokens > 0 { return resp.Usage.PromptTokens, resp.Usage.CompletionTokens } + return 0, estimateOpenAIOutputTokens(resp) +} + +// estimateOpenAIOutputTokens extracts assistant text from a non-streaming +// OpenAI chat response and returns an estimated output token count. Used at +// sites that need to estimate output when the upstream reported prompt tokens +// but omitted completion_tokens. +func estimateOpenAIOutputTokens(resp api.OpenAIChatResponse) int { output := "" if len(resp.Choices) > 0 && resp.Choices[0].Message != nil { output = contentAsString(resp.Choices[0].Message.Content) } - return 0, estimateAnthropicTokens(output) + return estimateAnthropicTokens(output) } diff --git a/openapi/local-api.json b/openapi/local-api.json index d0013282..bf7e0df3 100644 --- a/openapi/local-api.json +++ b/openapi/local-api.json @@ -8159,6 +8159,10 @@ "type": "integer", "minimum": 0 }, + "estimated_requests": { + "type": "integer", + "minimum": 0 + }, "input_tokens": { "type": "integer", "minimum": 0 diff --git a/pkg/api/types.go b/pkg/api/types.go index d3c4a539..9418633e 100644 --- a/pkg/api/types.go +++ b/pkg/api/types.go @@ -776,27 +776,28 @@ type APIUsageTotals struct { } type APIUsageRow struct { - APIKeyID string `json:"api_key_id"` - APIKeyName string `json:"api_key_name"` - Model string `json:"model"` - Source string `json:"source"` - SourceType string `json:"source_type"` - SourceName string `json:"source_name,omitempty"` - PoolID string `json:"pool_id,omitempty"` - PoolName string `json:"pool_name,omitempty"` - PoolModel string `json:"pool_model,omitempty"` - ActualMemberID string `json:"actual_member_id,omitempty"` - MemberModel string `json:"member_model,omitempty"` - EstimatedCost float64 `json:"estimated_cost,omitempty"` - CostCurrency string `json:"cost_currency,omitempty"` - CostKnown bool `json:"cost_known"` - FallbackCount int64 `json:"fallback_count,omitempty"` - LimitedCount int64 `json:"limited_count,omitempty"` - Requests int64 `json:"requests"` - InputTokens int64 `json:"input_tokens"` - OutputTokens int64 `json:"output_tokens"` - TotalTokens int64 `json:"total_tokens"` - LastUsedAt time.Time `json:"last_used_at"` + APIKeyID string `json:"api_key_id"` + APIKeyName string `json:"api_key_name"` + Model string `json:"model"` + Source string `json:"source"` + SourceType string `json:"source_type"` + SourceName string `json:"source_name,omitempty"` + PoolID string `json:"pool_id,omitempty"` + PoolName string `json:"pool_name,omitempty"` + PoolModel string `json:"pool_model,omitempty"` + ActualMemberID string `json:"actual_member_id,omitempty"` + MemberModel string `json:"member_model,omitempty"` + EstimatedCost float64 `json:"estimated_cost,omitempty"` + CostCurrency string `json:"cost_currency,omitempty"` + CostKnown bool `json:"cost_known"` + FallbackCount int64 `json:"fallback_count,omitempty"` + LimitedCount int64 `json:"limited_count,omitempty"` + Requests int64 `json:"requests"` + EstimatedRequests int64 `json:"estimated_requests,omitempty"` + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + TotalTokens int64 `json:"total_tokens"` + LastUsedAt time.Time `json:"last_used_at"` } type APIUsageSourceTotal struct { @@ -1415,29 +1416,29 @@ type ProviderPoolRouterBaselines struct { } type ProviderPoolRouterMetrics struct { - QueryCount int `json:"query_count"` - CellCount int `json:"cell_count"` - TrialCount int `json:"trial_count"` - Repeats int `json:"repeats"` - ResponseOutcomes map[string]int `json:"response_outcomes"` - WinRate float64 `json:"win_rate"` - Spend float64 `json:"spend"` - TotalCost float64 `json:"total_cost"` - Currency string `json:"currency,omitempty"` - CostUnit string `json:"cost_unit"` - MonetarySpendKnown bool `json:"monetary_spend_known"` - UnknownMonetarySpend bool `json:"unknown_monetary_spend"` - TrainQueryCount int `json:"train_query_count"` - HeldOutQueryCount int `json:"held_out_query_count"` - CVFoldCount int `json:"cv_fold_count"` - TrainUtility float64 `json:"train_utility"` - TrainQuality float64 `json:"train_quality"` - TrainCost float64 `json:"train_cost_score"` - HeldOutUtility float64 `json:"held_out_utility"` - HeldOutQuality float64 `json:"held_out_quality"` - HeldOutCost float64 `json:"held_out_cost_score"` - AllClustersOneMember bool `json:"all_clusters_one_member"` - SemanticDifferentiation bool `json:"semantic_differentiation"` + QueryCount int `json:"query_count"` + CellCount int `json:"cell_count"` + TrialCount int `json:"trial_count"` + Repeats int `json:"repeats"` + ResponseOutcomes map[string]int `json:"response_outcomes"` + WinRate float64 `json:"win_rate"` + Spend float64 `json:"spend"` + TotalCost float64 `json:"total_cost"` + Currency string `json:"currency,omitempty"` + CostUnit string `json:"cost_unit"` + MonetarySpendKnown bool `json:"monetary_spend_known"` + UnknownMonetarySpend bool `json:"unknown_monetary_spend"` + TrainQueryCount int `json:"train_query_count"` + HeldOutQueryCount int `json:"held_out_query_count"` + CVFoldCount int `json:"cv_fold_count"` + TrainUtility float64 `json:"train_utility"` + TrainQuality float64 `json:"train_quality"` + TrainCost float64 `json:"train_cost_score"` + HeldOutUtility float64 `json:"held_out_utility"` + HeldOutQuality float64 `json:"held_out_quality"` + HeldOutCost float64 `json:"held_out_cost_score"` + AllClustersOneMember bool `json:"all_clusters_one_member"` + SemanticDifferentiation bool `json:"semantic_differentiation"` Baselines ProviderPoolRouterBaselines `json:"baselines"` } diff --git a/web/src/api/client.ts b/web/src/api/client.ts index edcc667c..bc32cf2e 100644 --- a/web/src/api/client.ts +++ b/web/src/api/client.ts @@ -516,6 +516,7 @@ export interface LocalAPIUsageRow { fallback_count?: number; limited_count?: number; requests: number; + estimated_requests?: number; input_tokens: number; output_tokens: number; total_tokens: number; diff --git a/web/src/i18n.ts b/web/src/i18n.ts index a9ef3fcd..8700550e 100644 --- a/web/src/i18n.ts +++ b/web/src/i18n.ts @@ -1153,6 +1153,7 @@ export const en: Record = { "settings.apiUsagePeriodMonth": "Last 30 days", "settings.apiUsagePeriodYear": "Last 12 months", "settings.apiUsageRequests": "Requests", + "settings.apiUsageEstimated": "estimated", "settings.apiUsagePoolRequests": "Pool requests", "settings.apiUsageFallbacks": "Fallbacks", "settings.apiUsageLimited": "Limited attempts", @@ -2877,6 +2878,7 @@ export const zh: Record = { "settings.apiUsagePeriodMonth": "最近 30 天", "settings.apiUsagePeriodYear": "最近 12 个月", "settings.apiUsageRequests": "请求数", + "settings.apiUsageEstimated": "估算", "settings.apiUsagePoolRequests": "池请求", "settings.apiUsageFallbacks": "回退次数", "settings.apiUsageLimited": "限流次数", diff --git a/web/src/pages/AIGateway.tsx b/web/src/pages/AIGateway.tsx index 6205823c..676bbdbe 100644 --- a/web/src/pages/AIGateway.tsx +++ b/web/src/pages/AIGateway.tsx @@ -2193,7 +2193,12 @@ function UsageStatisticsSection() { {row.member_model ? `${row.model} → ${row.member_model}` : row.model} - {formatNumber(row.requests)} + + {formatNumber(row.requests)} + {row.estimated_requests ? ( + {t("settings.apiUsageEstimated")} + ) : null} + {formatNumber(row.input_tokens)} {formatNumber(row.output_tokens)} {formatNumber(row.total_tokens)} @@ -2266,7 +2271,12 @@ function UsageKeyBreakdown({ usage }: { usage: LocalAPIUsageResponse | null }) { {row.member_model ? `${row.model} → ${row.member_model}` : row.model} {apiUsageSourceRowLabel(row.source_type, row.source_name, row.pool_name)} - {formatNumber(row.requests)} + + {formatNumber(row.requests)} + {row.estimated_requests ? ( + {t("settings.apiUsageEstimated")} + ) : null} + — {formatNumber(row.input_tokens)} {formatNumber(row.output_tokens)}