Files
fire-safety-ymd/internal/integration/superagent/client.go
2026-09-06 01:40:15 +08:00

667 lines
19 KiB
Go

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
}