diff options
Diffstat (limited to 'internal/usage')
| -rw-r--r-- | internal/usage/observer.go | 88 | ||||
| -rw-r--r-- | internal/usage/observer_test.go | 71 |
2 files changed, 157 insertions, 2 deletions
diff --git a/internal/usage/observer.go b/internal/usage/observer.go index 71c1379..4c7b2b4 100644 --- a/internal/usage/observer.go +++ b/internal/usage/observer.go @@ -4,6 +4,7 @@ import ( "bytes" "encoding/json" "strings" + "time" "aigw/internal/domain" ) @@ -18,21 +19,32 @@ type Observer struct { 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} + 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 { @@ -85,16 +97,88 @@ func (o *Observer) observeSSE(p []byte) { } func (o *Observer) parseSSELine(line []byte) { - if !bytes.HasPrefix(line, []byte("data:")) || !bytes.Contains(line, []byte("\"usage\"")) { + 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"` diff --git a/internal/usage/observer_test.go b/internal/usage/observer_test.go index 4a4eba6..e96a38a 100644 --- a/internal/usage/observer_test.go +++ b/internal/usage/observer_test.go @@ -2,6 +2,7 @@ package usage import ( "testing" + "time" "aigw/internal/domain" ) @@ -57,3 +58,73 @@ func TestObserverReadsResponsesUsage(t *testing.T) { t.Fatalf("unexpected streaming Responses usage: %+v reported=%v", got, stream.Reported()) } } + +func TestObserverReadsEmbeddingsUsage(t *testing.T) { + observer := NewObserver(domain.ProtocolOpenAIEmbeddings, false) + _, _ = observer.Write([]byte(`{"object":"list","data":[],"usage":{"prompt_tokens":17,"total_tokens":17}}`)) + got := observer.Usage() + if !observer.Reported() || got.InputTokens != 17 || got.OutputTokens != 0 || got.TotalTokens != 17 { + t.Fatalf("unexpected Embeddings usage: %+v reported=%v", got, observer.Reported()) + } +} + +func TestObserverMarksFirstVisibleStreamingOutput(t *testing.T) { + base := time.Date(2026, 8, 6, 1, 2, 3, 0, time.UTC) + tests := []struct { + name string + protocol domain.Protocol + metadata string + output string + }{ + { + name: "openai chat", protocol: domain.ProtocolOpenAI, + metadata: "data: {\"choices\":[{\"delta\":{\"role\":\"assistant\",\"content\":null}}]}\n\n", + output: "data: {\"choices\":[{\"delta\":{\"content\":\"Hello\"}}]}\n\n", + }, + { + name: "responses", protocol: domain.ProtocolOpenAIResponses, + metadata: "data: {\"type\":\"response.created\"}\n\n", + output: "data: {\"type\":\"response.output_text.delta\",\"delta\":\"Hello\"}\n\n", + }, + { + name: "anthropic", protocol: domain.ProtocolAnthropic, + metadata: "data: {\"type\":\"content_block_start\",\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n", + output: "data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"Hello\"}}\n\n", + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + observer := NewObserver(test.protocol, true) + current := base + observer.now = func() time.Time { return current } + _, _ = observer.Write([]byte(test.metadata)) + if got := observer.FirstOutputAt(); !got.IsZero() { + t.Fatalf("metadata marked first output at %v", got) + } + current = base.Add(275 * time.Millisecond) + _, _ = observer.Write([]byte(test.output)) + if got := observer.FirstOutputAt(); !got.Equal(current) { + t.Fatalf("first output = %v, want %v", got, current) + } + current = base.Add(time.Second) + _, _ = observer.Write([]byte(test.output)) + if got := observer.FirstOutputAt(); !got.Equal(base.Add(275 * time.Millisecond)) { + t.Fatalf("first output changed to %v", got) + } + }) + } +} + +func TestObserverMarksFirstNonStreamingBodyWrite(t *testing.T) { + base := time.Date(2026, 8, 6, 1, 2, 3, 0, time.UTC) + observer := NewObserver(domain.ProtocolOpenAI, false) + observer.now = func() time.Time { return base } + _, _ = observer.Write(nil) + if !observer.FirstOutputAt().IsZero() { + t.Fatal("empty write must not mark first output") + } + _, _ = observer.Write([]byte(`{"choices":[]}`)) + if got := observer.FirstOutputAt(); !got.Equal(base) { + t.Fatalf("first output = %v, want %v", got, base) + } +} |
