package superagent import ( "bytes" "context" "crypto/rand" "encoding/base64" "encoding/json" "errors" "fmt" "io" "mime" "net" "net/http" "net/url" "strconv" "strings" "time" ) const ( defaultConnectTimeout = 15 * time.Second defaultRecoveryInitialBackoff = 250 * time.Millisecond defaultMaxMessageBytes int64 = 64 * 1024 maximumMessageBytes int64 = 16 * 1024 * 1024 maxControlResponseBytes int64 = 1024 * 1024 maxRequestEnvelopeBytes = 256 * 1024 maxIdentifierBytes = 512 maximumRecoveryAttempts = 20 maxRecoveryBackoff = 30 * time.Second ) // HTTPClient implements Client with the SuperAgent Open API HTTP protocol. type HTTPClient struct { enabled bool baseURL *url.URL apiKey string includeTrace bool httpClient *http.Client recoveryMaxAttempts int recoveryInitialBackoff time.Duration maxMessageBytes int64 } var _ Client = (*HTTPClient)(nil) // NewHTTPClient validates the provider boundary and constructs a client. It // deliberately has no overall HTTP timeout because an SSE run can be long-lived; // callers must provide a bounded context. func NewHTTPClient(cfg Config) (*HTTPClient, error) { if cfg.ConnectTimeout == 0 { cfg.ConnectTimeout = defaultConnectTimeout } if cfg.RecoveryInitialBackoff == 0 { cfg.RecoveryInitialBackoff = defaultRecoveryInitialBackoff } if cfg.MaxMessageBytes == 0 { cfg.MaxMessageBytes = defaultMaxMessageBytes } if cfg.ConnectTimeout < 0 || cfg.RecoveryInitialBackoff < 0 || cfg.MaxMessageBytes < 0 || cfg.RecoveryMaxAttempts < 0 { return nil, fmt.Errorf("%w: timing, size, and attempt limits must not be negative", ErrInvalidConfig) } if cfg.MaxMessageBytes > maximumMessageBytes { return nil, fmt.Errorf("%w: message byte limit exceeds the supported maximum", ErrInvalidConfig) } if cfg.RecoveryMaxAttempts > maximumRecoveryAttempts { return nil, fmt.Errorf("%w: recovery attempt limit exceeds the supported maximum", ErrInvalidConfig) } if !cfg.Enabled { return &HTTPClient{ enabled: false, recoveryMaxAttempts: cfg.RecoveryMaxAttempts, recoveryInitialBackoff: cfg.RecoveryInitialBackoff, maxMessageBytes: cfg.MaxMessageBytes, }, nil } baseURL, err := validateBaseURL(cfg.BaseURL) if err != nil { return nil, err } if !validAPIKey(cfg.APIKey) { return nil, fmt.Errorf("%w: API key is required and must be a valid header token", ErrInvalidConfig) } transport := http.DefaultTransport.(*http.Transport).Clone() transport.DialContext = (&net.Dialer{ Timeout: cfg.ConnectTimeout, KeepAlive: 30 * time.Second, }).DialContext return &HTTPClient{ enabled: true, baseURL: baseURL, apiKey: cfg.APIKey, includeTrace: cfg.IncludeTrace, httpClient: &http.Client{ Transport: transport, CheckRedirect: func(_ *http.Request, _ []*http.Request) error { return http.ErrUseLastResponse }, }, recoveryMaxAttempts: cfg.RecoveryMaxAttempts, recoveryInitialBackoff: cfg.RecoveryInitialBackoff, maxMessageBytes: cfg.MaxMessageBytes, }, nil } // CreateSession creates a provider session without sending a message. func (c *HTTPClient) CreateSession(ctx context.Context, request CreateSessionRequest) (Session, error) { if !c.enabled { return Session{}, ErrDisabled } if err := validateCreateSessionRequest(request); err != nil { return Session{}, err } body, err := json.Marshal(struct { ExternalSubjectID string `json:"external_subject_id"` IdempotencyKey string `json:"idempotency_key"` Metadata map[string]any `json:"metadata,omitempty"` }{ ExternalSubjectID: request.ExternalSubjectID, IdempotencyKey: request.IdempotencyKey, Metadata: request.Metadata, }) if err != nil { return Session{}, fmt.Errorf("%w: metadata is not JSON encodable", ErrInvalidRequest) } if len(body) > maxRequestEnvelopeBytes { return Session{}, fmt.Errorf("%w: create-session payload exceeds size limit", ErrInvalidRequest) } httpRequest, err := c.newRequest( ctx, http.MethodPost, c.endpoint("/api/open/agent-sessions"), bytes.NewReader(body), "application/json", request.RequestID, "", ) if err != nil { return Session{}, err } httpRequest.Header.Set("Content-Type", "application/json") response, err := c.httpClient.Do(httpRequest) if err != nil { if ctx.Err() != nil { return Session{}, ctx.Err() } return Session{}, fmt.Errorf("create superagent session: %w", err) } defer response.Body.Close() if err := requireHTTPSuccess(response); err != nil { return Session{}, err } payload, err := readJSONObject(response.Body) if err != nil { return Session{}, err } sessionID := stringValue(payload, "session_id") if sessionID == "" { sessionID = stringValue(payload, "id") } if !validResourceID(sessionID) { return Session{}, fmt.Errorf("%w: create-session response is missing a valid session ID", ErrProtocol) } return Session{ID: sessionID, Status: safeValue(stringValue(payload, "status"))}, nil } // StreamMessage sends one message and returns only after strict stream success. func (c *HTTPClient) StreamMessage( ctx context.Context, request StreamMessageRequest, traceHandler TraceHandler, ) (Result, error) { if !c.enabled { return Result{}, ErrDisabled } if err := c.validateStreamMessageRequest(request); err != nil { return Result{}, err } body, err := json.Marshal(struct { Message string `json:"message"` IdempotencyKey string `json:"idempotency_key"` Metadata map[string]any `json:"metadata,omitempty"` }{ Message: request.Message, IdempotencyKey: request.IdempotencyKey, Metadata: request.Metadata, }) if err != nil { return Result{}, fmt.Errorf("%w: metadata is not JSON encodable", ErrInvalidRequest) } if int64(len(body)) > c.maxMessageBytes+maxRequestEnvelopeBytes { return Result{}, fmt.Errorf("%w: message payload exceeds size limit", ErrInvalidRequest) } streamURL := c.endpoint("/api/open/agent-sessions/" + request.SessionID + "/messages/stream") streamQuery := streamURL.Query() streamQuery.Set("include_trace", strconv.FormatBool(c.includeTrace)) streamURL.RawQuery = streamQuery.Encode() httpRequest, err := c.newRequest( ctx, http.MethodPost, streamURL, bytes.NewReader(body), "text/event-stream", request.RequestID, "", ) if err != nil { return Result{}, err } httpRequest.Header.Set("Content-Type", "application/json") response, err := c.httpClient.Do(httpRequest) if err != nil { if ctx.Err() != nil { return Result{}, ctx.Err() } return Result{}, fmt.Errorf("stream superagent message: %w", err) } if err := requireHTTPSuccess(response); err != nil { response.Body.Close() return Result{}, err } if err := requireEventStream(response); err != nil { response.Body.Close() return Result{}, err } state := newStreamState(request.SessionID, c.includeTrace) if contentLocation := strings.TrimSpace(response.Header.Get("Content-Location")); contentLocation != "" { runURL, locationErr := c.resolveRunURL(httpRequest.URL, request.SessionID, contentLocation) if locationErr != nil { response.Body.Close() return Result{}, locationErr } state.runURL = runURL.String() state.runID = runIDFromURL(runURL) } consumeErr := state.consume(response.Body, traceHandler) closeErr := response.Body.Close() if consumeErr == nil && closeErr != nil { consumeErr = ErrStreamRead } if consumeErr == nil { return state.result() } if ctx.Err() != nil { return Result{}, ctx.Err() } if !errorsIsRecoverableStream(consumeErr) { return Result{}, consumeErr } if state.runURL == "" && validResourceID(state.runID) { state.runURL = c.endpoint("/api/open/agent-sessions/" + request.SessionID + "/runs/" + state.runID).String() } if state.runURL == "" { return Result{}, consumeErr } if err := c.recover(ctx, request.RequestID, state, traceHandler); err != nil { return Result{}, err } return state.result() } func (c *HTTPClient) recover( ctx context.Context, requestID string, state *streamState, traceHandler TraceHandler, ) error { runURL, err := url.Parse(state.runURL) if err != nil || !c.sameOrigin(runURL) || runURL.User != nil || runURL.Fragment != "" { return fmt.Errorf("%w: invalid recovery URL", ErrProtocol) } backoff := c.recoveryInitialBackoff var lastErr error for attempt := 0; attempt < c.recoveryMaxAttempts; attempt++ { if err := waitForRecovery(ctx, backoff); err != nil { return err } status, statusErr := c.queryRunStatus(ctx, runURL, requestID) if statusErr != nil { if ctx.Err() != nil { return ctx.Err() } if !errors.Is(statusErr, ErrStreamRead) && !isRetryableHTTPError(statusErr) { return statusErr } lastErr = statusErr backoff = nextBackoff(backoff) continue } if isFailedRunStatus(status) { return &RunError{Status: status} } eventsErr := c.subscribeRunEvents(ctx, runURL, requestID, state, traceHandler) if eventsErr == nil && state.endSeen { return nil } if ctx.Err() != nil { return ctx.Err() } if eventsErr != nil && !errorsIsRecoverableStream(eventsErr) && !isRetryableHTTPError(eventsErr) { return eventsErr } lastErr = eventsErr if lastErr == nil { lastErr = ErrStreamIncomplete } backoff = nextBackoff(backoff) } if lastErr == nil { lastErr = ErrStreamIncomplete } return fmt.Errorf("%w: %w", ErrRecoveryExhausted, lastErr) } func (c *HTTPClient) queryRunStatus(ctx context.Context, runURL *url.URL, requestID string) (string, error) { request, err := c.newRequest(ctx, http.MethodGet, runURL, nil, "application/json", requestID, "") if err != nil { return "", err } response, err := c.httpClient.Do(request) if err != nil { if ctx.Err() != nil { return "", ctx.Err() } return "", fmt.Errorf("%w: query run request failed", ErrStreamRead) } defer response.Body.Close() if err := requireHTTPSuccess(response); err != nil { return "", err } payload, err := readJSONObject(response.Body) if err != nil { return "", err } status := safeValue(stringValue(payload, "status")) if status == "" { status = "running" } return strings.ToLower(status), nil } func (c *HTTPClient) subscribeRunEvents( ctx context.Context, runURL *url.URL, requestID string, state *streamState, traceHandler TraceHandler, ) error { eventsURL := *runURL eventsURL.Path = strings.TrimRight(eventsURL.Path, "/") + "/events" eventsURL.RawPath = "" request, err := c.newRequest( ctx, http.MethodGet, &eventsURL, nil, "text/event-stream", requestID, state.lastEventID, ) if err != nil { return err } response, err := c.httpClient.Do(request) if err != nil { if ctx.Err() != nil { return ctx.Err() } return fmt.Errorf("%w: subscribe request failed", ErrStreamRead) } defer response.Body.Close() if err := requireHTTPSuccess(response); err != nil { return err } if err := requireEventStream(response); err != nil { return err } return state.consume(response.Body, traceHandler) } func (c *HTTPClient) newRequest( ctx context.Context, method string, target *url.URL, body io.Reader, accept string, requestID string, lastEventID string, ) (*http.Request, error) { if target == nil || !c.sameOrigin(target) || target.User != nil || target.Fragment != "" { return nil, fmt.Errorf("%w: outbound URL is outside the configured origin", ErrInvalidRequest) } if !validOptionalHeader(requestID) || !validOptionalHeader(lastEventID) { return nil, fmt.Errorf("%w: request contains an invalid header value", ErrInvalidRequest) } csrfToken, err := newCSRFToken() if err != nil { return nil, fmt.Errorf("%w: could not create request token", ErrProtocol) } request, err := http.NewRequestWithContext(ctx, method, target.String(), body) if err != nil { return nil, fmt.Errorf("%w: could not construct HTTP request", ErrInvalidRequest) } request.Header.Set("Authorization", "Bearer "+c.apiKey) request.Header.Set("Cache-Control", "no-cache") request.Header.Set("Accept", accept) request.Header.Set("X-CSRF-Token", csrfToken) request.AddCookie(&http.Cookie{Name: "csrf_token", Value: csrfToken}) if requestID != "" { request.Header.Set("X-Request-ID", requestID) } if lastEventID != "" { request.Header.Set("Last-Event-ID", lastEventID) } return request, nil } func (c *HTTPClient) endpoint(path string) *url.URL { target := *c.baseURL target.Path = strings.TrimRight(target.Path, "/") + path target.RawPath = "" target.RawQuery = "" target.Fragment = "" return &target } func (c *HTTPClient) resolveRunURL(requestURL *url.URL, sessionID, contentLocation string) (*url.URL, error) { reference, err := url.Parse(contentLocation) if err != nil || reference.User != nil || reference.Fragment != "" { return nil, fmt.Errorf("%w: invalid Content-Location", ErrProtocol) } resolved := requestURL.ResolveReference(reference) if !c.sameOrigin(resolved) { return nil, fmt.Errorf("%w: cross-origin Content-Location", ErrProtocol) } runID := runIDFromURL(resolved) if !validResourceID(runID) { return nil, fmt.Errorf("%w: Content-Location does not identify a valid run", ErrProtocol) } expectedSuffix := "/agent-sessions/" + sessionID + "/runs/" + runID if !strings.HasSuffix(strings.TrimRight(resolved.Path, "/"), expectedSuffix) { return nil, fmt.Errorf("%w: Content-Location identifies a different session", ErrProtocol) } return resolved, nil } func (c *HTTPClient) sameOrigin(target *url.URL) bool { if c.baseURL == nil || target == nil { return false } return strings.EqualFold(c.baseURL.Scheme, target.Scheme) && strings.EqualFold(c.baseURL.Hostname(), target.Hostname()) && effectivePort(c.baseURL) == effectivePort(target) } func validateBaseURL(raw string) (*url.URL, error) { parsed, err := url.Parse(strings.TrimSpace(raw)) if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { return nil, fmt.Errorf("%w: base URL must be an absolute HTTP(S) URL", ErrInvalidConfig) } if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { return nil, fmt.Errorf("%w: base URL must not contain user information, query, or fragment", ErrInvalidConfig) } parsed.Path = strings.TrimRight(parsed.Path, "/") parsed.RawPath = "" return parsed, nil } func validateCreateSessionRequest(request CreateSessionRequest) error { if strings.TrimSpace(request.ExternalSubjectID) == "" || len(request.ExternalSubjectID) > maxIdentifierBytes { return fmt.Errorf("%w: external subject ID is required and bounded", ErrInvalidRequest) } if strings.TrimSpace(request.IdempotencyKey) == "" || len(request.IdempotencyKey) > maxIdentifierBytes { return fmt.Errorf("%w: idempotency key is required and bounded", ErrInvalidRequest) } if !validOptionalHeader(request.RequestID) { return fmt.Errorf("%w: request ID is not a valid header value", ErrInvalidRequest) } return nil } func (c *HTTPClient) validateStreamMessageRequest(request StreamMessageRequest) error { if !validResourceID(request.SessionID) { return fmt.Errorf("%w: session ID is invalid", ErrInvalidRequest) } if strings.TrimSpace(request.Message) == "" { return fmt.Errorf("%w: message is required", ErrInvalidRequest) } if int64(len([]byte(request.Message))) > c.maxMessageBytes { return fmt.Errorf("%w: message exceeds configured byte limit", ErrInvalidRequest) } if strings.TrimSpace(request.IdempotencyKey) == "" || len(request.IdempotencyKey) > maxIdentifierBytes { return fmt.Errorf("%w: idempotency key is required and bounded", ErrInvalidRequest) } if !validOptionalHeader(request.RequestID) { return fmt.Errorf("%w: request ID is not a valid header value", ErrInvalidRequest) } return nil } func requireHTTPSuccess(response *http.Response) error { if response.StatusCode >= http.StatusOK && response.StatusCode < http.StatusMultipleChoices { return nil } body, _ := readBounded(response.Body, maxControlResponseBytes) return &HTTPStatusError{StatusCode: response.StatusCode, ProviderCode: providerCode(body)} } func requireEventStream(response *http.Response) error { mediaType, _, err := mime.ParseMediaType(response.Header.Get("Content-Type")) if err != nil || !strings.EqualFold(mediaType, "text/event-stream") { return fmt.Errorf("%w: expected text/event-stream response", ErrProtocol) } return nil } func readJSONObject(reader io.Reader) (map[string]any, error) { body, err := readBounded(reader, maxControlResponseBytes) if err != nil { return nil, err } return decodeJSONObject(string(body)) } func readBounded(reader io.Reader, limit int64) ([]byte, error) { body, err := io.ReadAll(io.LimitReader(reader, limit+1)) if err != nil { return nil, fmt.Errorf("%w: could not read control response", ErrProtocol) } if int64(len(body)) > limit { return nil, fmt.Errorf("%w: control response exceeds size limit", ErrProtocol) } return body, nil } func providerCode(body []byte) string { var payload map[string]any if err := json.Unmarshal(body, &payload); err != nil { return "" } code := stringValue(payload, "code") if code == "" { if nested, ok := payload["error"].(map[string]any); ok { code = stringValue(nested, "code") } } return safeValue(code) } func runIDFromURL(runURL *url.URL) string { if runURL == nil { return "" } path := strings.TrimRight(runURL.Path, "/") marker := strings.LastIndex(path, "/runs/") if marker < 0 { return "" } runID := path[marker+len("/runs/"):] if strings.Contains(runID, "/") { return "" } return runID } func newCSRFToken() (string, error) { bytes := make([]byte, 32) if _, err := rand.Read(bytes); err != nil { return "", err } return base64.RawURLEncoding.EncodeToString(bytes), nil } func validResourceID(value string) bool { return safeProviderValue.MatchString(value) } func validOptionalHeader(value string) bool { if len(value) > maxIdentifierBytes { return false } for _, character := range value { if character < 0x20 || character == 0x7f { return false } } return true } func validAPIKey(value string) bool { if value == "" || len(value) > 4096 { return false } for _, character := range value { if character <= 0x20 || character >= 0x7f { return false } } return true } func effectivePort(value *url.URL) string { if port := value.Port(); port != "" { return port } if strings.EqualFold(value.Scheme, "https") { return "443" } return "80" } func isFailedRunStatus(status string) bool { switch strings.ToLower(status) { case "error", "failed", "timeout", "interrupted", "cancelled", "canceled": return true default: return false } } func errorsIsRecoverableStream(err error) bool { return errors.Is(err, ErrStreamIncomplete) || errors.Is(err, ErrStreamRead) } func isRetryableHTTPError(err error) bool { var statusError *HTTPStatusError if !errors.As(err, &statusError) { return false } return statusError.StatusCode == http.StatusRequestTimeout || statusError.StatusCode == http.StatusConflict || statusError.StatusCode == http.StatusTooEarly || statusError.StatusCode == http.StatusTooManyRequests || statusError.StatusCode >= http.StatusInternalServerError } func waitForRecovery(ctx context.Context, duration time.Duration) error { timer := time.NewTimer(duration) defer timer.Stop() select { case <-ctx.Done(): return ctx.Err() case <-timer.C: return nil } } func nextBackoff(current time.Duration) time.Duration { if current >= maxRecoveryBackoff/2 { return maxRecoveryBackoff } return current * 2 }