Files
CLIProxyAPI/internal/runtime/executor/helps/usage_helpers.go
hkfires 6a489fa84d fix(auth): prefer errors from upstream attempts
Track when executor calls cross an upstream transport boundary and use that
signal to keep model/provider errors from being replaced by later local
preparation, selection, or internal failures.

Mark HTTP, websocket, relay, and usage-tracked transports as upstream
attempts, while avoiding marks for local validation, logging, missing
sessions, and successful websocket handshakes before request send.

Parse relative auth expiry metadata and adjust Antigravity refresh timing.
2026-08-29 12:50:46 +08:00

1288 lines
36 KiB
Go

package helps
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"net/http"
"reflect"
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
"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"
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"
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
)
type UsageReporter struct {
provider string
executorType string
model string
alias string
authID string
authIndex string
authMu sync.RWMutex
accessTokenHash string
authType string
apiKey string
source string
reasoning string
serviceTier string
generate bool
requestedAt time.Time
ttftMu sync.RWMutex
ttft time.Duration
firstPacketDuration time.Duration
firstPacketSet bool
ttftStart time.Time
ttftSet bool
once sync.Once
}
type usageExecutor interface {
Identifier() string
}
func NewExecutorUsageReporter(ctx context.Context, executor usageExecutor, model string, auth *cliproxyauth.Auth) *UsageReporter {
provider := ""
if executor != nil {
provider = executor.Identifier()
}
reporter := NewUsageReporter(ctx, provider, model, auth)
reporter.executorType = ExecutorTypeName(executor)
return reporter
}
func NewUsageReporter(ctx context.Context, provider, model string, auth *cliproxyauth.Auth) *UsageReporter {
apiKey := APIKeyFromContext(ctx)
alias := usage.RequestedModelAliasFromContext(ctx)
if alias == "" {
alias = model
}
reporter := &UsageReporter{
provider: provider,
model: model,
alias: strings.TrimSpace(alias),
requestedAt: time.Now(),
apiKey: apiKey,
source: resolveUsageSource(auth, apiKey),
authType: resolveUsageAuthType(auth),
reasoning: usage.ReasoningEffortFromContext(ctx),
serviceTier: usage.ServiceTierFromContext(ctx),
generate: usage.GenerateFromContext(ctx),
}
if auth != nil {
reporter.authID = auth.ID
reporter.authIndex = auth.EnsureIndex()
reporter.accessTokenHash = authAccessTokenSHA256(auth)
}
return reporter
}
// UpdateAccessTokenFingerprint records the token version actually used upstream.
func (r *UsageReporter) UpdateAccessTokenFingerprint(auth *cliproxyauth.Auth) {
if r == nil {
return
}
r.authMu.Lock()
r.accessTokenHash = authAccessTokenSHA256(auth)
r.authMu.Unlock()
}
func (r *UsageReporter) accessTokenFingerprint() string {
if r == nil {
return ""
}
r.authMu.RLock()
defer r.authMu.RUnlock()
return r.accessTokenHash
}
func ExecutorTypeName(executor any) string {
if executor == nil {
return ""
}
executorType := reflect.TypeOf(executor)
for executorType.Kind() == reflect.Pointer {
executorType = executorType.Elem()
}
return strings.TrimSpace(executorType.Name())
}
func (r *UsageReporter) Publish(ctx context.Context, detail usage.Detail) {
r.publishWithOutcome(ctx, detail, false, usage.Failure{})
}
func (r *UsageReporter) PublishAdditionalModel(ctx context.Context, model string, detail usage.Detail) {
record, ok := r.buildAdditionalModelRecord(model, detail)
if !ok {
return
}
r.publishRecord(ctx, record)
}
func (r *UsageReporter) SetTranslatedReasoningEffort(payload []byte, format string) {
if r == nil {
return
}
r.reasoning = thinking.ExtractTranslatedReasoningEffort(payload, format)
}
func (r *UsageReporter) TrackHTTPClient(client *http.Client) *http.Client {
return r.trackHTTPClient(client, false)
}
// TrackHTTPClientRoundTripOnly records the TTFT start time upon sending the request
// and captures first-packet arrival fallback on initial body reads, while keeping
// effective TTFT unset. This allows protocol-aware streaming executors (like Codex SSE)
// to mark effective TTFT explicitly upon receiving substantive token events, while
// preserving first-packet fallback metrics for non-2xx error bodies or keepalive streams.
func (r *UsageReporter) TrackHTTPClientRoundTripOnly(client *http.Client) *http.Client {
return r.trackHTTPClient(client, true)
}
func (r *UsageReporter) trackHTTPClient(client *http.Client, packetOnly bool) *http.Client {
if r == nil || client == nil {
return client
}
tracked := *client
transport := tracked.Transport
if transport == nil {
transport = http.DefaultTransport
}
tracked.Transport = usageTTFTRoundTripper{
base: transport,
reporter: r,
packetOnly: packetOnly,
}
return &tracked
}
func (r *UsageReporter) ObserveResponse(resp *http.Response) {
if r == nil || resp == nil || resp.Body == nil {
return
}
r.StartResponseTTFT()
resp.Body = &usageTTFTReadCloser{
ReadCloser: resp.Body,
mark: func() {
r.MarkFirstResponseByte()
},
}
}
func (r *UsageReporter) ObserveResponsePacketOnly(resp *http.Response) {
if r == nil || resp == nil || resp.Body == nil {
return
}
r.StartResponseTTFT()
resp.Body = &usageTTFTReadCloser{
ReadCloser: resp.Body,
mark: func() {
r.RecordFirstPacket()
},
}
}
func (r *UsageReporter) StartResponseTTFT() {
if r == nil {
return
}
r.ttftMu.Lock()
if !r.ttftSet && r.ttftStart.IsZero() {
r.ttftStart = time.Now()
}
r.ttftMu.Unlock()
}
func (r *UsageReporter) IsTTFTSet() bool {
if r == nil {
return false
}
r.ttftMu.RLock()
defer r.ttftMu.RUnlock()
return r.ttftSet
}
func (r *UsageReporter) IsFirstPacketSet() bool {
if r == nil {
return false
}
r.ttftMu.RLock()
defer r.ttftMu.RUnlock()
return r.firstPacketSet
}
// RecordFirstPacket records the arrival time of the first packet/chunk from upstream as a fallback.
func (r *UsageReporter) RecordFirstPacket() {
r.ObserveTokenEvent(false)
}
// ObserveTokenEvent records the first packet fallback time on the first frame,
// and if isToken is true, records the effective TTFT. It uses a fast-path
// read check to return immediately with zero write lock contention once TTFT is set
// or when subsequent non-token metadata frames arrive after the first packet.
func (r *UsageReporter) ObserveTokenEvent(isToken bool) {
if r == nil {
return
}
r.ttftMu.RLock()
if r.ttftSet {
r.ttftMu.RUnlock()
return
}
start := r.ttftStart
alreadyRecordedPacket := r.firstPacketSet
r.ttftMu.RUnlock()
if start.IsZero() {
return
}
if !isToken && alreadyRecordedPacket {
return
}
r.ttftMu.Lock()
defer r.ttftMu.Unlock()
if r.ttftSet {
return
}
if !r.firstPacketSet {
r.firstPacketDuration = time.Since(start)
r.firstPacketSet = true
}
if isToken {
r.ttft = time.Since(start)
r.ttftSet = true
r.ttftStart = time.Time{}
}
}
func (r *UsageReporter) MarkFirstResponseByte() {
if r == nil {
return
}
r.ttftMu.Lock()
if r.ttftSet {
r.ttftMu.Unlock()
return
}
start := r.ttftStart
r.ttftStart = time.Time{}
r.ttftMu.Unlock()
if start.IsZero() {
return
}
r.setTTFT(time.Since(start))
}
func (r *UsageReporter) buildAdditionalModelRecord(model string, detail usage.Detail) (usage.Record, bool) {
if r == nil {
return usage.Record{}, false
}
model = strings.TrimSpace(model)
if model == "" {
return usage.Record{}, false
}
detail = normalizeUsageDetailTotal(detail, r.provider, r.executorType)
if !hasNonZeroTokenUsage(detail) {
return usage.Record{}, false
}
return r.buildRecordForModel(model, detail, false, usage.Failure{}), true
}
func (r *UsageReporter) PublishFailure(ctx context.Context, errs ...error) {
r.publishWithOutcome(ctx, usage.Detail{}, true, failFromErrors(errs...))
}
func (r *UsageReporter) PublishFailureWithDetail(ctx context.Context, detail usage.Detail, errs ...error) {
r.publishWithOutcome(ctx, detail, true, failFromErrors(errs...))
}
func (r *UsageReporter) TrackFailure(ctx context.Context, errPtr *error) {
if r == nil || errPtr == nil {
return
}
if *errPtr != nil {
r.PublishFailure(ctx, *errPtr)
}
}
func (r *UsageReporter) publishWithOutcome(ctx context.Context, detail usage.Detail, failed bool, fail usage.Failure) {
if r == nil {
return
}
detail = normalizeUsageDetailTotal(detail, r.provider, r.executorType)
r.once.Do(func() {
r.publishRecord(ctx, r.buildRecord(detail, failed, fail))
})
}
func normalizeUsageDetailTotal(detail usage.Detail, provider, executorType string) usage.Detail {
return usage.EnsureTokenBreakdownForProvider(detail, provider, executorType)
}
func hasNonZeroTokenUsage(detail usage.Detail) bool {
return detail.InputTokens != 0 ||
detail.OutputTokens != 0 ||
detail.ReasoningTokens != 0 ||
detail.CachedTokens != 0 ||
detail.CacheReadTokens != 0 ||
detail.CacheCreationTokens != 0 ||
detail.TotalTokens != 0 ||
detail.TokenBreakdown.TotalTokens != 0
}
// ensurePublished guarantees that a usage record is emitted exactly once.
// It is safe to call multiple times; only the first call wins due to once.Do.
// This is used to ensure request counting even when upstream responses do not
// include any usage fields (tokens), especially for streaming paths.
func (r *UsageReporter) EnsurePublished(ctx context.Context) {
if r == nil {
return
}
r.once.Do(func() {
r.publishRecord(ctx, r.buildRecord(usage.Detail{}, false, usage.Failure{}))
})
}
func (r *UsageReporter) publishRecord(ctx context.Context, record usage.Record) {
record.ResponseHeaders = internallogging.GetResponseHeaders(ctx)
usage.PublishRecord(ctx, record)
}
func (r *UsageReporter) buildRecord(detail usage.Detail, failed bool, failures ...usage.Failure) usage.Record {
var fail usage.Failure
if len(failures) > 0 {
fail = failures[0]
}
if r == nil {
return usage.Record{Detail: detail, Failed: failed, Fail: fail, Generate: usage.GenerateFlag(true)}
}
return r.buildRecordForModel(r.model, detail, failed, fail)
}
func (r *UsageReporter) buildRecordForModel(model string, detail usage.Detail, failed bool, fail usage.Failure) usage.Record {
if r == nil {
return usage.Record{Model: model, Detail: detail, Failed: failed, Fail: fail, Generate: usage.GenerateFlag(true)}
}
return usage.Record{
Provider: r.provider,
ExecutorType: r.executorType,
Model: model,
Alias: r.alias,
Source: r.source,
APIKey: r.apiKey,
AuthID: r.authID,
AuthIndex: r.authIndex,
AccessTokenSHA256: r.accessTokenFingerprint(),
AuthType: r.authType,
ReasoningEffort: r.reasoning,
ServiceTier: r.serviceTier,
ResponseServiceTier: strings.TrimSpace(detail.ResponseServiceTier),
Generate: usage.GenerateFlag(r.generate),
RequestedAt: r.requestedAt,
Latency: r.latency(),
TTFT: r.ttftDuration(),
Failed: failed,
Fail: fail,
Detail: detail,
}
}
func failFromErrors(errs ...error) usage.Failure {
for _, err := range errs {
if err == nil {
continue
}
body := strings.TrimSpace(err.Error())
type responseBodyProvider interface {
ResponseBody() []byte
}
var responseErr responseBodyProvider
if errors.As(err, &responseErr) && responseErr != nil {
if responseBody := responseErr.ResponseBody(); len(responseBody) > 0 {
body = string(responseBody)
}
}
return usage.Failure{
Body: body,
StatusCode: clienterror.HTTPStatusFromError(err),
}
}
return usage.Failure{}
}
func (r *UsageReporter) latency() time.Duration {
if r == nil || r.requestedAt.IsZero() {
return 0
}
latency := time.Since(r.requestedAt)
if latency < 0 {
return 0
}
return latency
}
func (r *UsageReporter) setTTFT(ttft time.Duration) {
if r == nil {
return
}
if ttft < 0 {
ttft = 0
}
r.ttftMu.Lock()
if r.ttftSet {
r.ttftMu.Unlock()
return
}
r.ttft = ttft
r.ttftSet = true
r.ttftStart = time.Time{}
r.ttftMu.Unlock()
}
func (r *UsageReporter) ttftDuration() time.Duration {
if r == nil {
return 0
}
r.ttftMu.RLock()
defer r.ttftMu.RUnlock()
if r.ttftSet {
return r.ttft
}
if r.firstPacketSet {
return r.firstPacketDuration
}
return 0
}
type usageTTFTRoundTripper struct {
base http.RoundTripper
reporter *UsageReporter
packetOnly bool
}
func (t usageTTFTRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
cliproxyexecutor.MarkUpstreamAttempt(req.Context())
t.reporter.StartResponseTTFT()
resp, errRoundTrip := t.base.RoundTrip(req)
if errRoundTrip != nil {
return resp, errRoundTrip
}
if t.packetOnly {
t.reporter.ObserveResponsePacketOnly(resp)
} else {
t.reporter.ObserveResponse(resp)
}
return resp, nil
}
type usageTTFTReadCloser struct {
io.ReadCloser
once sync.Once
mark func()
}
func (r *usageTTFTReadCloser) Read(p []byte) (int, error) {
if r == nil || r.ReadCloser == nil {
return 0, io.ErrClosedPipe
}
n, errRead := r.ReadCloser.Read(p)
if n > 0 && r.mark != nil {
r.once.Do(r.mark)
}
return n, errRead
}
func APIKeyFromContext(ctx context.Context) string {
if ctx == nil {
return ""
}
ginCtx, ok := ctx.Value("gin").(*gin.Context)
if !ok || ginCtx == nil {
return ""
}
if v, exists := ginCtx.Get("userApiKey"); exists {
switch value := v.(type) {
case string:
return value
case fmt.Stringer:
return value.String()
default:
return fmt.Sprintf("%v", value)
}
}
return ""
}
func resolveUsageSource(auth *cliproxyauth.Auth, ctxAPIKey string) string {
if auth != nil {
provider := strings.TrimSpace(auth.Provider)
if strings.EqualFold(provider, "vertex") {
if auth.Metadata != nil {
if projectID, ok := auth.Metadata["project_id"].(string); ok {
if trimmed := strings.TrimSpace(projectID); trimmed != "" {
return trimmed
}
}
if project, ok := auth.Metadata["project"].(string); ok {
if trimmed := strings.TrimSpace(project); trimmed != "" {
return trimmed
}
}
}
}
if _, value := auth.AccountInfo(); value != "" {
return strings.TrimSpace(value)
}
if auth.Metadata != nil {
if email, ok := auth.Metadata["email"].(string); ok {
if trimmed := strings.TrimSpace(email); trimmed != "" {
return trimmed
}
}
}
if auth.Attributes != nil {
if key := strings.TrimSpace(auth.Attributes["api_key"]); key != "" {
return key
}
}
}
if trimmed := strings.TrimSpace(ctxAPIKey); trimmed != "" {
return trimmed
}
return ""
}
func resolveUsageAuthType(auth *cliproxyauth.Auth) string {
if auth == nil {
return ""
}
return auth.AuthKind()
}
// StreamUsageBuffer keeps the latest usage detail observed in a stream.
type StreamUsageBuffer struct {
detail usage.Detail
ok bool
}
var (
openAIStreamUsageMarker = []byte(`"usage"`)
openAIStreamServiceTierMarker = []byte(`"service_tier"`)
)
// Observe records detail when ok is true, allowing the final stream usage to win.
func (b *StreamUsageBuffer) Observe(detail usage.Detail, ok bool) {
if b == nil || !ok {
return
}
responseServiceTier := strings.TrimSpace(detail.ResponseServiceTier)
if responseServiceTier == "" || hasNonZeroTokenUsage(detail) {
preservedTier := b.detail.ResponseServiceTier
b.detail = detail
if b.detail.ResponseServiceTier == "" {
b.detail.ResponseServiceTier = preservedTier
}
} else {
b.detail.ResponseServiceTier = responseServiceTier
}
b.ok = true
}
// 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
}
payload := jsonPayload(line)
if len(payload) == 0 {
return
}
hasUsageCandidate := bytes.Contains(payload, openAIStreamUsageMarker)
needTier := b.detail.ResponseServiceTier == "" || hasUsageCandidate
hasTierCandidate := needTier && bytes.Contains(payload, openAIStreamServiceTierMarker)
if !hasUsageCandidate && !hasTierCandidate {
return
}
if !gjson.ValidBytes(payload) {
return
}
detail := usage.Detail{}
usageOK := false
if hasUsageCandidate {
usageNode := gjson.GetBytes(payload, "usage")
if hasOpenAIStyleUsageTokenFields(usageNode) {
detail = parseOpenAIStyleUsageNode(usageNode)
usageOK = true
}
}
if hasTierCandidate {
detail.ResponseServiceTier = extractResponseServiceTierFromValidJSON(payload)
}
b.Observe(detail, usageOK || detail.ResponseServiceTier != "")
}
// Publish emits the latest observed usage detail, if any.
func (b *StreamUsageBuffer) Publish(ctx context.Context, reporter *UsageReporter) bool {
if b == nil || !b.ok || reporter == nil {
return false
}
reporter.Publish(ctx, b.detail)
return true
}
// 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 {
return false
}
reporter.PublishFailureWithDetail(ctx, b.detail, errs...)
return true
}
// Detail returns the latest observed usage detail.
func (b *StreamUsageBuffer) Detail() (usage.Detail, bool) {
if b == nil || !b.ok {
return usage.Detail{}, false
}
return b.detail, true
}
func ParseCodexUsage(data []byte) (usage.Detail, bool) {
responseServiceTier := extractResponseServiceTier(data)
usageNode := gjson.ParseBytes(data).Get("response.usage")
if !hasOpenAIStyleUsageTokenFields(usageNode) {
if responseServiceTier == "" {
return usage.Detail{}, false
}
return usage.Detail{ResponseServiceTier: responseServiceTier}, true
}
detail := parseOpenAIStyleUsageNode(usageNode)
detail.ResponseServiceTier = responseServiceTier
return detail, true
}
func ParseCodexImageToolUsage(data []byte) (usage.Detail, bool) {
usageNode := gjson.ParseBytes(data).Get("response.tool_usage.image_gen")
if !hasOpenAIStyleUsageTokenFields(usageNode) {
return usage.Detail{}, false
}
return parseOpenAIStyleUsageNode(usageNode), true
}
func ParseOpenAIUsage(data []byte) usage.Detail {
responseServiceTier := extractResponseServiceTier(data)
usageNode := gjson.ParseBytes(data).Get("usage")
if !hasOpenAIStyleUsageTokenFields(usageNode) {
return usage.Detail{ResponseServiceTier: responseServiceTier}
}
detail := parseOpenAIStyleUsageNode(usageNode)
detail.ResponseServiceTier = responseServiceTier
return detail
}
func hasOpenAIStyleUsageTokenFields(usageNode gjson.Result) bool {
if !usageNode.Exists() || !usageNode.IsObject() {
return false
}
return usageNode.Get("total_tokens").Exists() || hasOpenAIStyleUsageBucketFields(usageNode)
}
func hasOpenAIStyleUsageBucketFields(usageNode gjson.Result) bool {
return usageNode.Get("prompt_tokens").Exists() ||
usageNode.Get("input_tokens").Exists() ||
usageNode.Get("completion_tokens").Exists() ||
usageNode.Get("output_tokens").Exists() ||
usageNode.Get("prompt_tokens_details.cached_tokens").Exists() ||
usageNode.Get("input_tokens_details.cached_tokens").Exists() ||
usageNode.Get("prompt_tokens_details.cache_write_tokens").Exists() ||
usageNode.Get("prompt_tokens_details.cache_creation_tokens").Exists() ||
usageNode.Get("input_tokens_details.cache_write_tokens").Exists() ||
usageNode.Get("input_tokens_details.cache_creation_tokens").Exists() ||
usageNode.Get("completion_tokens_details.reasoning_tokens").Exists() ||
usageNode.Get("output_tokens_details.reasoning_tokens").Exists()
}
func parseOpenAIStyleUsageNode(usageNode gjson.Result) usage.Detail {
inputNode := usageNode.Get("prompt_tokens")
if !inputNode.Exists() {
inputNode = usageNode.Get("input_tokens")
}
outputNode := usageNode.Get("completion_tokens")
if !outputNode.Exists() {
outputNode = usageNode.Get("output_tokens")
}
detail := usage.Detail{
InputTokens: inputNode.Int(),
OutputTokens: outputNode.Int(),
TotalTokens: usageNode.Get("total_tokens").Int(),
}
cached := usageNode.Get("prompt_tokens_details.cached_tokens")
if !cached.Exists() {
cached = usageNode.Get("input_tokens_details.cached_tokens")
}
if cached.Exists() {
detail.CachedTokens = cached.Int()
detail.CacheReadTokens = cached.Int()
}
cacheCreation := firstExistingUsageNode(
usageNode,
"input_tokens_details.cache_creation_tokens",
"input_tokens_details.cache_write_tokens",
"prompt_tokens_details.cache_creation_tokens",
"prompt_tokens_details.cache_write_tokens",
)
if cacheCreation.Exists() {
detail.CacheCreationTokens = cacheCreation.Int()
}
reasoning := usageNode.Get("completion_tokens_details.reasoning_tokens")
if !reasoning.Exists() {
reasoning = usageNode.Get("output_tokens_details.reasoning_tokens")
}
if reasoning.Exists() {
detail.ReasoningTokens = reasoning.Int()
}
if hasOpenAIStyleUsageBucketFields(usageNode) {
if inputNode.Exists() && outputNode.Exists() {
detail.TokenBreakdown = usage.NewSubsetTokenBreakdown(
detail.InputTokens,
detail.CacheReadTokens,
detail.CacheCreationTokens,
detail.OutputTokens,
detail.ReasoningTokens,
detail.TotalTokens,
)
} else {
cacheReadTokens := detail.CacheReadTokens
cacheCreationTokens := detail.CacheCreationTokens
if !inputNode.Exists() {
cacheReadTokens = 0
cacheCreationTokens = 0
}
reasoningTokens := detail.ReasoningTokens
if !outputNode.Exists() {
reasoningTokens = 0
}
detail.TokenBreakdown = usage.NewPartialSubsetTokenBreakdown(
detail.InputTokens,
cacheReadTokens,
cacheCreationTokens,
detail.OutputTokens,
reasoningTokens,
detail.TotalTokens,
)
}
} else {
detail.TokenBreakdown = usage.NewUnclassifiedTokenBreakdown(detail.TotalTokens)
}
if detail.TotalTokens == 0 {
detail.TotalTokens = detail.TokenBreakdown.TotalTokens
}
return detail
}
func ParseOpenAIStreamUsage(line []byte) (usage.Detail, bool) {
payload := jsonPayload(line)
if len(payload) == 0 || !gjson.ValidBytes(payload) {
return usage.Detail{}, false
}
responseServiceTier := extractResponseServiceTier(payload)
usageNode := gjson.GetBytes(payload, "usage")
if !hasOpenAIStyleUsageTokenFields(usageNode) {
if responseServiceTier == "" {
return usage.Detail{}, false
}
return usage.Detail{ResponseServiceTier: responseServiceTier}, true
}
detail := parseOpenAIStyleUsageNode(usageNode)
detail.ResponseServiceTier = responseServiceTier
return detail, true
}
func ParseClaudeUsage(data []byte) usage.Detail {
usageNode := gjson.ParseBytes(data).Get("usage")
if !usageNode.Exists() {
return usage.Detail{}
}
return parseClaudeUsageNode(usageNode)
}
func ParseClaudeStreamUsage(line []byte) (usage.Detail, bool) {
payload := jsonPayload(line)
if len(payload) == 0 || !gjson.ValidBytes(payload) {
return usage.Detail{}, false
}
usageNode := gjson.GetBytes(payload, "usage")
if !usageNode.Exists() {
return usage.Detail{}, false
}
return parseClaudeUsageNode(usageNode), true
}
func parseClaudeUsageNode(usageNode gjson.Result) usage.Detail {
cacheReadTokens := usageNode.Get("cache_read_input_tokens").Int()
cacheCreationTokens := usageNode.Get("cache_creation_input_tokens").Int()
rawOutputTokens := usageNode.Get("output_tokens").Int()
// Anthropic reports thinking as a subset of output_tokens. Prefer the official
// nested field, then fall back to legacy aliases used by some gateways.
reasoningNode := firstExistingUsageNode(
usageNode,
"output_tokens_details.thinking_tokens",
"output_tokens_details.reasoning_tokens",
"thinking_tokens",
)
reasoningTokens := reasoningNode.Int()
if reasoningTokens < 0 {
reasoningTokens = 0
}
nonReasoningOutput := rawOutputTokens
if reasoningTokens > 0 && reasoningTokens <= rawOutputTokens {
nonReasoningOutput = rawOutputTokens - reasoningTokens
} else if reasoningTokens > rawOutputTokens {
// Keep Detail.OutputTokens authoritative for keeper subset checks and
// avoid inventing extra non-reasoning output when the upstream payload
// is inconsistent.
nonReasoningOutput = 0
}
detail := usage.Detail{
InputTokens: usageNode.Get("input_tokens").Int(),
OutputTokens: rawOutputTokens,
ReasoningTokens: reasoningTokens,
CachedTokens: cacheReadTokens,
CacheReadTokens: cacheReadTokens,
CacheCreationTokens: cacheCreationTokens,
}
if detail.CachedTokens == 0 {
detail.CachedTokens = detail.CacheCreationTokens
}
// raw output_tokens already includes thinking; cache fields are independent
// from input_tokens in the Messages API.
detail.TotalTokens = detail.InputTokens + rawOutputTokens + detail.CacheReadTokens + detail.CacheCreationTokens
detail.TokenBreakdown = usage.NewIndependentTokenBreakdown(
detail.InputTokens,
detail.CacheReadTokens,
detail.CacheCreationTokens,
nonReasoningOutput,
detail.ReasoningTokens,
detail.TotalTokens,
)
return detail
}
func parseGeminiFamilyUsageDetail(node gjson.Result) usage.Detail {
cachedTokens := node.Get("cachedContentTokenCount").Int()
toolUseTokens := firstExistingUsageNode(node, "toolUsePromptTokenCount", "tool_use_prompt_token_count").Int()
inputTokens, okInput := safeUsageTokenSum(node.Get("promptTokenCount").Int(), toolUseTokens)
detail := usage.Detail{
InputTokens: inputTokens,
OutputTokens: node.Get("candidatesTokenCount").Int(),
ReasoningTokens: node.Get("thoughtsTokenCount").Int(),
TotalTokens: node.Get("totalTokenCount").Int(),
CachedTokens: cachedTokens,
CacheReadTokens: cachedTokens,
}
if !okInput {
detail.TokenBreakdown = invalidUsageTokenBreakdown(detail.TotalTokens)
return detail
}
if detail.TotalTokens == 0 {
var okTotal bool
detail.TotalTokens, okTotal = safeUsageTokenSum(detail.InputTokens, detail.OutputTokens, detail.ReasoningTokens)
if !okTotal {
detail.TotalTokens = 0
detail.TokenBreakdown = invalidUsageTokenBreakdown(0)
return detail
}
}
detail.TokenBreakdown = usage.NewSeparateReasoningTokenBreakdown(
detail.InputTokens,
detail.CacheReadTokens,
detail.CacheCreationTokens,
detail.OutputTokens,
detail.ReasoningTokens,
detail.TotalTokens,
)
return detail
}
func parseInteractionsUsageDetail(node gjson.Result) usage.Detail {
cacheRead := firstExistingUsageNode(node, "cache_read_tokens", "cacheReadTokens")
toolUseTokens := firstExistingUsageNode(node, "tool_use_tokens", "total_tool_use_tokens", "toolUseTokens", "totalToolUseTokens").Int()
inputTokens, okInput := safeUsageTokenSum(
firstExistingUsageNode(node, "input_tokens", "prompt_tokens", "total_input_tokens").Int(),
toolUseTokens,
)
detail := usage.Detail{
InputTokens: inputTokens,
OutputTokens: firstExistingUsageNode(node, "output_tokens", "completion_tokens", "total_output_tokens").Int(),
ReasoningTokens: firstExistingUsageNode(node, "reasoning_tokens", "thoughtsTokenCount", "total_thought_tokens").Int(),
TotalTokens: firstExistingUsageNode(node, "total_tokens", "totalTokenCount").Int(),
CachedTokens: firstExistingUsageNode(node, "cached_tokens", "cachedContentTokenCount", "total_cached_tokens").Int(),
CacheReadTokens: cacheRead.Int(),
CacheCreationTokens: firstExistingUsageNode(node, "cache_creation_tokens", "cacheCreationTokens", "cache_write_tokens", "cacheWriteTokens").Int(),
}
if !okInput {
detail.TokenBreakdown = invalidUsageTokenBreakdown(detail.TotalTokens)
return detail
}
if !cacheRead.Exists() && detail.CachedTokens > 0 {
detail.CacheReadTokens = detail.CachedTokens
}
if detail.TotalTokens == 0 {
var okTotal bool
detail.TotalTokens, okTotal = safeUsageTokenSum(detail.InputTokens, detail.OutputTokens, detail.ReasoningTokens)
if !okTotal {
detail.TotalTokens = 0
detail.TokenBreakdown = invalidUsageTokenBreakdown(0)
return detail
}
}
detail.TokenBreakdown = usage.NewSeparateReasoningTokenBreakdown(
detail.InputTokens,
detail.CacheReadTokens,
detail.CacheCreationTokens,
detail.OutputTokens,
detail.ReasoningTokens,
detail.TotalTokens,
)
return detail
}
func hasUsageDetail(detail usage.Detail) bool {
return hasNonZeroTokenUsage(detail)
}
func ParseInteractionsUsage(data []byte) usage.Detail {
root := gjson.ParseBytes(data)
node := firstExistingUsageNode(root, "usage", "total_usage", "metadata.total_usage", "metadata.usage", "usageMetadata", "usage_metadata", "interaction.usage", "interaction.total_usage", "interaction.metadata.total_usage")
if !node.Exists() {
return usage.Detail{}
}
if node.Get("promptTokenCount").Exists() || node.Get("candidatesTokenCount").Exists() {
detail := parseGeminiFamilyUsageDetail(node)
detail.ResponseServiceTier = extractResponseServiceTier(data)
return detail
}
detail := parseInteractionsUsageDetail(node)
detail.ResponseServiceTier = extractResponseServiceTier(data)
return detail
}
func extractResponseServiceTier(payload []byte) string {
if len(payload) == 0 || !gjson.ValidBytes(payload) {
return ""
}
return extractResponseServiceTierFromValidJSON(payload)
}
func extractResponseServiceTierFromValidJSON(payload []byte) string {
for _, path := range []string{"response.service_tier", "service_tier", "interaction.service_tier"} {
if tier := strings.TrimSpace(gjson.GetBytes(payload, path).String()); tier != "" {
return tier
}
}
return ""
}
func ParseInteractionsStreamUsage(line []byte) (usage.Detail, bool) {
payload := jsonPayload(line)
if len(payload) == 0 {
payload = line
}
if len(payload) == 0 || !gjson.ValidBytes(payload) {
return usage.Detail{}, false
}
detail := ParseInteractionsUsage(payload)
if !hasUsageDetail(detail) {
return usage.Detail{}, false
}
return detail, true
}
func ParseGeminiUsage(data []byte) usage.Detail {
usageNode := gjson.ParseBytes(data)
node := usageNode.Get("usageMetadata")
if !node.Exists() {
node = usageNode.Get("usage_metadata")
}
if !node.Exists() {
return usage.Detail{}
}
return parseGeminiFamilyUsageDetail(node)
}
func ParseGeminiStreamUsage(line []byte) (usage.Detail, bool) {
payload := jsonPayload(line)
if len(payload) == 0 || !gjson.ValidBytes(payload) {
return usage.Detail{}, false
}
node := gjson.GetBytes(payload, "usageMetadata")
if !node.Exists() {
node = gjson.GetBytes(payload, "usage_metadata")
}
if !node.Exists() {
return usage.Detail{}, false
}
detail := parseGeminiFamilyUsageDetail(node)
if !hasNonZeroTokenUsage(detail) {
return usage.Detail{}, false
}
return detail, true
}
func firstExistingUsageNode(root gjson.Result, paths ...string) gjson.Result {
for _, path := range paths {
node := root.Get(path)
if node.Exists() {
return node
}
}
return gjson.Result{}
}
func safeUsageTokenSum(values ...int64) (int64, bool) {
var total int64
for _, value := range values {
if value < 0 || total > int64(^uint64(0)>>1)-value {
return 0, false
}
total += value
}
return total, true
}
func invalidUsageTokenBreakdown(total int64) usage.TokenBreakdown {
if total < 0 {
total = 0
}
return usage.TokenBreakdown{
SchemaVersion: usage.TokenAccountingSchemaVersion,
Quality: usage.TokenAccountingQualityInconsistent,
TotalTokens: total,
UnclassifiedTokens: total,
}
}
func ParseAntigravityUsage(data []byte) usage.Detail {
usageNode := gjson.ParseBytes(data)
node := usageNode.Get("response.usageMetadata")
if !node.Exists() {
node = usageNode.Get("usageMetadata")
}
if !node.Exists() {
node = usageNode.Get("usage_metadata")
}
if !node.Exists() {
return usage.Detail{}
}
return parseGeminiFamilyUsageDetail(node)
}
func ParseAntigravityStreamUsage(line []byte) (usage.Detail, bool) {
payload := jsonPayload(line)
if len(payload) == 0 || !gjson.ValidBytes(payload) {
return usage.Detail{}, false
}
node := gjson.GetBytes(payload, "response.usageMetadata")
if !node.Exists() {
node = gjson.GetBytes(payload, "usageMetadata")
}
if !node.Exists() {
node = gjson.GetBytes(payload, "usage_metadata")
}
if !node.Exists() {
return usage.Detail{}, false
}
return parseGeminiFamilyUsageDetail(node), true
}
var stopChunkWithoutUsage sync.Map
func rememberStopWithoutUsage(traceID string) {
stopChunkWithoutUsage.Store(traceID, struct{}{})
time.AfterFunc(10*time.Minute, func() { stopChunkWithoutUsage.Delete(traceID) })
}
// FilterSSEUsageMetadata removes usageMetadata from SSE events that are not
// terminal (finishReason != "stop"). Stop chunks are left untouched. This
// function is shared between aistudio and antigravity executors.
func FilterSSEUsageMetadata(payload []byte) []byte {
if len(payload) == 0 {
return payload
}
lines := bytes.Split(payload, []byte("\n"))
modified := false
foundData := false
for idx, line := range lines {
trimmed := bytes.TrimSpace(line)
if len(trimmed) == 0 || !bytes.HasPrefix(trimmed, []byte("data:")) {
continue
}
foundData = true
dataIdx := bytes.Index(line, []byte("data:"))
if dataIdx < 0 {
continue
}
rawJSON := bytes.TrimSpace(line[dataIdx+5:])
traceID := gjson.GetBytes(rawJSON, "traceId").String()
if isStopChunkWithoutUsage(rawJSON) && traceID != "" {
rememberStopWithoutUsage(traceID)
continue
}
if traceID != "" {
if _, ok := stopChunkWithoutUsage.Load(traceID); ok && hasUsageMetadata(rawJSON) {
stopChunkWithoutUsage.Delete(traceID)
continue
}
}
cleaned, changed := StripUsageMetadataFromJSON(rawJSON)
if !changed {
continue
}
var rebuilt []byte
rebuilt = append(rebuilt, line[:dataIdx]...)
rebuilt = append(rebuilt, []byte("data:")...)
if len(cleaned) > 0 {
rebuilt = append(rebuilt, ' ')
rebuilt = append(rebuilt, cleaned...)
}
lines[idx] = rebuilt
modified = true
}
if !modified {
if !foundData {
// Handle payloads that are raw JSON without SSE data: prefix.
trimmed := bytes.TrimSpace(payload)
cleaned, changed := StripUsageMetadataFromJSON(trimmed)
if !changed {
return payload
}
return cleaned
}
return payload
}
return bytes.Join(lines, []byte("\n"))
}
// StripUsageMetadataFromJSON drops usageMetadata unless finishReason is present (terminal).
// It handles both formats:
// - Aistudio: candidates.0.finishReason
// - Antigravity: response.candidates.0.finishReason
func StripUsageMetadataFromJSON(rawJSON []byte) ([]byte, bool) {
jsonBytes := bytes.TrimSpace(rawJSON)
if len(jsonBytes) == 0 || !gjson.ValidBytes(jsonBytes) {
return rawJSON, false
}
// Check for finishReason in both aistudio and antigravity formats
finishReason := gjson.GetBytes(jsonBytes, "candidates.0.finishReason")
if !finishReason.Exists() {
finishReason = gjson.GetBytes(jsonBytes, "response.candidates.0.finishReason")
}
terminalReason := finishReason.Exists() && strings.TrimSpace(finishReason.String()) != ""
usageMetadata := gjson.GetBytes(jsonBytes, "usageMetadata")
if !usageMetadata.Exists() {
usageMetadata = gjson.GetBytes(jsonBytes, "response.usageMetadata")
}
// Terminal chunk: keep as-is.
if terminalReason {
return rawJSON, false
}
// Nothing to strip
if !usageMetadata.Exists() {
return rawJSON, false
}
// Remove usageMetadata from both possible locations
cleaned := jsonBytes
var changed bool
if usageMetadata = gjson.GetBytes(cleaned, "usageMetadata"); usageMetadata.Exists() {
// Rename usageMetadata to cpaUsageMetadata
cleaned, _ = sjson.SetRawBytes(cleaned, "cpaUsageMetadata", []byte(usageMetadata.Raw))
cleaned, _ = sjson.DeleteBytes(cleaned, "usageMetadata")
changed = true
}
if usageMetadata = gjson.GetBytes(cleaned, "response.usageMetadata"); usageMetadata.Exists() {
// Rename usageMetadata to cpaUsageMetadata
cleaned, _ = sjson.SetRawBytes(cleaned, "response.cpaUsageMetadata", []byte(usageMetadata.Raw))
cleaned, _ = sjson.DeleteBytes(cleaned, "response.usageMetadata")
changed = true
}
return cleaned, changed
}
func hasUsageMetadata(jsonBytes []byte) bool {
if len(jsonBytes) == 0 || !gjson.ValidBytes(jsonBytes) {
return false
}
if gjson.GetBytes(jsonBytes, "usageMetadata").Exists() {
return true
}
if gjson.GetBytes(jsonBytes, "response.usageMetadata").Exists() {
return true
}
return false
}
func isStopChunkWithoutUsage(jsonBytes []byte) bool {
if len(jsonBytes) == 0 || !gjson.ValidBytes(jsonBytes) {
return false
}
finishReason := gjson.GetBytes(jsonBytes, "candidates.0.finishReason")
if !finishReason.Exists() {
finishReason = gjson.GetBytes(jsonBytes, "response.candidates.0.finishReason")
}
trimmed := strings.TrimSpace(finishReason.String())
if !finishReason.Exists() || trimmed == "" {
return false
}
return !hasUsageMetadata(jsonBytes)
}
func JSONPayload(line []byte) []byte {
return jsonPayload(line)
}
func jsonPayload(line []byte) []byte {
trimmed := bytes.TrimSpace(line)
if len(trimmed) == 0 {
return nil
}
if bytes.Equal(trimmed, []byte("[DONE]")) {
return nil
}
if bytes.HasPrefix(trimmed, []byte("event:")) {
return nil
}
if bytes.HasPrefix(trimmed, []byte("data:")) {
trimmed = bytes.TrimSpace(trimmed[len("data:"):])
}
if len(trimmed) == 0 || trimmed[0] != '{' {
return nil
}
return trimmed
}