132 lines
2.9 KiB
Go
132 lines
2.9 KiB
Go
package openai
|
|
|
|
import (
|
|
"os"
|
|
"strings"
|
|
"testing"
|
|
|
|
"nub/internal/llm"
|
|
)
|
|
|
|
func TestDecoder_FragmentedToolCall(t *testing.T) {
|
|
f, err := os.Open("testdata/tool_call_fragmented.sse")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer f.Close()
|
|
|
|
dec := newDecoder(f)
|
|
var events []llm.Event
|
|
for {
|
|
evs, more, err := dec.next()
|
|
if err != nil {
|
|
t.Fatalf("decode: %v", err)
|
|
}
|
|
events = append(events, evs...)
|
|
if !more {
|
|
break
|
|
}
|
|
}
|
|
|
|
var (
|
|
gotTextStart bool
|
|
gotTextDelta string
|
|
gotToolStart *llm.BlockStart
|
|
toolArgs string
|
|
gotDone *llm.Done
|
|
)
|
|
for _, ev := range events {
|
|
switch e := ev.(type) {
|
|
case llm.BlockStart:
|
|
if e.Block.Kind == llm.KindText {
|
|
gotTextStart = true
|
|
}
|
|
if e.Block.Kind == llm.KindToolUse {
|
|
cp := e
|
|
gotToolStart = &cp
|
|
}
|
|
case llm.BlockDelta:
|
|
if gotToolStart != nil && e.Index == gotToolStart.Index {
|
|
toolArgs += e.PartialJSON
|
|
} else {
|
|
gotTextDelta += e.Text
|
|
}
|
|
case llm.Done:
|
|
cp := e
|
|
gotDone = &cp
|
|
}
|
|
}
|
|
|
|
if !gotTextStart {
|
|
t.Error("expected a text BlockStart")
|
|
}
|
|
if gotTextDelta != "Ich lese die Datei." {
|
|
t.Errorf("text delta = %q", gotTextDelta)
|
|
}
|
|
if gotToolStart == nil {
|
|
t.Fatal("expected a tool_use BlockStart")
|
|
}
|
|
if gotToolStart.Block.ID != "call_abc" || gotToolStart.Block.Name != "read" {
|
|
t.Errorf("tool start = %+v", gotToolStart.Block)
|
|
}
|
|
if toolArgs != `{"path":"a.go"}` {
|
|
t.Errorf("accumulated tool args = %q", toolArgs)
|
|
}
|
|
if gotDone == nil {
|
|
t.Fatal("expected a Done event")
|
|
}
|
|
if gotDone.Stop != llm.StopToolUse {
|
|
t.Errorf("stop reason = %q, want tool_use", gotDone.Stop)
|
|
}
|
|
if gotDone.Usage.InputTokens != 42 || gotDone.Usage.OutputTokens != 7 {
|
|
t.Errorf("usage = %+v", gotDone.Usage)
|
|
}
|
|
}
|
|
|
|
// TestDecoder_ParsesCachedTokens verifiziert usage.prompt_tokens_details.
|
|
// cached_tokens (M7: Grundlage für die Prompt-Cache-Verifikation) landet
|
|
// korrekt in llm.Usage.CacheReadTokens.
|
|
func TestDecoder_ParsesCachedTokens(t *testing.T) {
|
|
f, err := os.Open("testdata/cached_usage.sse")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer f.Close()
|
|
|
|
dec := newDecoder(f)
|
|
var gotDone *llm.Done
|
|
for {
|
|
evs, more, err := dec.next()
|
|
if err != nil {
|
|
t.Fatalf("decode: %v", err)
|
|
}
|
|
for _, ev := range evs {
|
|
if d, ok := ev.(llm.Done); ok {
|
|
cp := d
|
|
gotDone = &cp
|
|
}
|
|
}
|
|
if !more {
|
|
break
|
|
}
|
|
}
|
|
|
|
if gotDone == nil {
|
|
t.Fatal("expected a Done event")
|
|
}
|
|
if gotDone.Usage.InputTokens != 1200 || gotDone.Usage.OutputTokens != 5 {
|
|
t.Errorf("usage = %+v", gotDone.Usage)
|
|
}
|
|
if gotDone.Usage.CacheReadTokens != 896 {
|
|
t.Errorf("cache read tokens = %d, want 896", gotDone.Usage.CacheReadTokens)
|
|
}
|
|
}
|
|
|
|
func TestDecoder_StreamErrorEvent(t *testing.T) {
|
|
body := `data: {"error":{"message":"rate limited","type":"rate_limit_error"}}` + "\n\ndata: [DONE]\n"
|
|
dec := newDecoder(strings.NewReader(body))
|
|
_, _, err := dec.next()
|
|
if err == nil {
|
|
t.Fatal("expected error from stream error event")
|
|
}
|
|
}
|