summaryrefslogtreecommitdiff
path: root/internal/usage
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/usage/observer.go52
-rw-r--r--internal/usage/observer_test.go22
2 files changed, 62 insertions, 12 deletions
diff --git a/internal/usage/observer.go b/internal/usage/observer.go
index 808781f..71c1379 100644
--- a/internal/usage/observer.go
+++ b/internal/usage/observer.go
@@ -95,18 +95,28 @@ func (o *Observer) parseSSELine(line []byte) {
o.parseJSON(payload)
}
+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"`
+ 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"`
+ Usage *usageFields `json:"usage"`
+ Response *struct {
+ Usage *usageFields `json:"usage"`
+ } `json:"response"`
Message *struct {
Usage *usageFields `json:"usage"`
} `json:"message"`
@@ -131,16 +141,22 @@ func (o *Observer) parseJSON(payload []byte) {
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
@@ -163,6 +179,26 @@ func (o *Observer) apply(fields *usageFields) {
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
}
diff --git a/internal/usage/observer_test.go b/internal/usage/observer_test.go
index 3104acb..4a4eba6 100644
--- a/internal/usage/observer_test.go
+++ b/internal/usage/observer_test.go
@@ -8,25 +8,25 @@ import (
func TestObserverReadsOpenAIJSONUsage(t *testing.T) {
observer := NewObserver(domain.ProtocolOpenAI, false)
- _, _ = observer.Write([]byte(`{"choices":[],"usage":{"prompt_tokens":11,"completion_tokens":7,"total_tokens":18}}`))
+ _, _ = observer.Write([]byte(`{"choices":[],"usage":{"prompt_tokens":11,"prompt_tokens_details":{"cached_tokens":3,"cache_write_tokens":2},"completion_tokens":7,"total_tokens":18}}`))
got := observer.Usage()
if !observer.Reported() {
t.Fatal("expected usage to be marked as reported")
}
- if got.InputTokens != 11 || got.OutputTokens != 7 || got.TotalTokens != 18 {
+ if got.InputTokens != 6 || got.CacheReadInputTokens != 3 || got.CacheCreationInputTokens != 2 || got.OutputTokens != 7 || got.TotalTokens != 18 {
t.Fatalf("unexpected usage: %+v", got)
}
}
func TestObserverCombinesAnthropicSSEUsage(t *testing.T) {
observer := NewObserver(domain.ProtocolAnthropic, true)
- _, _ = observer.Write([]byte("event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":12,\"output_tokens\":1}}}\n\n"))
+ _, _ = observer.Write([]byte("event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":12,\"output_tokens\":1,\"cache_read_input_tokens\":5,\"cache_creation_input_tokens\":2}}}\n\n"))
_, _ = observer.Write([]byte("event: message_delta\ndata: {\"type\":\"message_delta\",\"usage\":{\"output_tokens\":8}}\n\n"))
got := observer.Usage()
if !observer.Reported() {
t.Fatal("expected streaming usage to be marked as reported")
}
- if got.InputTokens != 12 || got.OutputTokens != 8 || got.TotalTokens != 20 {
+ if got.InputTokens != 12 || got.OutputTokens != 8 || got.TotalTokens != 20 || got.CacheReadInputTokens != 5 || got.CacheCreationInputTokens != 2 {
t.Fatalf("unexpected usage: %+v", got)
}
}
@@ -43,3 +43,17 @@ func TestObserverDistinguishesMissingUsageFromReportedZero(t *testing.T) {
t.Fatal("explicit zero usage must be distinguished from a missing usage object")
}
}
+
+func TestObserverReadsResponsesUsage(t *testing.T) {
+ nonStream := NewObserver(domain.ProtocolOpenAIResponses, false)
+ _, _ = nonStream.Write([]byte(`{"object":"response","usage":{"input_tokens":11,"input_tokens_details":{"cached_tokens":3,"cache_write_tokens":2},"output_tokens":7,"total_tokens":18}}`))
+ if got := nonStream.Usage(); got.InputTokens != 6 || got.OutputTokens != 7 || got.TotalTokens != 18 || got.CacheReadInputTokens != 3 || got.CacheCreationInputTokens != 2 || !nonStream.Reported() {
+ t.Fatalf("unexpected non-stream Responses usage: %+v reported=%v", got, nonStream.Reported())
+ }
+
+ stream := NewObserver(domain.ProtocolOpenAIResponses, true)
+ _, _ = stream.Write([]byte("event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":13,\"output_tokens\":5,\"total_tokens\":18}}}\n\n"))
+ if got := stream.Usage(); got.InputTokens != 13 || got.OutputTokens != 5 || got.TotalTokens != 18 || !stream.Reported() {
+ t.Fatalf("unexpected streaming Responses usage: %+v reported=%v", got, stream.Reported())
+ }
+}