263 lines
9.3 KiB
Go
263 lines
9.3 KiB
Go
package superagent
|
|
|
|
import (
|
|
"errors"
|
|
"strings"
|
|
"testing"
|
|
)
|
|
|
|
func TestStreamStateConsumesCurrentTraceProtocol(t *testing.T) {
|
|
state := newStreamState("session-1", true)
|
|
var traces []TraceEvent
|
|
stream := strings.Join([]string{
|
|
"event: metadata\nid: 1\ndata: {\"run_id\":\"run-1\",\"resolved_profile_id\":\"profile-1\",\"resolved_profile_version_id\":\"version-1\"}\n",
|
|
"event: trace\nid: 2\ndata: {\"event\":\"message.delta\",\"run_id\":\"run-1\",\"text\":\"Hel\"}\n",
|
|
"event: trace\nid: 3\ndata: {\"event\":\"message.delta\",\"run_id\":\"run-1\",\"text\":\"lo\"}\n",
|
|
"event: trace\nid: 4\ndata: {\"event\":\"message.final\",\"run_id\":\"run-1\",\"message_id\":\"message-1\",\"text\":\"Hello\"}\n",
|
|
"event: trace\nid: 5\ndata: {\"event\":\"run.completed\",\"run_id\":\"run-1\",\"status\":\"success\"}\n",
|
|
"event: end\nid: 6\n",
|
|
}, "\n")
|
|
|
|
if err := state.consume(strings.NewReader(stream), func(event TraceEvent) {
|
|
traces = append(traces, event)
|
|
}); err != nil {
|
|
t.Fatalf("consume() error = %v", err)
|
|
}
|
|
result, err := state.result()
|
|
if err != nil {
|
|
t.Fatalf("result() error = %v", err)
|
|
}
|
|
if result.Answer != "Hello" || result.RunID != "run-1" || result.ProfileID != "profile-1" || result.ProfileVersionID != "version-1" {
|
|
t.Fatalf("unexpected result: %#v", result)
|
|
}
|
|
if result.LastEventID != "6" {
|
|
t.Fatalf("LastEventID = %q, want 6", result.LastEventID)
|
|
}
|
|
if len(traces) != 4 || traces[2].MessageID != "message-1" {
|
|
t.Fatalf("unexpected traces: %#v", traces)
|
|
}
|
|
}
|
|
|
|
func TestStreamStateConsumesLegacyValuesAndStructuredContent(t *testing.T) {
|
|
state := newStreamState("session-legacy", true)
|
|
stream := "event: values\n" +
|
|
"data: {\"messages\":[{\"type\":\"human\",\"content\":\"ignored\"},{\"type\":\"ai\",\"content\":[{\"text\":\"first\"},{\"content\":\"second\"}],\"response_metadata\":{\"finish_reason\":\"stop\",\"model_name\":\"model-1\"},\"usage_metadata\":{\"input_tokens\":2,\"output_tokens\":3,\"total_tokens\":5}}]}\n\n" +
|
|
"event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n" +
|
|
"event: end\n\n"
|
|
|
|
if err := state.consume(strings.NewReader(stream), nil); err != nil {
|
|
t.Fatalf("consume() error = %v", err)
|
|
}
|
|
result, err := state.result()
|
|
if err != nil {
|
|
t.Fatalf("result() error = %v", err)
|
|
}
|
|
if result.Answer != "first\nsecond" || result.ModelName != "model-1" {
|
|
t.Fatalf("unexpected legacy result: %#v", result)
|
|
}
|
|
if result.Usage != (TokenUsage{Input: 2, Output: 3, Total: 5}) {
|
|
t.Fatalf("Usage = %#v", result.Usage)
|
|
}
|
|
}
|
|
|
|
func TestStreamStateConsumesNoTraceFinalMessage(t *testing.T) {
|
|
state := newStreamState("session-no-trace", false)
|
|
stream := "event: messages\n" +
|
|
"data: [{\"type\":\"AIMessageChunk\",\"content\":\"\",\"response_metadata\":{\"finish_reason\":\"stop\",\"model_name\":\"model-1\"},\"usage_metadata\":{\"input_tokens\":2,\"output_tokens\":3,\"total_tokens\":5}}]\n\n" +
|
|
"event: message.final\n" +
|
|
"data: {\"run_id\":\"run-no-trace\",\"text\":\"No trace answer\"}\n\n" +
|
|
"event: end\n\n"
|
|
|
|
if err := state.consume(strings.NewReader(stream), nil); err != nil {
|
|
t.Fatalf("consume() error = %v", err)
|
|
}
|
|
result, err := state.result()
|
|
if err != nil {
|
|
t.Fatalf("result() error = %v", err)
|
|
}
|
|
if result.Answer != "No trace answer" || result.ModelName != "model-1" || result.RunID != "run-no-trace" {
|
|
t.Fatalf("unexpected result: %#v", result)
|
|
}
|
|
if result.Usage != (TokenUsage{Input: 2, Output: 3, Total: 5}) {
|
|
t.Fatalf("Usage = %#v", result.Usage)
|
|
}
|
|
}
|
|
|
|
func TestStreamStateNoTraceRequiresStopAndEnd(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
stream string
|
|
wantErr error
|
|
result bool
|
|
}{
|
|
{
|
|
name: "missing stop",
|
|
stream: "event: messages\n" +
|
|
"data: [{\"type\":\"AIMessageChunk\",\"content\":\"partial\",\"response_metadata\":{\"finish_reason\":null}}]\n\n" +
|
|
"event: message.final\ndata: {\"run_id\":\"run-1\",\"text\":\"complete\"}\n\n" +
|
|
"event: end\n\n",
|
|
wantErr: ErrProtocol,
|
|
result: true,
|
|
},
|
|
{
|
|
name: "missing message final",
|
|
stream: "event: messages\n" +
|
|
"data: [{\"type\":\"AIMessageChunk\",\"content\":\"\",\"response_metadata\":{\"finish_reason\":\"stop\"}}]\n\n" +
|
|
"event: end\n\n",
|
|
wantErr: ErrProtocol,
|
|
result: true,
|
|
},
|
|
{
|
|
name: "missing end",
|
|
stream: "event: messages\n" +
|
|
"data: [{\"type\":\"AIMessageChunk\",\"content\":\"\",\"response_metadata\":{\"finish_reason\":\"stop\"}}]\n\n" +
|
|
"event: message.final\ndata: {\"run_id\":\"run-1\",\"text\":\"complete\"}\n\n",
|
|
wantErr: ErrStreamIncomplete,
|
|
},
|
|
{
|
|
name: "top-level error",
|
|
stream: "event: error\ndata: {\"code\":\"provider_error\",\"message\":\"do not expose me\"}\n\n",
|
|
wantErr: ErrRunFailed,
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
state := newStreamState("session-no-trace", false)
|
|
err := state.consume(strings.NewReader(tt.stream), nil)
|
|
if tt.result && err == nil {
|
|
_, err = state.result()
|
|
}
|
|
if !errors.Is(err, tt.wantErr) {
|
|
t.Fatalf("error = %v, want errors.Is(..., %v)", err, tt.wantErr)
|
|
}
|
|
if err != nil && strings.Contains(strings.ToLower(err.Error()), "do not expose") {
|
|
t.Fatalf("error leaked provider message: %v", err)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestStreamStateTraceStillRequiresRunCompleted(t *testing.T) {
|
|
state := newStreamState("session-trace", true)
|
|
stream := "event: values\n" +
|
|
"data: {\"messages\":[{\"type\":\"ai\",\"content\":\"complete\",\"response_metadata\":{\"finish_reason\":\"stop\"}}]}\n\n" +
|
|
"event: end\n\n"
|
|
|
|
if err := state.consume(strings.NewReader(stream), nil); err != nil {
|
|
t.Fatalf("consume() error = %v", err)
|
|
}
|
|
_, err := state.result()
|
|
if !errors.Is(err, ErrProtocol) || !strings.Contains(err.Error(), "run.completed") {
|
|
t.Fatalf("result() error = %v, want missing run.completed protocol error", err)
|
|
}
|
|
}
|
|
|
|
func TestStreamStateSupportsMultilineDataAndDeduplicatesIDs(t *testing.T) {
|
|
state := newStreamState("session-1", true)
|
|
stream := ": heartbeat\n\n" +
|
|
"event: trace\nid: delta-1\ndata: {\"event\":\"message.delta\",\ndata: \"text\":\"A\"}\n\n" +
|
|
"event: trace\nid: delta-1\ndata: {\"event\":\"message.delta\",\"text\":\"duplicate\"}\n\n" +
|
|
"event: trace\nid: done-1\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n" +
|
|
"event: end\nid: end-1\n\n"
|
|
|
|
if err := state.consume(strings.NewReader(stream), nil); err != nil {
|
|
t.Fatalf("consume() error = %v", err)
|
|
}
|
|
result, err := state.result()
|
|
if err != nil {
|
|
t.Fatalf("result() error = %v", err)
|
|
}
|
|
if result.Answer != "A" {
|
|
t.Fatalf("Answer = %q, want A", result.Answer)
|
|
}
|
|
}
|
|
|
|
func TestStreamStateRejectsIncompleteOrFailedStreams(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
stream string
|
|
wantErr error
|
|
result bool
|
|
}{
|
|
{
|
|
name: "missing end",
|
|
stream: "event: trace\ndata: {\"event\":\"message.final\",\"text\":\"partial\"}\n\n" + "event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n",
|
|
wantErr: ErrStreamIncomplete,
|
|
},
|
|
{
|
|
name: "missing completed",
|
|
stream: "event: trace\ndata: {\"event\":\"message.final\",\"text\":\"partial\"}\n\n" + "event: end\n\n",
|
|
wantErr: ErrProtocol,
|
|
result: true,
|
|
},
|
|
{
|
|
name: "missing answer",
|
|
stream: "event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n" + "event: end\n\n",
|
|
wantErr: ErrProtocol,
|
|
result: true,
|
|
},
|
|
{
|
|
name: "top-level error",
|
|
stream: "event: error\ndata: {\"code\":\"provider_error\",\"message\":\"do not expose me\"}\n\n",
|
|
wantErr: ErrRunFailed,
|
|
},
|
|
{
|
|
name: "run failed",
|
|
stream: "event: trace\ndata: {\"event\":\"run.failed\",\"error\":{\"code\":\"tool_failed\",\"message\":\"private\"}}\n\n",
|
|
wantErr: ErrRunFailed,
|
|
},
|
|
{
|
|
name: "completed non-success",
|
|
stream: "event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"error\"}\n\n",
|
|
wantErr: ErrRunFailed,
|
|
},
|
|
{
|
|
name: "invalid JSON",
|
|
stream: "event: trace\ndata: not-json\n\n",
|
|
wantErr: ErrProtocol,
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
state := newStreamState("session-1", true)
|
|
err := state.consume(strings.NewReader(tt.stream), nil)
|
|
if tt.result && err == nil {
|
|
_, err = state.result()
|
|
}
|
|
if !errors.Is(err, tt.wantErr) {
|
|
t.Fatalf("error = %v, want errors.Is(..., %v)", err, tt.wantErr)
|
|
}
|
|
if strings.Contains(strings.ToLower(err.Error()), "do not expose") || strings.Contains(strings.ToLower(err.Error()), "private") {
|
|
t.Fatalf("error leaked provider message: %v", err)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestTraceProjectionRedactsAndBoundsText(t *testing.T) {
|
|
state := newStreamState("session-1", true)
|
|
secretText := "token=highly-secret " + strings.Repeat("x", maxTraceText+100)
|
|
stream := "event: trace\ndata: {\"event\":\"message.final\",\"text\":" + quotedJSON(secretText) + "}\n\n" +
|
|
"event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n" +
|
|
"event: end\n\n"
|
|
|
|
if err := state.consume(strings.NewReader(stream), nil); err != nil {
|
|
t.Fatalf("consume() error = %v", err)
|
|
}
|
|
if len(state.traceEvents) == 0 {
|
|
t.Fatal("expected trace event")
|
|
}
|
|
text := state.traceEvents[0].Text
|
|
if strings.Contains(text, "highly-secret") || len(text) > maxTraceText {
|
|
t.Fatalf("trace text was not safely projected: length=%d text=%q", len(text), text[:min(len(text), 80)])
|
|
}
|
|
}
|
|
|
|
func quotedJSON(value string) string {
|
|
value = strings.ReplaceAll(value, "\\", "\\\\")
|
|
value = strings.ReplaceAll(value, "\"", "\\\"")
|
|
return "\"" + value + "\""
|
|
}
|