package usage import ( "bytes" "encoding/json" "strings" "time" "aigw/internal/domain" ) const maxCaptureBytes = 64 << 10 type Observer struct { protocol domain.Protocol stream bool buffer []byte line []byte usage domain.Usage found bool explicitTotal bool firstOutputAt time.Time now func() time.Time } func NewObserver(protocol domain.Protocol, stream bool) *Observer { return &Observer{protocol: protocol, stream: stream, now: time.Now} } func (o *Observer) Write(p []byte) (int, error) { if o.stream { o.observeSSE(p) } else { if len(p) > 0 { o.markFirstOutput() } o.captureTail(p) } return len(p), nil } // FirstOutputAt is the arrival time of the first user-visible output. For // streaming responses, metadata-only and heartbeat events are ignored. func (o *Observer) FirstOutputAt() time.Time { return o.firstOutputAt } func (o *Observer) Usage() domain.Usage { if o.stream { if len(o.line) > 0 { o.parseSSELine(o.line) } return o.usage } o.parseJSON(o.buffer) return o.usage } // Reported distinguishes a real zero-token usage object from a response that // omitted usage entirely. Billing must never infer this from token totals. func (o *Observer) Reported() bool { if !o.stream { o.parseJSON(o.buffer) } else if len(o.line) > 0 { o.parseSSELine(o.line) } return o.found } func (o *Observer) captureTail(p []byte) { if len(p) >= maxCaptureBytes { o.buffer = append(o.buffer[:0], p[len(p)-maxCaptureBytes:]...) return } if len(o.buffer)+len(p) > maxCaptureBytes { drop := len(o.buffer) + len(p) - maxCaptureBytes copy(o.buffer, o.buffer[drop:]) o.buffer = o.buffer[:len(o.buffer)-drop] } o.buffer = append(o.buffer, p...) } func (o *Observer) observeSSE(p []byte) { o.line = append(o.line, p...) for { index := bytes.IndexByte(o.line, '\n') if index < 0 { if len(o.line) > maxCaptureBytes { o.line = append(o.line[:0], o.line[len(o.line)-maxCaptureBytes:]...) } return } line := bytes.TrimSpace(o.line[:index]) o.parseSSELine(line) o.line = o.line[index+1:] } } func (o *Observer) parseSSELine(line []byte) { if !bytes.HasPrefix(line, []byte("data:")) { return } payload := bytes.TrimSpace(bytes.TrimPrefix(line, []byte("data:"))) if bytes.Equal(payload, []byte("[DONE]")) { return } o.observeStreamOutput(payload) if !bytes.Contains(payload, []byte("\"usage\"")) { return } o.parseJSON(payload) } type streamOutputEvent struct { Type string `json:"type"` Delta json.RawMessage `json:"delta"` Choices []struct { Delta struct { Content json.RawMessage `json:"content"` } `json:"delta"` } `json:"choices"` } func (o *Observer) observeStreamOutput(payload []byte) { if !o.firstOutputAt.IsZero() { return } var event streamOutputEvent if json.Unmarshal(payload, &event) != nil { return } for _, choice := range event.Choices { if rawContainsVisibleText(choice.Delta.Content) { o.markFirstOutput() return } } switch event.Type { case "response.output_text.delta": if rawContainsVisibleText(event.Delta) { o.markFirstOutput() } case "content_block_delta": var delta struct { Text string `json:"text"` } if json.Unmarshal(event.Delta, &delta) == nil && delta.Text != "" { o.markFirstOutput() } } } func rawContainsVisibleText(raw json.RawMessage) bool { if len(raw) == 0 || bytes.Equal(raw, []byte("null")) { return false } var text string if json.Unmarshal(raw, &text) == nil { return text != "" } var parts []struct { Text string `json:"text"` } if json.Unmarshal(raw, &parts) != nil { return false } for _, part := range parts { if part.Text != "" { return true } } return false } func (o *Observer) markFirstOutput() { if !o.firstOutputAt.IsZero() { return } o.firstOutputAt = o.now() } type tokenDetails struct { CachedTokens *int64 `json:"cached_tokens"` CacheWriteTokens *int64 `json:"cache_write_tokens"` } type usageFields struct { PromptTokens *int64 `json:"prompt_tokens"` CompletionTokens *int64 `json:"completion_tokens"` TotalTokens *int64 `json:"total_tokens"` InputTokens *int64 `json:"input_tokens"` OutputTokens *int64 `json:"output_tokens"` CacheCreationInputTokens *int64 `json:"cache_creation_input_tokens"` CacheReadInputTokens *int64 `json:"cache_read_input_tokens"` PromptTokensDetails *tokenDetails `json:"prompt_tokens_details"` InputTokensDetails *tokenDetails `json:"input_tokens_details"` } type responseEnvelope struct { Usage *usageFields `json:"usage"` Response *struct { Usage *usageFields `json:"usage"` } `json:"response"` Message *struct { Usage *usageFields `json:"usage"` } `json:"message"` } func (o *Observer) parseJSON(payload []byte) { var envelope responseEnvelope if err := json.Unmarshal(payload, &envelope); err != nil { payload = extractUsageObject(payload) if len(payload) == 0 { return } var fields usageFields if json.Unmarshal(payload, &fields) == nil { o.apply(&fields) } return } if envelope.Usage != nil { o.apply(envelope.Usage) } if envelope.Message != nil && envelope.Message.Usage != nil { o.apply(envelope.Message.Usage) } if envelope.Response != nil && envelope.Response.Usage != nil { o.apply(envelope.Response.Usage) } } func (o *Observer) apply(fields *usageFields) { inputReported := false if fields.PromptTokens != nil { o.usage.InputTokens = *fields.PromptTokens o.found = true inputReported = true } if fields.InputTokens != nil { o.usage.InputTokens = *fields.InputTokens o.found = true inputReported = true } if fields.CompletionTokens != nil { o.usage.OutputTokens = *fields.CompletionTokens o.found = true } if fields.OutputTokens != nil { o.usage.OutputTokens = *fields.OutputTokens o.found = true } if fields.TotalTokens != nil { o.usage.TotalTokens = *fields.TotalTokens o.found = true o.explicitTotal = true } if fields.CacheCreationInputTokens != nil { o.usage.CacheCreationInputTokens = *fields.CacheCreationInputTokens o.found = true } if fields.CacheReadInputTokens != nil { o.usage.CacheReadInputTokens = *fields.CacheReadInputTokens o.found = true } details := fields.InputTokensDetails if details == nil { details = fields.PromptTokensDetails } if details != nil { if details.CachedTokens != nil { o.usage.CacheReadInputTokens = *details.CachedTokens o.found = true } if details.CacheWriteTokens != nil { o.usage.CacheCreationInputTokens = *details.CacheWriteTokens o.found = true } // OpenAI reports cached token details as subsets of prompt/input_tokens. // Normalize them into mutually exclusive buckets before billing. Anthropic // reports its top-level cache fields separately, so they are not adjusted. if inputReported && (o.protocol == domain.ProtocolOpenAI || o.protocol == domain.ProtocolOpenAIResponses) { o.usage.InputTokens = max(0, o.usage.InputTokens-o.usage.CacheReadInputTokens-o.usage.CacheCreationInputTokens) } } if !o.explicitTotal && o.found { o.usage.TotalTokens = o.usage.InputTokens + o.usage.OutputTokens } } func extractUsageObject(payload []byte) []byte { index := strings.LastIndex(string(payload), `"usage"`) if index < 0 { return nil } rest := payload[index+len(`"usage"`):] start := bytes.IndexByte(rest, '{') if start < 0 { return nil } rest = rest[start:] depth := 0 inString := false escaped := false for i, b := range rest { if inString { if escaped { escaped = false continue } if b == '\\' { escaped = true } else if b == '"' { inString = false } continue } switch b { case '"': inString = true case '{': depth++ case '}': depth-- if depth == 0 { return rest[:i+1] } } } return nil }