Files
fire-safety-ymd/internal/integration/superagent/chat_test.go
2026-09-05 15:46:37 +08:00

116 lines
3.8 KiB
Go

package superagent
import (
"context"
"errors"
"testing"
"fire-safety-ymd/internal/service"
)
func TestChatAgentAdapterMapsRequestsAndRemovesSensitiveTraceFields(t *testing.T) {
client := &fakeProviderClient{
session: Session{ID: "provider-session-1"},
result: Result{
RunID: "run-1",
ModelName: "qwen-plus-latest",
Answer: "safe final answer",
Usage: TokenUsage{Input: 1, Output: 2, Total: 3},
},
trace: TraceEvent{
Event: "tool.started",
RunID: "provider-run-id",
MessageID: "provider-message-id",
ToolCallID: "provider-tool-call-id",
ToolName: "fire_safety_test",
Text: "raw text must not cross the adapter",
Status: "running",
Timestamp: "2026-09-05T00:00:00Z",
},
}
adapter, err := NewChatAgentAdapter(client)
if err != nil {
t.Fatalf("NewChatAgentAdapter() error = %v", err)
}
session, err := adapter.CreateSession(context.Background(), service.AgentCreateSessionRequest{
ExternalSubjectID: "test-subject",
IdempotencyKey: "session-key",
RequestID: "request-1",
Metadata: map[string]any{"source": "test"},
})
if err != nil || session.ID != "provider-session-1" {
t.Fatalf("CreateSession() = %#v, %v", session, err)
}
var trace service.AgentTraceEvent
result, err := adapter.StreamMessage(context.Background(), service.AgentMessageRequest{
SessionID: session.ID,
Message: "question",
IdempotencyKey: "turn-key",
RequestID: "request-1",
}, func(event service.AgentTraceEvent) { trace = event })
if err != nil {
t.Fatalf("StreamMessage() error = %v", err)
}
if trace != (service.AgentTraceEvent{Event: "tool.started", ToolName: "fire_safety_test", Status: "running"}) {
t.Fatalf("projected trace = %#v", trace)
}
if result.Answer != "safe final answer" || result.RunID != "run-1" || result.ModelID != "qwen-plus-latest" || result.Usage.Total != 3 {
t.Fatalf("result = %#v", result)
}
if client.createRequest.ExternalSubjectID != "test-subject" || client.messageRequest.SessionID != "provider-session-1" || client.messageRequest.Message != "question" {
t.Fatalf("provider requests = %#v / %#v", client.createRequest, client.messageRequest)
}
}
func TestMapChatErrorUsesStableServiceCategories(t *testing.T) {
tests := []struct {
input error
want error
}{
{input: ErrDisabled, want: service.ErrChatUpstreamUnavailable},
{input: &HTTPStatusError{StatusCode: 503, ProviderCode: "unavailable"}, want: service.ErrChatUpstreamUnavailable},
{input: ErrProtocol, want: service.ErrChatUpstreamProtocol},
{input: ErrStreamIncomplete, want: service.ErrChatUpstreamProtocol},
{input: ErrRecoveryExhausted, want: service.ErrChatUpstreamProtocol},
{input: &RunError{Code: "failed"}, want: service.ErrChatRunFailed},
{input: context.DeadlineExceeded, want: context.DeadlineExceeded},
{input: context.Canceled, want: context.Canceled},
}
for _, tt := range tests {
mapped := mapChatError("test", tt.input)
if !errors.Is(mapped, tt.want) {
t.Fatalf("mapChatError(%v) = %v, want %v", tt.input, mapped, tt.want)
}
}
}
func TestNewChatAgentAdapterRejectsNilClient(t *testing.T) {
if _, err := NewChatAgentAdapter(nil); err == nil {
t.Fatal("NewChatAgentAdapter(nil) error = nil")
}
}
type fakeProviderClient struct {
session Session
result Result
trace TraceEvent
createErr error
streamErr error
createRequest CreateSessionRequest
messageRequest StreamMessageRequest
}
func (f *fakeProviderClient) CreateSession(_ context.Context, request CreateSessionRequest) (Session, error) {
f.createRequest = request
return f.session, f.createErr
}
func (f *fakeProviderClient) StreamMessage(_ context.Context, request StreamMessageRequest, trace TraceHandler) (Result, error) {
f.messageRequest = request
if trace != nil {
trace(f.trace)
}
return f.result, f.streamErr
}