summaryrefslogtreecommitdiff
path: root/internal/usage/observer_test.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/usage/observer_test.go26
1 files changed, 26 insertions, 0 deletions
diff --git a/internal/usage/observer_test.go b/internal/usage/observer_test.go
new file mode 100644
index 0000000..8fcf408
--- /dev/null
+++ b/internal/usage/observer_test.go
@@ -0,0 +1,26 @@
+package usage
+
+import (
+ "testing"
+
+ "aigw/internal/domain"
+)
+
+func TestObserverReadsOpenAIJSONUsage(t *testing.T) {
+ observer := NewObserver(domain.ProtocolOpenAI, false)
+ _, _ = observer.Write([]byte(`{"choices":[],"usage":{"prompt_tokens":11,"completion_tokens":7,"total_tokens":18}}`))
+ got := observer.Usage()
+ if got.InputTokens != 11 || 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_delta\ndata: {\"type\":\"message_delta\",\"usage\":{\"output_tokens\":8}}\n\n"))
+ got := observer.Usage()
+ if got.InputTokens != 12 || got.OutputTokens != 8 || got.TotalTokens != 20 {
+ t.Fatalf("unexpected usage: %+v", got)
+ }
+}