116 lines
3.8 KiB
Go
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
|
|
}
|