486 lines
20 KiB
Go
486 lines
20 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 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"
|
|
}
|