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 TestHTTPClientStreamsWithoutTraceWhenDisabled(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { if got := request.URL.Query().Get("include_trace"); got != "false" { t.Errorf("include_trace = %q, want false", got) } writer.Header().Set("Content-Type", "text/event-stream") fmt.Fprint(writer, "event: messages\ndata: [{\"type\":\"AIMessageChunk\",\"content\":\"\",\"response_metadata\":{\"finish_reason\":\"stop\"}}]\n\n") fmt.Fprint(writer, "event: message.final\ndata: {\"run_id\":\"run-no-trace\",\"text\":\"no-trace answer\"}\n\n") fmt.Fprint(writer, "event: end\n\n") })) defer server.Close() client, err := NewHTTPClient(Config{ Enabled: true, BaseURL: server.URL, APIKey: testAPIKey, IncludeTrace: false, ConnectTimeout: time.Second, RecoveryInitialBackoff: time.Millisecond, MaxMessageBytes: 1024, }) if err != nil { t.Fatalf("NewHTTPClient() error = %v", err) } result, err := client.StreamMessage(context.Background(), validMessageRequest(), nil) if err != nil { t.Fatalf("StreamMessage() error = %v", err) } if result.Answer != "no-trace answer" { t.Fatalf("Answer = %q, want no-trace answer", result.Answer) } } 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, IncludeTrace: true, 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" }