summaryrefslogtreecommitdiff
path: root/internal/usage/observer.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/usage/observer.go88
1 files changed, 86 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"`