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

452 lines
18 KiB
Go

package superagent
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
)
const testAPIKey = "test-open-api-key"
func TestHTTPClientCreatesSessionAndStreamsMessage(t *testing.T) {
var csrfTokens []string
var tokensMu sync.Mutex
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
csrf := requireProviderHeaders(t, request, "request-1")
tokensMu.Lock()
csrfTokens = append(csrfTokens, csrf)
tokensMu.Unlock()
switch {
case request.Method == http.MethodPost && request.URL.Path == "/api/open/agent-sessions":
if request.Header.Get("Accept") != "application/json" || request.Header.Get("Content-Type") != "application/json" {
t.Errorf("unexpected create content headers: %#v", request.Header)
}
var payload map[string]any
if err := json.NewDecoder(request.Body).Decode(&payload); err != nil {
t.Errorf("decode create body: %v", err)
writer.WriteHeader(http.StatusBadRequest)
return
}
if payload["external_subject_id"] != "probe-subject" || payload["idempotency_key"] != "session-key" {
t.Errorf("unexpected create body: %#v", payload)
}
writer.Header().Set("Content-Type", "application/json")
fmt.Fprint(writer, `{"session_id":"session-1","status":"active"}`)
case request.Method == http.MethodPost && request.URL.Path == "/api/open/agent-sessions/session-1/messages/stream":
if request.URL.Query().Get("include_trace") != "true" || request.Header.Get("Accept") != "text/event-stream" {
t.Errorf("unexpected stream request: %s %#v", request.URL.String(), request.Header)
}
var payload map[string]any
if err := json.NewDecoder(request.Body).Decode(&payload); err != nil {
t.Errorf("decode stream body: %v", err)
writer.WriteHeader(http.StatusBadRequest)
return
}
if payload["message"] != "safe probe" || payload["idempotency_key"] != "message-key" {
t.Errorf("unexpected stream body: %#v", payload)
}
writer.Header().Set("Content-Type", "text/event-stream; charset=utf-8")
writer.Header().Set("Content-Location", "/api/open/agent-sessions/session-1/runs/run-1")
fmt.Fprint(writer, successfulSSE("run-1", "connectivity OK"))
default:
t.Errorf("unexpected request: %s %s", request.Method, request.URL.String())
writer.WriteHeader(http.StatusNotFound)
}
}))
defer server.Close()
client := newTestHTTPClient(t, server.URL, 2, time.Millisecond)
session, err := client.CreateSession(context.Background(), CreateSessionRequest{
ExternalSubjectID: "probe-subject",
IdempotencyKey: "session-key",
RequestID: "request-1",
Metadata: map[string]any{"purpose": "test"},
})
if err != nil {
t.Fatalf("CreateSession() error = %v", err)
}
if session != (Session{ID: "session-1", Status: "active"}) {
t.Fatalf("unexpected session: %#v", session)
}
var traces []TraceEvent
result, err := client.StreamMessage(context.Background(), StreamMessageRequest{
SessionID: session.ID,
Message: "safe probe",
IdempotencyKey: "message-key",
RequestID: "request-1",
Metadata: map[string]any{"purpose": "test"},
}, func(event TraceEvent) {
traces = append(traces, event)
})
if err != nil {
t.Fatalf("StreamMessage() error = %v", err)
}
if result.Answer != "connectivity OK" || result.SessionID != "session-1" || result.RunID != "run-1" {
t.Fatalf("unexpected result: %#v", result)
}
if len(traces) != 2 {
t.Fatalf("trace count = %d, want 2", len(traces))
}
tokensMu.Lock()
defer tokensMu.Unlock()
if len(csrfTokens) != 2 || csrfTokens[0] == csrfTokens[1] {
t.Fatalf("CSRF tokens must be non-empty and unique per request: %#v", csrfTokens)
}
}
func TestHTTPClientRecoversWithoutRepostingMessage(t *testing.T) {
var messagePosts atomic.Int32
var runQueries atomic.Int32
var eventQueries atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
requireProviderHeaders(t, request, "request-recovery")
switch {
case request.Method == http.MethodPost && request.URL.Path == "/api/open/agent-sessions/session-1/messages/stream":
messagePosts.Add(1)
writer.Header().Set("Content-Type", "text/event-stream")
writer.Header().Set("Content-Location", "/api/open/agent-sessions/session-1/runs/run-1")
fmt.Fprint(writer, "event: trace\nid: event-1\ndata: {\"event\":\"message.delta\",\"run_id\":\"run-1\",\"text\":\"partial\"}\n\n")
case request.Method == http.MethodGet && request.URL.Path == "/api/open/agent-sessions/session-1/runs/run-1":
runQueries.Add(1)
writer.Header().Set("Content-Type", "application/json")
fmt.Fprint(writer, `{"status":"running"}`)
case request.Method == http.MethodGet && request.URL.Path == "/api/open/agent-sessions/session-1/runs/run-1/events":
eventQueries.Add(1)
if request.Header.Get("Last-Event-ID") != "event-1" {
t.Errorf("Last-Event-ID = %q, want event-1", request.Header.Get("Last-Event-ID"))
}
writer.Header().Set("Content-Type", "text/event-stream")
fmt.Fprint(writer, "event: trace\nid: event-1\ndata: {\"event\":\"message.delta\",\"text\":\"duplicate\"}\n\n")
fmt.Fprint(writer, "event: trace\nid: event-2\ndata: {\"event\":\"message.final\",\"text\":\"recovered answer\"}\n\n")
fmt.Fprint(writer, "event: trace\nid: event-3\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n")
fmt.Fprint(writer, "event: end\nid: event-4\n\n")
default:
t.Errorf("unexpected request: %s %s", request.Method, request.URL.Path)
writer.WriteHeader(http.StatusNotFound)
}
}))
defer server.Close()
client := newTestHTTPClient(t, server.URL, 2, time.Millisecond)
result, err := client.StreamMessage(context.Background(), StreamMessageRequest{
SessionID: "session-1",
Message: "safe probe",
IdempotencyKey: "message-key",
RequestID: "request-recovery",
}, nil)
if err != nil {
t.Fatalf("StreamMessage() error = %v", err)
}
if result.Answer != "recovered answer" || result.LastEventID != "event-4" {
t.Fatalf("unexpected recovered result: %#v", result)
}
if messagePosts.Load() != 1 || runQueries.Load() != 1 || eventQueries.Load() != 1 {
t.Fatalf("request counts: message=%d run=%d events=%d", messagePosts.Load(), runQueries.Load(), eventQueries.Load())
}
}
func TestHTTPClientDerivesRecoveryURLFromMetadataRunID(t *testing.T) {
var messagePosts atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
switch {
case request.Method == http.MethodPost && request.URL.Path == "/api/open/agent-sessions/session-1/messages/stream":
messagePosts.Add(1)
writer.Header().Set("Content-Type", "text/event-stream")
fmt.Fprint(writer, "event: metadata\nid: meta-1\ndata: {\"run_id\":\"run-derived\"}\n\n")
case request.Method == http.MethodGet && request.URL.Path == "/api/open/agent-sessions/session-1/runs/run-derived":
writer.Header().Set("Content-Type", "application/json")
fmt.Fprint(writer, `{"status":"running"}`)
case request.Method == http.MethodGet && request.URL.Path == "/api/open/agent-sessions/session-1/runs/run-derived/events":
if request.Header.Get("Last-Event-ID") != "meta-1" {
t.Errorf("Last-Event-ID = %q, want meta-1", request.Header.Get("Last-Event-ID"))
}
writer.Header().Set("Content-Type", "text/event-stream")
fmt.Fprint(writer, successfulSSE("run-derived", "derived recovery"))
default:
t.Errorf("unexpected request: %s %s", request.Method, request.URL.Path)
writer.WriteHeader(http.StatusNotFound)
}
}))
defer server.Close()
client := newTestHTTPClient(t, server.URL, 1, time.Millisecond)
result, err := client.StreamMessage(context.Background(), validMessageRequest(), nil)
if err != nil {
t.Fatalf("StreamMessage() error = %v", err)
}
if result.Answer != "derived recovery" || result.RunID != "run-derived" || messagePosts.Load() != 1 {
t.Fatalf("unexpected recovery result: %#v, posts=%d", result, messagePosts.Load())
}
}
func TestHTTPClientStopsRecoveryOnFailedRunStatus(t *testing.T) {
var eventQueries atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
switch {
case request.Method == http.MethodPost:
writer.Header().Set("Content-Type", "text/event-stream")
writer.Header().Set("Content-Location", "/api/open/agent-sessions/session-1/runs/run-1")
fmt.Fprint(writer, "event: trace\ndata: {\"event\":\"message.delta\",\"text\":\"partial\"}\n\n")
case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/run-1"):
writer.Header().Set("Content-Type", "application/json")
fmt.Fprint(writer, `{"status":"timeout"}`)
case request.Method == http.MethodGet:
eventQueries.Add(1)
writer.WriteHeader(http.StatusInternalServerError)
}
}))
defer server.Close()
client := newTestHTTPClient(t, server.URL, 2, time.Millisecond)
_, err := client.StreamMessage(context.Background(), validMessageRequest(), nil)
if !errors.Is(err, ErrRunFailed) {
t.Fatalf("StreamMessage() error = %v, want ErrRunFailed", err)
}
if eventQueries.Load() != 0 {
t.Fatalf("events endpoint queried %d times after failed run", eventQueries.Load())
}
}
func TestHTTPClientReturnsRecoveryExhaustedAfterBoundedAttempts(t *testing.T) {
var messagePosts atomic.Int32
var runQueries atomic.Int32
var eventQueries atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
switch {
case request.Method == http.MethodPost:
messagePosts.Add(1)
writer.Header().Set("Content-Type", "text/event-stream")
writer.Header().Set("Content-Location", "/api/open/agent-sessions/session-1/runs/run-1")
fmt.Fprint(writer, "event: trace\nid: initial-1\ndata: {\"event\":\"message.delta\",\"text\":\"partial\"}\n\n")
case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/run-1"):
runQueries.Add(1)
writer.Header().Set("Content-Type", "application/json")
fmt.Fprint(writer, `{"status":"running"}`)
case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/events"):
eventQueries.Add(1)
writer.Header().Set("Content-Type", "text/event-stream")
fmt.Fprint(writer, ": still running\n\n")
}
}))
defer server.Close()
client := newTestHTTPClient(t, server.URL, 2, time.Millisecond)
result, err := client.StreamMessage(context.Background(), validMessageRequest(), nil)
if !errors.Is(err, ErrRecoveryExhausted) {
t.Fatalf("StreamMessage() error = %v, want ErrRecoveryExhausted", err)
}
if result.Answer != "" || messagePosts.Load() != 1 || runQueries.Load() != 2 || eventQueries.Load() != 2 {
t.Fatalf("unbounded or partial recovery: result=%#v messages=%d runs=%d events=%d", result, messagePosts.Load(), runQueries.Load(), eventQueries.Load())
}
}
func TestHTTPClientRejectsCrossOriginContentLocation(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Content-Type", "text/event-stream")
writer.Header().Set("Content-Location", "https://attacker.example/api/open/agent-sessions/session-1/runs/run-1")
fmt.Fprint(writer, successfulSSE("run-1", "must not be accepted"))
}))
defer server.Close()
client := newTestHTTPClient(t, server.URL, 1, time.Millisecond)
_, err := client.StreamMessage(context.Background(), validMessageRequest(), nil)
if !errors.Is(err, ErrProtocol) || !strings.Contains(err.Error(), "cross-origin") {
t.Fatalf("StreamMessage() error = %v, want cross-origin protocol error", err)
}
}
func TestHTTPClientNeverReturnsPartialAnswerWithoutRecoveryLocation(t *testing.T) {
var messagePosts atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
messagePosts.Add(1)
writer.Header().Set("Content-Type", "text/event-stream")
fmt.Fprint(writer, "event: trace\ndata: {\"event\":\"message.delta\",\"text\":\"partial\"}\n\n")
}))
defer server.Close()
client := newTestHTTPClient(t, server.URL, 2, time.Millisecond)
result, err := client.StreamMessage(context.Background(), validMessageRequest(), nil)
if !errors.Is(err, ErrStreamIncomplete) {
t.Fatalf("StreamMessage() error = %v, want ErrStreamIncomplete", err)
}
if result.Answer != "" || messagePosts.Load() != 1 {
t.Fatalf("partial response escaped or message was retried: result=%#v posts=%d", result, messagePosts.Load())
}
}
func TestHTTPClientRedactsHTTPErrorBodyAndKey(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Content-Type", "application/json")
writer.WriteHeader(http.StatusTooManyRequests)
fmt.Fprint(writer, `{"code":"rate_limited","message":"provider-private-value"}`)
}))
defer server.Close()
client := newTestHTTPClient(t, server.URL, 0, time.Millisecond)
_, err := client.CreateSession(context.Background(), CreateSessionRequest{
ExternalSubjectID: "subject-1",
IdempotencyKey: "session-key",
})
if !errors.Is(err, ErrHTTPStatus) {
t.Fatalf("CreateSession() error = %v, want ErrHTTPStatus", err)
}
if !strings.Contains(err.Error(), "rate_limited") || strings.Contains(err.Error(), "provider-private-value") || strings.Contains(err.Error(), testAPIKey) {
t.Fatalf("unsafe or incomplete HTTP error: %v", err)
}
}
func TestHTTPClientLimitsControlResponseBody(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Content-Type", "application/json")
fmt.Fprint(writer, `{"session_id":"`)
fmt.Fprint(writer, strings.Repeat("x", int(maxControlResponseBytes)))
fmt.Fprint(writer, `"}`)
}))
defer server.Close()
client := newTestHTTPClient(t, server.URL, 0, time.Millisecond)
_, err := client.CreateSession(context.Background(), CreateSessionRequest{
ExternalSubjectID: "subject-1",
IdempotencyKey: "session-key",
})
if !errors.Is(err, ErrProtocol) || !strings.Contains(err.Error(), "size limit") {
t.Fatalf("CreateSession() error = %v, want bounded protocol error", err)
}
}
func TestHTTPClientDisabledDoesNotUseNetwork(t *testing.T) {
client, err := NewHTTPClient(Config{Enabled: false})
if err != nil {
t.Fatalf("NewHTTPClient() error = %v", err)
}
if _, err := client.CreateSession(context.Background(), CreateSessionRequest{}); !errors.Is(err, ErrDisabled) {
t.Fatalf("CreateSession() error = %v, want ErrDisabled", err)
}
if _, err := client.StreamMessage(context.Background(), StreamMessageRequest{}, nil); !errors.Is(err, ErrDisabled) {
t.Fatalf("StreamMessage() error = %v, want ErrDisabled", err)
}
}
func TestNewHTTPClientRejectsUnsafeAPIKey(t *testing.T) {
_, err := NewHTTPClient(Config{
Enabled: true,
BaseURL: "https://superagent.example.test",
APIKey: "unsafe key",
})
if !errors.Is(err, ErrInvalidConfig) || strings.Contains(err.Error(), "unsafe key") {
t.Fatalf("NewHTTPClient() error = %v, want redacted config error", err)
}
}
func TestHTTPClientHonorsContextDuringRecoveryBackoff(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if request.Method != http.MethodPost {
t.Errorf("unexpected recovery request after timeout: %s", request.Method)
}
writer.Header().Set("Content-Type", "text/event-stream")
writer.Header().Set("Content-Location", "/api/open/agent-sessions/session-1/runs/run-1")
fmt.Fprint(writer, "event: trace\nid: event-1\ndata: {\"event\":\"message.delta\",\"text\":\"partial\"}\n\n")
}))
defer server.Close()
client := newTestHTTPClient(t, server.URL, 2, 100*time.Millisecond)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
defer cancel()
_, err := client.StreamMessage(ctx, validMessageRequest(), nil)
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("StreamMessage() error = %v, want context deadline", err)
}
}
func TestHTTPClientValidatesMessageLimitAndEventContentType(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Content-Type", "application/json")
fmt.Fprint(writer, `{}`)
}))
defer server.Close()
client, err := NewHTTPClient(Config{
Enabled: true,
BaseURL: server.URL,
APIKey: testAPIKey,
ConnectTimeout: time.Second,
RecoveryInitialBackoff: time.Millisecond,
MaxMessageBytes: 4,
})
if err != nil {
t.Fatalf("NewHTTPClient() error = %v", err)
}
tooLarge := validMessageRequest()
tooLarge.Message = "12345"
if _, err := client.StreamMessage(context.Background(), tooLarge, nil); !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("large message error = %v, want ErrInvalidRequest", err)
}
valid := validMessageRequest()
valid.Message = "1234"
if _, err := client.StreamMessage(context.Background(), valid, nil); !errors.Is(err, ErrProtocol) {
t.Fatalf("wrong content-type error = %v, want ErrProtocol", err)
}
}
func newTestHTTPClient(t *testing.T, baseURL string, recoveryAttempts int, backoff time.Duration) *HTTPClient {
t.Helper()
client, err := NewHTTPClient(Config{
Enabled: true,
BaseURL: baseURL,
APIKey: testAPIKey,
ConnectTimeout: time.Second,
RecoveryMaxAttempts: recoveryAttempts,
RecoveryInitialBackoff: backoff,
MaxMessageBytes: 1024,
})
if err != nil {
t.Fatalf("NewHTTPClient() error = %v", err)
}
return client
}
func requireProviderHeaders(t *testing.T, request *http.Request, requestID string) string {
t.Helper()
if request.Header.Get("Authorization") != "Bearer "+testAPIKey {
t.Errorf("unexpected Authorization header")
}
if request.Header.Get("Cache-Control") != "no-cache" {
t.Errorf("Cache-Control = %q", request.Header.Get("Cache-Control"))
}
if request.Header.Get("X-Request-ID") != requestID {
t.Errorf("X-Request-ID = %q, want %q", request.Header.Get("X-Request-ID"), requestID)
}
csrf := request.Header.Get("X-CSRF-Token")
cookie, err := request.Cookie("csrf_token")
if err != nil || csrf == "" || cookie.Value != csrf {
t.Errorf("invalid CSRF double-submit values: header-present=%t cookie-error=%v", csrf != "", err)
}
return csrf
}
func validMessageRequest() StreamMessageRequest {
return StreamMessageRequest{
SessionID: "session-1",
Message: "safe",
IdempotencyKey: "message-key",
RequestID: "request-1",
}
}
func successfulSSE(runID, answer string) string {
return "event: trace\nid: event-1\ndata: {\"event\":\"message.final\",\"run_id\":\"" + runID + "\",\"text\":\"" + answer + "\"}\n\n" +
"event: trace\nid: event-2\ndata: {\"event\":\"run.completed\",\"run_id\":\"" + runID + "\",\"status\":\"success\"}\n\n" +
"event: end\nid: event-3\n\n"
}