summaryrefslogtreecommitdiff
path: root/internal/usage
diff options
context:
space:
mode:
Diffstat (limited to 'internal/usage')
-rw-r--r--internal/usage/observer.go88
-rw-r--r--internal/usage/observer_test.go71
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)
+ }
+}