From 1c3e926819fa1ba017e7dc7751c99ae1bf225457 Mon Sep 17 00:00:00 2001 From: naiba Date: Mon, 20 Jul 2026 04:48:15 +0000 Subject: [PATCH] test(agentcompat): add protocol clients Co-authored-by: naiba/CloudCode --- .../agentcompat/internal/client/client.go | 158 ++++++++++ .../internal/client/client_test.go | 199 ++++++++++++ .../agentcompat/internal/client/errors.go | 101 ++++++ .../agentcompat/internal/client/http.go | 144 +++++++++ .../internal/client/http_capability_test.go | 35 +++ .../internal/client/io_stream_capability.go | 106 +++++++ .../client/io_stream_capability_test.go | 162 ++++++++++ .../client/io_stream_capability_types.go | 126 ++++++++ .../internal/client/io_stream_state.go | 29 ++ .../internal/client/io_stream_state_test.go | 114 +++++++ .../agentcompat/internal/client/mcp.go | 156 ++++++++++ .../internal/client/mcp_filesystem.go | 104 +++++++ .../agentcompat/internal/client/mcp_test.go | 290 ++++++++++++++++++ .../internal/client/sqlite_hold.go | 90 ++++++ .../internal/client/sqlite_hold_test.go | 78 +++++ .../agentcompat/internal/client/transfer.go | 184 +++++++++++ .../internal/client/transfer_rejection.go | 46 +++ .../internal/client/transfer_test.go | 249 +++++++++++++++ .../agentcompat/internal/client/websocket.go | 220 +++++++++++++ .../internal/client/websocket_close_test.go | 169 ++++++++++ .../client/websocket_read_contract_test.go | 236 ++++++++++++++ .../internal/client/websocket_test.go | 230 ++++++++++++++ .../internal/client/websocket_until_test.go | 103 +++++++ 23 files changed, 3329 insertions(+) create mode 100644 integration/agentcompat/internal/client/client.go create mode 100644 integration/agentcompat/internal/client/client_test.go create mode 100644 integration/agentcompat/internal/client/errors.go create mode 100644 integration/agentcompat/internal/client/http.go create mode 100644 integration/agentcompat/internal/client/http_capability_test.go create mode 100644 integration/agentcompat/internal/client/io_stream_capability.go create mode 100644 integration/agentcompat/internal/client/io_stream_capability_test.go create mode 100644 integration/agentcompat/internal/client/io_stream_capability_types.go create mode 100644 integration/agentcompat/internal/client/io_stream_state.go create mode 100644 integration/agentcompat/internal/client/io_stream_state_test.go create mode 100644 integration/agentcompat/internal/client/mcp.go create mode 100644 integration/agentcompat/internal/client/mcp_filesystem.go create mode 100644 integration/agentcompat/internal/client/mcp_test.go create mode 100644 integration/agentcompat/internal/client/sqlite_hold.go create mode 100644 integration/agentcompat/internal/client/sqlite_hold_test.go create mode 100644 integration/agentcompat/internal/client/transfer.go create mode 100644 integration/agentcompat/internal/client/transfer_rejection.go create mode 100644 integration/agentcompat/internal/client/transfer_test.go create mode 100644 integration/agentcompat/internal/client/websocket.go create mode 100644 integration/agentcompat/internal/client/websocket_close_test.go create mode 100644 integration/agentcompat/internal/client/websocket_read_contract_test.go create mode 100644 integration/agentcompat/internal/client/websocket_test.go create mode 100644 integration/agentcompat/internal/client/websocket_until_test.go diff --git a/integration/agentcompat/internal/client/client.go b/integration/agentcompat/internal/client/client.go new file mode 100644 index 00000000..0839b2cf --- /dev/null +++ b/integration/agentcompat/internal/client/client.go @@ -0,0 +1,158 @@ +package client + +import ( + "context" + "errors" + "fmt" + "net/http" + "net/http/cookiejar" + "net/url" + "strings" + "sync/atomic" + "time" + + "github.com/gorilla/websocket" +) + +const ( + defaultRequestTimeout = 10 * time.Second + defaultTransferTimeout = 5 * time.Minute + defaultMaxResponseBytes = int64(8 << 20) + defaultMaxTransferBytes = int64(100 << 20) +) + +var ErrInvalidConfig = errors.New("client: invalid configuration") + +type Config struct { + BaseURL string + HTTPClient *http.Client + WebSocketDialer *websocket.Dialer + BearerToken string + Origin string + RequestTimeout time.Duration + TransferTimeout time.Duration + MaxResponseBytes int64 + MaxTransferBytes int64 +} + +type Client struct { + baseURL *url.URL + httpClient *http.Client + webSocketDialer *websocket.Dialer + bearerToken string + origin string + requestTimeout time.Duration + transferTimeout time.Duration + maxResponseBytes int64 + maxTransferBytes int64 + nextRequestID atomic.Uint64 + requestCount atomic.Uint64 +} + +func (client *Client) RequestCount() uint64 { + return client.requestCount.Load() +} + +func New(config Config) (*Client, error) { + baseURL, err := url.Parse(config.BaseURL) + if err != nil || baseURL.Host == "" || (baseURL.Scheme != "http" && baseURL.Scheme != "https") { + return nil, fmt.Errorf("base URL: %w", ErrInvalidConfig) + } + + requestTimeout := config.RequestTimeout + if requestTimeout == 0 { + requestTimeout = defaultRequestTimeout + } + transferTimeout := config.TransferTimeout + if transferTimeout == 0 { + transferTimeout = defaultTransferTimeout + } + maxResponseBytes := config.MaxResponseBytes + if maxResponseBytes == 0 { + maxResponseBytes = defaultMaxResponseBytes + } + maxTransferBytes := config.MaxTransferBytes + if maxTransferBytes == 0 { + maxTransferBytes = defaultMaxTransferBytes + } + if requestTimeout < 0 || transferTimeout < 0 || maxResponseBytes < 1 || maxTransferBytes < 1 { + return nil, fmt.Errorf("request limits: %w", ErrInvalidConfig) + } + + httpClient := &http.Client{} + if config.HTTPClient != nil { + clone := *config.HTTPClient + httpClient = &clone + } + if httpClient.Jar == nil { + jar, jarErr := cookiejar.New(nil) + if jarErr != nil { + return nil, fmt.Errorf("cookie jar: %w", jarErr) + } + httpClient.Jar = jar + } + httpClient.CheckRedirect = rejectRedirect + + dialer := websocket.DefaultDialer + if config.WebSocketDialer != nil { + dialer = config.WebSocketDialer + } + dialerClone := *dialer + if dialerClone.HandshakeTimeout == 0 || dialerClone.HandshakeTimeout > requestTimeout { + dialerClone.HandshakeTimeout = requestTimeout + } + + origin := strings.TrimSpace(config.Origin) + if origin == "" { + origin = baseURL.Scheme + "://" + baseURL.Host + } + return &Client{ + baseURL: baseURL, + httpClient: httpClient, + webSocketDialer: &dialerClone, + bearerToken: strings.TrimSpace(config.BearerToken), + origin: origin, + requestTimeout: requestTimeout, + transferTimeout: transferTimeout, + maxResponseBytes: maxResponseBytes, + maxTransferBytes: maxTransferBytes, + }, nil +} + +func rejectRedirect(*http.Request, []*http.Request) error { + return ErrRedirect +} + +func (client *Client) requestContext(parent context.Context) (context.Context, context.CancelFunc) { + return context.WithTimeout(parent, client.requestTimeout) +} + +func (client *Client) transferContext(parent context.Context) (context.Context, context.CancelFunc) { + return context.WithTimeout(parent, client.transferTimeout) +} + +func (client *Client) resolvePath(path string) (*url.URL, error) { + reference, err := url.Parse(path) + if err != nil || reference.IsAbs() || reference.Host != "" { + return nil, fmt.Errorf("request path: %w", ErrInvalidConfig) + } + return client.baseURL.ResolveReference(reference), nil +} + +func (client *Client) applyAuthenticatedHeaders(request *http.Request, includeCSRF bool) { + if client.bearerToken != "" { + request.Header.Set("Authorization", "Bearer "+client.bearerToken) + } + if client.origin != "" { + request.Header.Set("Origin", client.origin) + } + if !includeCSRF || client.httpClient.Jar == nil { + return + } + for _, cookie := range client.httpClient.Jar.Cookies(client.baseURL) { + if cookie.Name == "nz-csrf" { + request.Header.Set("X-CSRF-Token", cookie.Value) + return + } + } +} diff --git a/integration/agentcompat/internal/client/client_test.go b/integration/agentcompat/internal/client/client_test.go new file mode 100644 index 00000000..152995c4 --- /dev/null +++ b/integration/agentcompat/internal/client/client_test.go @@ -0,0 +1,199 @@ +package client + +import ( + "context" + "net/http" + "net/http/cookiejar" + "net/http/httptest" + "net/url" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +type semanticRequest struct { + Name string `json:"name"` +} + +type semanticResult struct { + ID uint64 `json:"id"` +} + +func TestClient_RESTSemanticSuccess(t *testing.T) { + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + require.Equal(t, http.MethodPost, request.Method) + require.Equal(t, "Bearer test-token", request.Header.Get("Authorization")) + require.Equal(t, "csrf-value", request.Header.Get("X-CSRF-Token")) + require.NotEmpty(t, request.Header.Get("Origin")) + writer.Header().Set("Content-Type", "application/json") + if requests.Add(1) == 1 { + _, _ = writer.Write([]byte(`{"success":false,"error":"semantic failure"}`)) + return + } + writer.WriteHeader(http.StatusCreated) + _, _ = writer.Write([]byte(`{"success":true,"data":{"id":42}}`)) + })) + t.Cleanup(server.Close) + + jar, err := cookiejar.New(nil) + require.NoError(t, err) + baseURL, err := url.Parse(server.URL) + require.NoError(t, err) + jar.SetCookies(baseURL, []*http.Cookie{{Name: "nz-csrf", Value: "csrf-value"}}) + httpClient := server.Client() + httpClient.Jar = jar + client := newTestClient(t, Config{ + BaseURL: server.URL, + HTTPClient: httpClient, + BearerToken: "test-token", + Origin: server.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + }) + + requestBody := semanticRequest{Name: "probe"} + _, err = REST[semanticRequest, semanticResult](context.Background(), client, RESTRequest[semanticRequest]{ + Method: http.MethodPost, + Path: "/semantic", + Body: &requestBody, + }) + require.ErrorIs(t, err, ErrSemanticFailure) + + result, err := REST[semanticRequest, semanticResult](context.Background(), client, RESTRequest[semanticRequest]{ + Method: http.MethodPost, + Path: "/semantic", + Body: &requestBody, + }) + require.NoError(t, err) + require.Equal(t, uint64(42), result.ID) +} + +func TestClient_RESTRejectsOversize(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"success":true,"data":{"id":42},"padding":"` + strings.Repeat("x", 256) + `"}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 64}) + _, err := REST[struct{}, semanticResult](context.Background(), client, RESTRequest[struct{}]{ + Method: http.MethodGet, + Path: "/oversize", + }) + require.ErrorIs(t, err, ErrResponseTooLarge) +} + +func TestClient_RESTClassifiesNonJSONStatus(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.WriteHeader(http.StatusUnauthorized) + _, _ = writer.Write([]byte("unauthorized")) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + _, err := REST[struct{}, semanticResult](context.Background(), client, RESTRequest[struct{}]{ + Method: http.MethodGet, + Path: "/unauthorized", + }) + require.ErrorIs(t, err, ErrUnauthorized) +} + +func TestClient_RESTRejectsCrossOriginRedirect(t *testing.T) { + var foreignRequests atomic.Int32 + foreignServer := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + foreignRequests.Add(1) + })) + t.Cleanup(foreignServer.Close) + dashboardServer := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + http.Redirect(writer, request, foreignServer.URL+"/credentials", http.StatusTemporaryRedirect) + })) + t.Cleanup(dashboardServer.Close) + + client := newTestClient(t, Config{ + BaseURL: dashboardServer.URL, + BearerToken: "redirect-secret", + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + }) + _, err := REST[struct{}, semanticResult](context.Background(), client, RESTRequest[struct{}]{Method: http.MethodGet, Path: "/redirect"}) + require.ErrorIs(t, err, ErrRedirect) + require.Zero(t, foreignRequests.Load()) +} + +func TestClient_RESTDeadline(t *testing.T) { + // A tiny deadline races handler scheduling under -race; this barrier proves + // the REST request is in flight before evaluating deadline behavior. + requestEntered := make(chan struct{}) + requestCancelled := make(chan struct{}) + releaseHandler := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(_ http.ResponseWriter, request *http.Request) { + close(requestEntered) + select { + case <-request.Context().Done(): + close(requestCancelled) + case <-releaseHandler: + } + })) + t.Cleanup(func() { + close(releaseHandler) + server.Close() + }) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + result := make(chan error, 1) + go func() { + _, err := REST[struct{}, semanticResult](context.Background(), client, RESTRequest[struct{}]{ + Method: http.MethodGet, + Path: "/deadline", + }) + result <- err + }() + waitContext, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + select { + case <-requestEntered: + case <-waitContext.Done(): + t.Fatal("server did not receive REST request") + } + var err error + select { + case err = <-result: + case <-waitContext.Done(): + t.Fatal("REST request did not reach its deadline") + } + require.ErrorIs(t, err, context.DeadlineExceeded) + select { + case <-requestCancelled: + case <-waitContext.Done(): + t.Fatal("server request context was not cancelled") + } +} + +func TestClient_RESTRejectsRedirectWithoutForwardingSensitiveHeaders(t *testing.T) { + var redirected atomic.Int32 + foreignServer := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + redirected.Add(1) + })) + t.Cleanup(foreignServer.Close) + + dashboardServer := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + http.Redirect(writer, &http.Request{}, foreignServer.URL+"/capture", http.StatusTemporaryRedirect) + })) + t.Cleanup(dashboardServer.Close) + + client := newTestClient(t, Config{BaseURL: dashboardServer.URL, BearerToken: "test-token", RequestTimeout: time.Second, MaxResponseBytes: 1024}) + _, err := REST[struct{}, semanticResult](context.Background(), client, RESTRequest[struct{}]{Method: http.MethodGet, Path: "/redirect"}) + require.ErrorIs(t, err, ErrRedirect) + require.Zero(t, redirected.Load()) +} + +func newTestClient(t *testing.T, config Config) *Client { + t.Helper() + client, err := New(config) + require.NoError(t, err) + return client +} diff --git a/integration/agentcompat/internal/client/errors.go b/integration/agentcompat/internal/client/errors.go new file mode 100644 index 00000000..8711c05f --- /dev/null +++ b/integration/agentcompat/internal/client/errors.go @@ -0,0 +1,101 @@ +package client + +import ( + "encoding/json" + "errors" + "fmt" + "regexp" +) + +var ( + ErrHTTPStatus = errors.New("client: HTTP status failure") + ErrSemanticFailure = errors.New("client: semantic failure") + ErrResponseTooLarge = errors.New("client: response too large") + ErrTransferTooLarge = errors.New("client: transfer too large") + ErrTransferExpired = errors.New("client: transfer URL expired") + ErrUnauthorized = errors.New("client: unauthorized") + ErrRedirect = errors.New("client: redirect rejected") + ErrJSONRPC = errors.New("client: JSON-RPC failure") + ErrToolFailure = errors.New("client: MCP tool failure") +) + +var ( + authorizationPattern = regexp.MustCompile(`(?i)(authorization\s*[:=]\s*(?:bearer\s+)?)[^\s,;"']+`) + bearerPattern = regexp.MustCompile(`(?i)(\bbearer\s+)[A-Za-z0-9._~+/=-]+`) + transferTokenPattern = regexp.MustCompile(`(?i)(/mcp/(?:download|upload)/)[^?\s]+`) + credentialPattern = regexp.MustCompile(`(?i)(["']?(?:x-csrf-token|csrf|token|jwt[_-]?(?:secret(?:[_-]?key)?|token)?|pat|api[_-]?(?:key|token)|access[_-]?token|agent[_-]?secret(?:[_-]?key)?|client[_-]?secret|password|credential|signature)["']?\s*[:=]\s*["']?)[^"'\s,;&}]+(["']?)`) + querySecretPattern = regexp.MustCompile(`(?i)([?&](?:token|access_token|api_key|jwt|pat|secret|authorization|sig|signature|x-amz-signature)=)[^&#\s]+`) + jwtPattern = regexp.MustCompile(`\beyJ[A-Za-z0-9_-]+\.[A-Za-z0-9_-]+\.[A-Za-z0-9_-]+\b`) +) + +type HTTPError struct { + StatusCode int + Message string +} + +type WebSocketHandshakeError struct { + StatusCode int + Message string +} + +func (err *WebSocketHandshakeError) Error() string { + return fmt.Sprintf("WebSocket handshake: status %d: %s", err.StatusCode, Redact(err.Message)) +} + +type WebSocketCloseError struct { + Code int + Text string +} + +func (err *WebSocketCloseError) Error() string { + return fmt.Sprintf("WebSocket closed: code %d: %s", err.Code, Redact(err.Text)) +} + +func (err *HTTPError) Error() string { + if err.Message == "" { + return fmt.Sprintf("%s: status %d", ErrHTTPStatus, err.StatusCode) + } + return fmt.Sprintf("%s: status %d: %s", ErrHTTPStatus, err.StatusCode, Redact(err.Message)) +} + +func (err *HTTPError) Is(target error) bool { + return target == ErrHTTPStatus || (target == ErrUnauthorized && err.StatusCode == 401) +} + +type RPCError struct { + Code int `json:"code"` + Message string `json:"message"` +} + +type ToolFailure struct { + Message string + StructuredContent json.RawMessage +} + +func (err *ToolFailure) Error() string { + if err.Message == "" { + return ErrToolFailure.Error() + } + return fmt.Sprintf("%s: %s", ErrToolFailure, Redact(err.Message)) +} + +func (err *ToolFailure) Is(target error) bool { + return target == ErrToolFailure +} + +func (err *RPCError) Error() string { + return fmt.Sprintf("%s: code %d: %s", ErrJSONRPC, err.Code, Redact(err.Message)) +} + +func (err *RPCError) Is(target error) bool { + return target == ErrJSONRPC || (target == ErrUnauthorized && err.Code == -32001) +} + +func Redact(value string) string { + redacted := authorizationPattern.ReplaceAllString(value, `${1}[REDACTED]`) + redacted = bearerPattern.ReplaceAllString(redacted, `${1}[REDACTED]`) + redacted = credentialPattern.ReplaceAllString(redacted, `${1}[REDACTED]${2}`) + redacted = querySecretPattern.ReplaceAllString(redacted, `${1}[REDACTED]`) + redacted = jwtPattern.ReplaceAllString(redacted, `[REDACTED]`) + return transferTokenPattern.ReplaceAllString(redacted, `${1}[REDACTED]`) +} diff --git a/integration/agentcompat/internal/client/http.go b/integration/agentcompat/internal/client/http.go new file mode 100644 index 00000000..629e4746 --- /dev/null +++ b/integration/agentcompat/internal/client/http.go @@ -0,0 +1,144 @@ +package client + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + + "github.com/nezhahq/nezha/pkg/agentcompatcontract" +) + +type CommonResponse[T any] struct { + Success bool `json:"success"` + Data T `json:"data"` + Error string `json:"error"` +} + +type RESTRequest[T any] struct { + Method string + Path string + Body *T + IOStreamCapability IOStreamCapability +} + +type LoginRequest struct { + Username string `json:"username"` + Password string `json:"password"` +} + +type LoginResponse struct { + Token string `json:"token"` + Expire string `json:"expire"` +} + +func (client *Client) Login(ctx context.Context, request LoginRequest) (LoginResponse, error) { + return REST[LoginRequest, LoginResponse](ctx, client, RESTRequest[LoginRequest]{ + Method: http.MethodPost, + Path: "/api/v1/login", + Body: &request, + }) +} + +func REST[Request, Response any](ctx context.Context, client *Client, request RESTRequest[Request]) (Response, error) { + var zero Response + requestURL, err := client.resolvePath(request.Path) + if err != nil { + return zero, err + } + + var body io.Reader + if request.Body != nil { + encoded, marshalErr := json.Marshal(request.Body) + if marshalErr != nil { + return zero, fmt.Errorf("encode REST request: %w", marshalErr) + } + body = bytes.NewReader(encoded) + } + + requestContext, cancel := client.requestContext(ctx) + defer cancel() + httpRequest, err := http.NewRequestWithContext(requestContext, request.Method, requestURL.String(), body) + if err != nil { + return zero, fmt.Errorf("create REST request: %w", err) + } + if request.Body != nil { + httpRequest.Header.Set("Content-Type", "application/json") + } + client.applyAuthenticatedHeaders(httpRequest, true) + if request.IOStreamCapability.Value() != "" { + httpRequest.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, request.IOStreamCapability.Value()) + } + + status, responseBody, err := client.execute(httpRequest, client.maxResponseBytes) + if err != nil { + return zero, err + } + if status < 200 || status >= 300 { + var envelope CommonResponse[Response] + if json.Unmarshal(responseBody, &envelope) == nil { + return zero, &HTTPError{StatusCode: status, Message: Redact(envelope.Error)} + } + return zero, &HTTPError{StatusCode: status, Message: Redact(string(responseBody))} + } + var envelope CommonResponse[Response] + if err := json.Unmarshal(responseBody, &envelope); err != nil { + return zero, fmt.Errorf("decode REST response: %w", err) + } + if !envelope.Success { + return zero, fmt.Errorf("%w: %s", ErrSemanticFailure, Redact(envelope.Error)) + } + return envelope.Data, nil +} + +func DoREST[Request, Response any](ctx context.Context, client *Client, request RESTRequest[Request]) (Response, error) { + return REST[Request, Response](ctx, client, request) +} + +func (client *Client) execute(request *http.Request, maxBytes int64) (int, []byte, error) { + client.requestCount.Add(1) + response, err := client.httpClient.Do(request) + if err != nil { + if request.Context().Err() != nil { + return 0, nil, fmt.Errorf("HTTP request: %w", request.Context().Err()) + } + return 0, nil, errorsNewRedacted("HTTP request", err) + } + defer response.Body.Close() + + body, err := readBounded(response.Body, maxBytes) + if err != nil { + return response.StatusCode, nil, err + } + return response.StatusCode, body, nil +} + +func readBounded(reader io.Reader, maxBytes int64) ([]byte, error) { + body, err := io.ReadAll(io.LimitReader(reader, maxBytes+1)) + if err != nil { + return nil, fmt.Errorf("read response: %w", err) + } + if int64(len(body)) > maxBytes { + return nil, ErrResponseTooLarge + } + return body, nil +} + +func errorsNewRedacted(operation string, err error) error { + return &redactedOperationError{operation: operation, cause: err} +} + +type redactedOperationError struct { + operation string + cause error +} + +func (err *redactedOperationError) Error() string { + return fmt.Sprintf("%s: %s", err.operation, Redact(err.cause.Error())) +} + +func (err *redactedOperationError) Unwrap() error { + return err.cause +} diff --git a/integration/agentcompat/internal/client/http_capability_test.go b/integration/agentcompat/internal/client/http_capability_test.go new file mode 100644 index 00000000..14949cf1 --- /dev/null +++ b/integration/agentcompat/internal/client/http_capability_test.go @@ -0,0 +1,35 @@ +package client + +import ( + "context" + "encoding/base64" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + "github.com/stretchr/testify/require" +) + +func TestRESTTypedCapabilityHeaderOnlyAttachesWhenRequested(t *testing.T) { + raw := base64.RawURLEncoding.EncodeToString([]byte(strings.Repeat("h", 32))) + capability, err := ParseIOStreamCapability(raw) + require.NoError(t, err) + seen := make(chan string, 2) + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + seen <- request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader) + writer.Header().Set("Content-Type", "application/json") + _, err := writer.Write([]byte(`{"success":true,"data":{}}`)) + require.NoError(t, err) + })) + t.Cleanup(server.Close) + transport := newTestClient(t, Config{BaseURL: server.URL, BearerToken: "private-pat"}) + + _, err = DoREST[struct{}, struct{}](context.Background(), transport, RESTRequest[struct{}]{Method: http.MethodPost, Path: "/ordinary"}) + require.NoError(t, err) + _, err = DoREST[struct{}, struct{}](context.Background(), transport, RESTRequest[struct{}]{Method: http.MethodPost, Path: "/create", IOStreamCapability: capability}) + require.NoError(t, err) + require.Empty(t, <-seen) + require.Equal(t, raw, <-seen) +} diff --git a/integration/agentcompat/internal/client/io_stream_capability.go b/integration/agentcompat/internal/client/io_stream_capability.go new file mode 100644 index 00000000..2a2be009 --- /dev/null +++ b/integration/agentcompat/internal/client/io_stream_capability.go @@ -0,0 +1,106 @@ +package client + +import ( + "context" + "errors" + "net/http" + "strings" +) + +const ( + ioStreamCapabilityRegisterPath = "/agentcompat/io-stream-capability/register" + ioStreamCapabilityWaitPath = "/agentcompat/io-stream-capability/wait" + ioStreamCapabilityCancelPath = "/agentcompat/io-stream-capability/cancel" + ioStreamCapabilityUnregisterPath = "/agentcompat/io-stream-capability/unregister" +) + +const ( + ioStreamCapabilityInvalidMessage = "agentcompat capability request is invalid" + ioStreamCapabilityUnavailableMessage = "agentcompat capability is not available" + ioStreamCapabilityConflictMessage = "agentcompat capability is active" + ioStreamCapabilityCleanupMessage = "agentcompat capability cleanup failed" +) + +type IOStreamCapabilityClient struct { + transport *Client +} + +func (client *Client) IOStreamCapabilities() IOStreamCapabilityClient { + return IOStreamCapabilityClient{transport: client} +} + +func (client IOStreamCapabilityClient) Register(ctx context.Context, request IOStreamCapabilityRegisterRequest) (IOStreamCapabilityRegisterResponse, error) { + if err := validateIOStreamCapabilityIdentity(request.Purpose, request.ServerID, request.ResourceID); err != nil { + return IOStreamCapabilityRegisterResponse{}, err + } + response, err := DoREST[IOStreamCapabilityRegisterRequest, IOStreamCapabilityRegisterResponse](ctx, client.transport, RESTRequest[IOStreamCapabilityRegisterRequest]{ + Method: http.MethodPost, Path: ioStreamCapabilityRegisterPath, Body: &request, + }) + if err != nil { + return IOStreamCapabilityRegisterResponse{}, mapIOStreamCapabilityError(err) + } + if response.Capability.value == "" { + return IOStreamCapabilityRegisterResponse{}, ErrIOStreamCapabilityUnavailable + } + return response, nil +} + +func (client IOStreamCapabilityClient) Wait(ctx context.Context, request IOStreamCapabilityWaitRequest) (IOStreamCapabilityWaitResponse, error) { + access := IOStreamCapabilityAccessRequest(request) + if err := validateIOStreamCapabilityAccess(access); err != nil { + return IOStreamCapabilityWaitResponse{}, err + } + response, err := DoREST[IOStreamCapabilityWaitRequest, IOStreamCapabilityWaitResponse](ctx, client.transport, RESTRequest[IOStreamCapabilityWaitRequest]{ + Method: http.MethodPost, Path: ioStreamCapabilityWaitPath, Body: &request, + }) + if err != nil { + return IOStreamCapabilityWaitResponse{}, mapIOStreamCapabilityError(err) + } + if response.StreamID.value == "" { + return IOStreamCapabilityWaitResponse{}, ErrIOStreamCapabilityUnavailable + } + return response, nil +} + +func (client IOStreamCapabilityClient) Cancel(ctx context.Context, request IOStreamCapabilityAccessRequest) error { + if err := validateIOStreamCapabilityAccess(request); err != nil { + return err + } + _, err := DoREST[IOStreamCapabilityAccessRequest, ioStreamCapabilityEmptyResponse](ctx, client.transport, RESTRequest[IOStreamCapabilityAccessRequest]{ + Method: http.MethodPost, Path: ioStreamCapabilityCancelPath, Body: &request, + }) + return mapIOStreamCapabilityError(err) +} + +func (client IOStreamCapabilityClient) Unregister(ctx context.Context, request IOStreamCapabilityAccessRequest) error { + if err := validateIOStreamCapabilityAccess(request); err != nil { + return err + } + _, err := DoREST[IOStreamCapabilityAccessRequest, ioStreamCapabilityEmptyResponse](ctx, client.transport, RESTRequest[IOStreamCapabilityAccessRequest]{ + Method: http.MethodPost, Path: ioStreamCapabilityUnregisterPath, Body: &request, + }) + return mapIOStreamCapabilityError(err) +} + +func mapIOStreamCapabilityError(err error) error { + if err == nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return err + } + if errors.Is(err, ErrUnauthorized) { + return ErrUnauthorized + } + if errors.Is(err, ErrSemanticFailure) { + message := err.Error() + switch { + case strings.HasSuffix(message, ioStreamCapabilityInvalidMessage): + return ErrInvalidIOStreamCapabilityRequest + case strings.HasSuffix(message, ioStreamCapabilityConflictMessage): + return ErrIOStreamCapabilityConflict + case strings.HasSuffix(message, ioStreamCapabilityCleanupMessage): + return ErrIOStreamCapabilityCleanup + case strings.HasSuffix(message, ioStreamCapabilityUnavailableMessage): + return ErrIOStreamCapabilityUnavailable + } + } + return ErrIOStreamCapabilityUnavailable +} diff --git a/integration/agentcompat/internal/client/io_stream_capability_test.go b/integration/agentcompat/internal/client/io_stream_capability_test.go new file mode 100644 index 00000000..03a0b470 --- /dev/null +++ b/integration/agentcompat/internal/client/io_stream_capability_test.go @@ -0,0 +1,162 @@ +package client + +import ( + "context" + "encoding/base64" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestIOStreamCapabilityClientUsesTypedPATAuthenticatedWireContract(t *testing.T) { + rawCapability := base64.RawURLEncoding.EncodeToString([]byte(strings.Repeat("c", 32))) + requestNumber := 0 + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + requestNumber++ + require.Equal(t, "Bearer private-pat", request.Header.Get("Authorization")) + writer.Header().Set("Content-Type", "application/json") + switch requestNumber { + case 1: + require.Equal(t, "/agentcompat/io-stream-capability/register", request.URL.Path) + var body IOStreamCapabilityRegisterRequest + require.NoError(t, json.NewDecoder(request.Body).Decode(&body)) + require.Equal(t, IOStreamCapabilityPurposeTerminal, body.Purpose) + require.Equal(t, uint64(7), body.ServerID) + require.Zero(t, body.ResourceID) + _, err := writer.Write([]byte(`{"success":true,"data":{"capability":"` + rawCapability + `"}}`)) + require.NoError(t, err) + case 2: + require.Equal(t, "/agentcompat/io-stream-capability/wait", request.URL.Path) + var body IOStreamCapabilityWaitRequest + require.NoError(t, json.NewDecoder(request.Body).Decode(&body)) + require.Equal(t, rawCapability, body.Capability.Value()) + _, err := writer.Write([]byte(`{"success":true,"data":{"stream_id":"private-stream"}}`)) + require.NoError(t, err) + case 3, 4: + expectedPath := "/agentcompat/io-stream-capability/cancel" + if requestNumber == 4 { + expectedPath = "/agentcompat/io-stream-capability/unregister" + } + require.Equal(t, expectedPath, request.URL.Path) + _, err := writer.Write([]byte(`{"success":true,"data":{}}`)) + require.NoError(t, err) + default: + t.Fatalf("unexpected request %d", requestNumber) + } + })) + t.Cleanup(server.Close) + transport := newTestClient(t, Config{BaseURL: server.URL, BearerToken: "private-pat"}) + capabilities := transport.IOStreamCapabilities() + + registered, err := capabilities.Register(context.Background(), IOStreamCapabilityRegisterRequest{Purpose: IOStreamCapabilityPurposeTerminal, ServerID: 7}) + require.NoError(t, err) + require.Equal(t, rawCapability, registered.Capability.Value()) + waited, err := capabilities.Wait(context.Background(), IOStreamCapabilityWaitRequest{Capability: registered.Capability, Purpose: IOStreamCapabilityPurposeTerminal, ServerID: 7}) + require.NoError(t, err) + require.Equal(t, "private-stream", waited.StreamID.Value()) + access := IOStreamCapabilityAccessRequest{Capability: registered.Capability, Purpose: IOStreamCapabilityPurposeTerminal, ServerID: 7} + require.NoError(t, capabilities.Cancel(context.Background(), access)) + require.NoError(t, capabilities.Unregister(context.Background(), access)) + require.Equal(t, 4, requestNumber) +} + +func TestIOStreamCapabilityClientValidatesBeforeDispatch(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + t.Fatal("invalid request must not dispatch") + })) + t.Cleanup(server.Close) + transport := newTestClient(t, Config{BaseURL: server.URL}) + capabilities := transport.IOStreamCapabilities() + + _, err := capabilities.Register(context.Background(), IOStreamCapabilityRegisterRequest{Purpose: "unknown", ServerID: 7}) + require.ErrorIs(t, err, ErrInvalidIOStreamCapabilityRequest) + _, err = capabilities.Register(context.Background(), IOStreamCapabilityRegisterRequest{Purpose: IOStreamCapabilityPurposeNAT, ServerID: 7}) + require.ErrorIs(t, err, ErrInvalidIOStreamCapabilityRequest) + _, err = capabilities.Wait(context.Background(), IOStreamCapabilityWaitRequest{Purpose: IOStreamCapabilityPurposeTerminal, ServerID: 7}) + require.ErrorIs(t, err, ErrInvalidIOStreamCapabilityRequest) + require.Zero(t, transport.RequestCount()) +} + +func TestIOStreamCapabilityClientErrorsNeverEchoSensitiveValues(t *testing.T) { + rawCapability := base64.RawURLEncoding.EncodeToString([]byte(strings.Repeat("s", 32))) + capability, err := ParseIOStreamCapability(rawCapability) + require.NoError(t, err) + privateValues := []string{rawCapability, "private-stream", "private-pat", "Authorization", "private-creator", "private-server"} + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, writeErr := writer.Write([]byte(`{"success":false,"error":"capability ` + rawCapability + ` stream private-stream Authorization: Bearer private-pat creator private-creator server private-server"}`)) + require.NoError(t, writeErr) + })) + t.Cleanup(server.Close) + transport := newTestClient(t, Config{BaseURL: server.URL, BearerToken: "private-pat"}) + access := IOStreamCapabilityAccessRequest{Capability: capability, Purpose: IOStreamCapabilityPurposeTerminal, ServerID: 7} + + for _, invoke := range []func() error{ + func() error { + _, callErr := transport.IOStreamCapabilities().Wait(context.Background(), IOStreamCapabilityWaitRequest(access)) + return callErr + }, + func() error { return transport.IOStreamCapabilities().Cancel(context.Background(), access) }, + func() error { return transport.IOStreamCapabilities().Unregister(context.Background(), access) }, + } { + callErr := invoke() + require.ErrorIs(t, callErr, ErrIOStreamCapabilityUnavailable) + for _, privateValue := range privateValues { + require.NotContains(t, callErr.Error(), privateValue) + } + } +} + +func TestIOStreamCapabilityClientPreservesCancellationIdentity(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + select { + case <-request.Context().Done(): + case <-time.After(time.Second): + t.Error("request context was not canceled") + } + })) + t.Cleanup(server.Close) + transport := newTestClient(t, Config{BaseURL: server.URL}) + rawCapability := base64.RawURLEncoding.EncodeToString([]byte(strings.Repeat("x", 32))) + capability, err := ParseIOStreamCapability(rawCapability) + require.NoError(t, err) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err = transport.IOStreamCapabilities().Wait(ctx, IOStreamCapabilityWaitRequest{Capability: capability, Purpose: IOStreamCapabilityPurposeTerminal, ServerID: 7}) + require.True(t, errors.Is(err, context.Canceled)) +} + +func TestIOStreamCapabilityClientMapsTypedNonsecretErrors(t *testing.T) { + messages := []struct { + message string + target error + }{ + {message: ioStreamCapabilityConflictMessage, target: ErrIOStreamCapabilityConflict}, + {message: ioStreamCapabilityCleanupMessage, target: ErrIOStreamCapabilityCleanup}, + {message: ioStreamCapabilityUnavailableMessage, target: ErrIOStreamCapabilityUnavailable}, + } + for _, testCase := range messages { + t.Run(testCase.message, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, err := writer.Write([]byte(`{"success":false,"error":"` + testCase.message + `"}`)) + require.NoError(t, err) + })) + t.Cleanup(server.Close) + transport := newTestClient(t, Config{BaseURL: server.URL}) + rawCapability := base64.RawURLEncoding.EncodeToString([]byte(strings.Repeat("m", 32))) + capability, err := ParseIOStreamCapability(rawCapability) + require.NoError(t, err) + callErr := transport.IOStreamCapabilities().Unregister(context.Background(), IOStreamCapabilityAccessRequest{Capability: capability, Purpose: IOStreamCapabilityPurposeTerminal, ServerID: 7}) + require.ErrorIs(t, callErr, testCase.target) + require.Equal(t, testCase.target.Error(), callErr.Error()) + }) + } +} diff --git a/integration/agentcompat/internal/client/io_stream_capability_types.go b/integration/agentcompat/internal/client/io_stream_capability_types.go new file mode 100644 index 00000000..90fdc1c3 --- /dev/null +++ b/integration/agentcompat/internal/client/io_stream_capability_types.go @@ -0,0 +1,126 @@ +package client + +import ( + "encoding/base64" + "encoding/json" + "errors" +) + +var ( + ErrInvalidIOStreamCapabilityRequest = errors.New("client: invalid IOStream capability request") + ErrIOStreamCapabilityUnavailable = errors.New("client: IOStream capability unavailable") + ErrIOStreamCapabilityConflict = errors.New("client: IOStream capability active") + ErrIOStreamCapabilityCleanup = errors.New("client: IOStream capability cleanup failed") +) + +type IOStreamCapabilityPurpose string + +const ( + IOStreamCapabilityPurposeTerminal IOStreamCapabilityPurpose = "terminal" + IOStreamCapabilityPurposeFileManager IOStreamCapabilityPurpose = "file_manager" + IOStreamCapabilityPurposeNAT IOStreamCapabilityPurpose = "nat" +) + +type IOStreamCapability struct { + value string +} + +func ParseIOStreamCapability(value string) (IOStreamCapability, error) { + raw, err := base64.RawURLEncoding.DecodeString(value) + if err != nil || len(raw) != 32 { + return IOStreamCapability{}, ErrInvalidIOStreamCapabilityRequest + } + return IOStreamCapability{value: value}, nil +} + +func (capability IOStreamCapability) Value() string { + return capability.value +} + +func (capability IOStreamCapability) MarshalJSON() ([]byte, error) { + if capability.value == "" { + return nil, ErrInvalidIOStreamCapabilityRequest + } + return json.Marshal(capability.value) +} + +func (capability *IOStreamCapability) UnmarshalJSON(data []byte) error { + var value string + if err := json.Unmarshal(data, &value); err != nil { + return ErrInvalidIOStreamCapabilityRequest + } + parsed, err := ParseIOStreamCapability(value) + if err != nil { + return err + } + *capability = parsed + return nil +} + +type IOStreamID struct { + value string +} + +func (streamID IOStreamID) Value() string { + return streamID.value +} + +func (streamID *IOStreamID) UnmarshalJSON(data []byte) error { + var value string + if err := json.Unmarshal(data, &value); err != nil || value == "" { + return ErrIOStreamCapabilityUnavailable + } + streamID.value = value + return nil +} + +type IOStreamCapabilityRegisterRequest struct { + Purpose IOStreamCapabilityPurpose `json:"purpose"` + ServerID uint64 `json:"server_id"` + ResourceID uint64 `json:"resource_id,omitempty"` +} + +type IOStreamCapabilityRegisterResponse struct { + Capability IOStreamCapability `json:"capability"` +} + +type IOStreamCapabilityAccessRequest struct { + Capability IOStreamCapability `json:"capability"` + Purpose IOStreamCapabilityPurpose `json:"purpose"` + ServerID uint64 `json:"server_id"` + ResourceID uint64 `json:"resource_id,omitempty"` +} + +type IOStreamCapabilityWaitRequest IOStreamCapabilityAccessRequest + +type IOStreamCapabilityWaitResponse struct { + StreamID IOStreamID `json:"stream_id"` +} + +type ioStreamCapabilityEmptyResponse struct{} + +func validateIOStreamCapabilityIdentity(purpose IOStreamCapabilityPurpose, serverID, resourceID uint64) error { + if serverID == 0 { + return ErrInvalidIOStreamCapabilityRequest + } + switch purpose { + case IOStreamCapabilityPurposeTerminal, IOStreamCapabilityPurposeFileManager: + if resourceID != 0 { + return ErrInvalidIOStreamCapabilityRequest + } + case IOStreamCapabilityPurposeNAT: + if resourceID == 0 { + return ErrInvalidIOStreamCapabilityRequest + } + default: + return ErrInvalidIOStreamCapabilityRequest + } + return nil +} + +func validateIOStreamCapabilityAccess(request IOStreamCapabilityAccessRequest) error { + if request.Capability.value == "" { + return ErrInvalidIOStreamCapabilityRequest + } + return validateIOStreamCapabilityIdentity(request.Purpose, request.ServerID, request.ResourceID) +} diff --git a/integration/agentcompat/internal/client/io_stream_state.go b/integration/agentcompat/internal/client/io_stream_state.go new file mode 100644 index 00000000..1321e03d --- /dev/null +++ b/integration/agentcompat/internal/client/io_stream_state.go @@ -0,0 +1,29 @@ +package client + +import ( + "context" + "net/http" +) + +type IOStreamState struct { + Count int `json:"count"` + Generation uint64 `json:"generation"` +} + +type IOStreamStateExpectation struct { + ExpectedCount *int `json:"expected_count,omitempty"` + PresentStreamID string `json:"present_stream_id,omitempty"` + AbsentStreamID string `json:"absent_stream_id,omitempty"` +} + +func ExpectedIOStreamCount(count int) *int { + return &count +} + +func (client *Client) IOStreamState(ctx context.Context) (IOStreamState, error) { + return DoREST[struct{}, IOStreamState](ctx, client, RESTRequest[struct{}]{Method: http.MethodGet, Path: "/agentcompat/io-stream-state"}) +} + +func (client *Client) WaitForIOStreamState(ctx context.Context, expectation IOStreamStateExpectation) (IOStreamState, error) { + return DoREST[IOStreamStateExpectation, IOStreamState](ctx, client, RESTRequest[IOStreamStateExpectation]{Method: http.MethodPost, Path: "/agentcompat/io-stream-state", Body: &expectation}) +} diff --git a/integration/agentcompat/internal/client/io_stream_state_test.go b/integration/agentcompat/internal/client/io_stream_state_test.go new file mode 100644 index 00000000..0a0bb71e --- /dev/null +++ b/integration/agentcompat/internal/client/io_stream_state_test.go @@ -0,0 +1,114 @@ +package client + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestClientIOStreamStateHelpersUseTypedRESTContracts(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.Header().Set("Content-Type", "application/json") + switch request.Method { + case http.MethodGet: + require.Equal(t, "/agentcompat/io-stream-state", request.URL.Path) + _, err := writer.Write([]byte(`{"success":true,"data":{"count":0,"generation":4}}`)) + require.NoError(t, err) + case http.MethodPost: + var payload map[string]json.RawMessage + require.NoError(t, json.NewDecoder(request.Body).Decode(&payload)) + var expectedCount int + require.NoError(t, json.Unmarshal(payload["expected_count"], &expectedCount)) + require.Equal(t, 0, expectedCount) + var absentStreamID string + require.NoError(t, json.Unmarshal(payload["absent_stream_id"], &absentStreamID)) + require.Equal(t, "stream-id", absentStreamID) + var presentStreamID string + require.NoError(t, json.Unmarshal(payload["present_stream_id"], &presentStreamID)) + require.Equal(t, "present-id", presentStreamID) + _, err := writer.Write([]byte(`{"success":true,"data":{"count":0,"generation":5}}`)) + require.NoError(t, err) + default: + writer.WriteHeader(http.StatusMethodNotAllowed) + } + })) + t.Cleanup(server.Close) + httpClient := newTestClient(t, Config{BaseURL: server.URL}) + snapshot, err := httpClient.IOStreamState(context.Background()) + require.NoError(t, err) + require.Equal(t, IOStreamState{Count: 0, Generation: 4}, snapshot) + waited, err := httpClient.WaitForIOStreamState(context.Background(), IOStreamStateExpectation{ExpectedCount: ExpectedIOStreamCount(0), PresentStreamID: "present-id", AbsentStreamID: "stream-id"}) + require.NoError(t, err) + require.Equal(t, IOStreamState{Count: 0, Generation: 5}, waited) +} + +func TestClientIOStreamStateExpectationJSONPresence(t *testing.T) { + requestCount := 0 + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + var payload map[string]json.RawMessage + require.NoError(t, json.NewDecoder(request.Body).Decode(&payload)) + requestCount++ + if requestCount == 2 { + require.NotContains(t, payload, "expected_count") + } else { + value, exists := payload["expected_count"] + require.True(t, exists) + var count int + require.NoError(t, json.Unmarshal(value, &count)) + require.Zero(t, count) + } + writer.Header().Set("Content-Type", "application/json") + _, err := writer.Write([]byte(`{"success":true,"data":{"count":0,"generation":1}}`)) + require.NoError(t, err) + })) + t.Cleanup(server.Close) + httpClient := newTestClient(t, Config{BaseURL: server.URL}) + _, err := httpClient.WaitForIOStreamState(context.Background(), IOStreamStateExpectation{ExpectedCount: ExpectedIOStreamCount(0)}) + require.NoError(t, err) + _, err = REST[IOStreamStateExpectation, IOStreamState](context.Background(), httpClient, RESTRequest[IOStreamStateExpectation]{ + Method: http.MethodPost, + Path: "/agentcompat/io-stream-state", + Body: &IOStreamStateExpectation{AbsentStreamID: "stream-id"}, + }) + require.NoError(t, err) +} + +func TestClientIOStreamStateMapsSemanticFailure(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, err := writer.Write([]byte(`{"success":false,"error":"invalid expectation"}`)) + require.NoError(t, err) + })) + t.Cleanup(server.Close) + httpClient := newTestClient(t, Config{BaseURL: server.URL}) + state, err := httpClient.WaitForIOStreamState(context.Background(), IOStreamStateExpectation{AbsentStreamID: "private-stream-id"}) + require.ErrorIs(t, err, ErrSemanticFailure) + require.Zero(t, state) +} + +func TestClientIOStreamStateMapsUnauthorizedGETAndPOST(t *testing.T) { + for _, method := range []string{http.MethodGet, http.MethodPost} { + t.Run(method, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.WriteHeader(http.StatusUnauthorized) + })) + t.Cleanup(server.Close) + httpClient := newTestClient(t, Config{BaseURL: server.URL}) + var state IOStreamState + var err error + if method == http.MethodGet { + state, err = httpClient.IOStreamState(context.Background()) + } else { + state, err = httpClient.WaitForIOStreamState(context.Background(), IOStreamStateExpectation{ExpectedCount: ExpectedIOStreamCount(0)}) + } + require.Error(t, err) + require.True(t, errors.Is(err, ErrUnauthorized)) + require.Zero(t, state) + }) + } +} diff --git a/integration/agentcompat/internal/client/mcp.go b/integration/agentcompat/internal/client/mcp.go new file mode 100644 index 00000000..ed56e13e --- /dev/null +++ b/integration/agentcompat/internal/client/mcp.go @@ -0,0 +1,156 @@ +package client + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "net/http" +) + +type MCPContent struct { + Type string `json:"type"` + Text string `json:"text"` +} + +type ToolCall[Arguments any] struct { + Name string + Arguments Arguments +} + +type ToolCallResult[Result any] struct { + Content []MCPContent `json:"content"` + StructuredContent Result `json:"structuredContent"` + IsError bool `json:"isError"` +} + +type toolCallWireResult struct { + Content []MCPContent `json:"content"` + StructuredContent json.RawMessage `json:"structuredContent"` + IsError bool `json:"isError"` +} + +type InitializeResult struct { + ProtocolVersion string `json:"protocolVersion"` + ServerInfo struct { + Name string `json:"name"` + Version string `json:"version"` + } `json:"serverInfo"` +} + +type Tool struct { + Name string `json:"name"` + Description string `json:"description"` +} + +type ToolsListResult struct { + Tools []Tool `json:"tools"` +} + +type jsonRPCRequest[Params any] struct { + JSONRPC string `json:"jsonrpc"` + ID uint64 `json:"id"` + Method string `json:"method"` + Params Params `json:"params"` +} + +type jsonRPCResponse struct { + JSONRPC string `json:"jsonrpc"` + ID uint64 `json:"id"` + Result *json.RawMessage `json:"result"` + Error *RPCError `json:"error"` +} + +type toolCallParams[Arguments any] struct { + Name string `json:"name"` + Arguments Arguments `json:"arguments"` +} + +func (client *Client) Initialize(ctx context.Context) (InitializeResult, error) { + return mcpCall[struct{}, InitializeResult](ctx, client, "initialize", struct{}{}) +} + +func (client *Client) ListTools(ctx context.Context) (ToolsListResult, error) { + return mcpCall[struct{}, ToolsListResult](ctx, client, "tools/list", struct{}{}) +} + +func CallTool[Arguments, Result any](ctx context.Context, client *Client, call ToolCall[Arguments]) (ToolCallResult[Result], error) { + wireResult, err := mcpCall[toolCallParams[Arguments], toolCallWireResult](ctx, client, "tools/call", toolCallParams[Arguments]{ + Name: call.Name, + Arguments: call.Arguments, + }) + if err != nil { + return ToolCallResult[Result]{}, err + } + if wireResult.IsError { + message := "tool returned an error" + if len(wireResult.Content) > 0 && wireResult.Content[0].Text != "" { + message = wireResult.Content[0].Text + } + return ToolCallResult[Result]{}, &ToolFailure{Message: message, StructuredContent: json.RawMessage(Redact(string(wireResult.StructuredContent)))} + } + var result Result + if len(wireResult.StructuredContent) > 0 && string(wireResult.StructuredContent) != "null" { + if err := json.Unmarshal(wireResult.StructuredContent, &result); err != nil { + return ToolCallResult[Result]{}, fmt.Errorf("decode MCP tool structured content: %w", err) + } + } + return ToolCallResult[Result]{Content: wireResult.Content, StructuredContent: result, IsError: wireResult.IsError}, nil +} + +func mcpCall[Params, Result any](ctx context.Context, client *Client, method string, params Params) (Result, error) { + var zero Result + requestID := client.nextRequestID.Add(1) + requestBody, err := json.Marshal(jsonRPCRequest[Params]{JSONRPC: "2.0", ID: requestID, Method: method, Params: params}) + if err != nil { + return zero, fmt.Errorf("encode MCP request: %w", err) + } + requestURL, err := client.resolvePath("/mcp") + if err != nil { + return zero, err + } + requestContext, cancel := client.requestContext(ctx) + defer cancel() + httpRequest, err := http.NewRequestWithContext(requestContext, http.MethodPost, requestURL.String(), bytes.NewReader(requestBody)) + if err != nil { + return zero, fmt.Errorf("create MCP request: %w", err) + } + httpRequest.Header.Set("Content-Type", "application/json") + client.applyAuthenticatedHeaders(httpRequest, false) + + status, responseBody, err := client.execute(httpRequest, client.maxResponseBytes) + if err != nil { + return zero, err + } + var envelope jsonRPCResponse + if status < 200 || status >= 300 { + message := "" + if json.Unmarshal(responseBody, &envelope) == nil && envelope.Error != nil { + message = Redact(envelope.Error.Message) + } else { + message = Redact(string(responseBody)) + } + return zero, &HTTPError{StatusCode: status, Message: message} + } + if err := json.Unmarshal(responseBody, &envelope); err != nil { + return zero, fmt.Errorf("decode MCP response: %w", err) + } + if envelope.JSONRPC != "2.0" || envelope.ID != requestID { + return zero, fmt.Errorf("%w: invalid response envelope", ErrJSONRPC) + } + if (envelope.Result == nil) == (envelope.Error == nil) { + return zero, fmt.Errorf("%w: response must contain exactly one of result or error", ErrJSONRPC) + } + if envelope.Error != nil { + envelope.Error.Message = Redact(envelope.Error.Message) + return zero, envelope.Error + } + if string(*envelope.Result) == "null" { + return zero, fmt.Errorf("%w: null result", ErrJSONRPC) + } + var result Result + if err := json.Unmarshal(*envelope.Result, &result); err != nil { + return zero, fmt.Errorf("decode MCP result: %w", err) + } + return result, nil +} diff --git a/integration/agentcompat/internal/client/mcp_filesystem.go b/integration/agentcompat/internal/client/mcp_filesystem.go new file mode 100644 index 00000000..9e1b0466 --- /dev/null +++ b/integration/agentcompat/internal/client/mcp_filesystem.go @@ -0,0 +1,104 @@ +package client + +import "encoding/json" + +type WhoAmIResult struct { + UserID uint64 `json:"user_id"` + IsAdmin bool `json:"is_admin"` + TokenID uint64 `json:"token_id"` + TokenName string `json:"token_name"` + Scopes []string `json:"scopes"` + ServerIDs []uint64 `json:"server_ids"` +} + +type ServerListArguments struct { + OnlineOnly bool `json:"online_only"` +} + +type ServerListItem struct { + ID uint64 `json:"id"` + Name string `json:"name"` + UUID string `json:"uuid"` + Online bool `json:"online"` +} + +type ServerListResult struct { + Servers []ServerListItem `json:"servers"` + Count int `json:"count"` +} + +type ServerGetArguments struct { + ServerID uint64 `json:"server_id"` +} + +type ServerGetResult struct { + ID uint64 `json:"id"` + Name string `json:"name"` + UUID string `json:"uuid"` + Host json.RawMessage `json:"host"` + State json.RawMessage `json:"state"` +} + +type FsListArguments struct { + ServerID uint64 `json:"server_id"` + Path string `json:"path"` + ShowHidden bool `json:"show_hidden"` +} + +type FsEntry struct { + Name string `json:"name"` + Type string `json:"type"` + Size int64 `json:"size"` + Mode string `json:"mode"` + MTime int64 `json:"mtime"` + IsSymlink bool `json:"is_symlink"` + LinkTarget string `json:"link_target"` +} + +type FsListResult struct { + Entries []FsEntry `json:"entries"` + Truncated bool `json:"truncated"` + Total int `json:"total"` +} + +type FsReadArguments struct { + ServerID uint64 `json:"server_id"` + Path string `json:"path"` + Offset int64 `json:"offset"` + Length int64 `json:"length"` + Encoding string `json:"encoding"` +} + +type FsReadResult struct { + Content string `json:"content"` + Encoding string `json:"encoding"` + Size int64 `json:"size"` + SHA256 string `json:"sha256"` + Truncated bool `json:"truncated"` +} + +type FsWriteArguments struct { + ServerID uint64 `json:"server_id"` + Path string `json:"path"` + Content string `json:"content"` + Encoding string `json:"encoding"` + Mode string `json:"mode"` + IfMatchSHA256 string `json:"if_match_sha256"` + CreateDirs bool `json:"create_dirs"` +} + +type FsWriteResult struct { + Size int64 `json:"size"` + SHA256 string `json:"sha256"` + Error string `json:"error"` +} + +type FsDeleteArguments struct { + ServerID uint64 `json:"server_id"` + Path string `json:"path"` + Recursive bool `json:"recursive"` +} + +type FsDeleteResult struct { + DeletedCount int `json:"deleted_count"` +} diff --git a/integration/agentcompat/internal/client/mcp_test.go b/integration/agentcompat/internal/client/mcp_test.go new file mode 100644 index 00000000..f785ee40 --- /dev/null +++ b/integration/agentcompat/internal/client/mcp_test.go @@ -0,0 +1,290 @@ +package client + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +type testJSONRPCRequest struct { + JSONRPC string `json:"jsonrpc"` + ID uint64 `json:"id"` + Method string `json:"method"` + Params json.RawMessage `json:"params"` +} + +type testToolCallParams struct { + Name string `json:"name"` +} + +type fileReadArguments struct { + ServerID uint64 `json:"server_id"` + Path string `json:"path"` +} + +type fileReadResult struct { + Size int64 `json:"size"` + SHA256 string `json:"sha256"` +} + +func TestClient_MCPStructuredContent(t *testing.T) { + requestIDs := make(chan uint64, 2) + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + require.Equal(t, "/mcp", request.URL.Path) + require.Equal(t, "Bearer mcp-token", request.Header.Get("Authorization")) + require.NotEmpty(t, request.Header.Get("Origin")) + + var rpcRequest testJSONRPCRequest + require.NoError(t, json.NewDecoder(request.Body).Decode(&rpcRequest)) + require.Equal(t, "2.0", rpcRequest.JSONRPC) + require.Equal(t, "tools/call", rpcRequest.Method) + var params testToolCallParams + require.NoError(t, json.Unmarshal(rpcRequest.Params, ¶ms)) + require.Equal(t, "fs.read", params.Name) + requestIDs <- rpcRequest.ID + + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"jsonrpc":"2.0","id":` + jsonNumber(rpcRequest.ID) + `,"result":{"content":[{"type":"text","text":"ok"}],"structuredContent":{"size":7,"sha256":"abc123"}}}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{ + BaseURL: server.URL, + BearerToken: "mcp-token", + Origin: server.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + }) + call := ToolCall[fileReadArguments]{Name: "fs.read", Arguments: fileReadArguments{ServerID: 7, Path: "/tmp/report"}} + first, err := CallTool[fileReadArguments, fileReadResult](context.Background(), client, call) + require.NoError(t, err) + require.Equal(t, int64(7), first.StructuredContent.Size) + require.Equal(t, "abc123", first.StructuredContent.SHA256) + + _, err = CallTool[fileReadArguments, fileReadResult](context.Background(), client, call) + require.NoError(t, err) + require.Equal(t, uint64(1), <-requestIDs) + require.Equal(t, uint64(2), <-requestIDs) +} + +func TestClient_MCPUnauthorized(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.Header().Set("Content-Type", "application/json") + writer.WriteHeader(http.StatusUnauthorized) + _, _ = writer.Write([]byte(`{"jsonrpc":"2.0","id":1,"error":{"code":-32001,"message":"unauthorized"}}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{ + BaseURL: server.URL, + BearerToken: "invalid-token", + Origin: server.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + }) + _, err := CallTool[fileReadArguments, fileReadResult](context.Background(), client, ToolCall[fileReadArguments]{ + Name: "fs.read", + Arguments: fileReadArguments{ServerID: 7, Path: "/tmp/report"}, + }) + require.Error(t, err) + require.True(t, errors.Is(err, ErrUnauthorized)) +} + +func TestClient_MCPClassifiesNonJSONUnauthorized(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.WriteHeader(http.StatusUnauthorized) + _, _ = writer.Write([]byte("unauthorized")) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + _, err := client.Initialize(context.Background()) + require.ErrorIs(t, err, ErrUnauthorized) +} + +func TestClient_MCPToolSemanticFailure(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + var rpcRequest testJSONRPCRequest + require.NoError(t, json.NewDecoder(request.Body).Decode(&rpcRequest)) + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"jsonrpc":"2.0","id":` + jsonNumber(rpcRequest.ID) + `,"result":{"content":[{"type":"text","text":"agent offline"}],"isError":true}}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + _, err := CallTool[fileReadArguments, fileReadResult](context.Background(), client, ToolCall[fileReadArguments]{ + Name: "fs.read", + Arguments: fileReadArguments{ServerID: 7, Path: "/tmp/report"}, + }) + require.ErrorIs(t, err, ErrToolFailure) +} + +func TestClient_MCPToolFailurePreservesStructuredContent(t *testing.T) { + // Given + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + var rpcRequest testJSONRPCRequest + require.NoError(t, json.NewDecoder(request.Body).Decode(&rpcRequest)) + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"jsonrpc":"2.0","id":` + jsonNumber(rpcRequest.ID) + `,"result":{"content":[{"type":"text","text":"command not found"}],"structuredContent":{"exit_code":127,"error":"command or working directory not found"},"isError":true}}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + _, err := CallTool[fileReadArguments, fileReadResult](context.Background(), client, ToolCall[fileReadArguments]{Name: "server.exec", Arguments: fileReadArguments{ServerID: 7}}) + + // Then + var toolFailure *ToolFailure + require.ErrorAs(t, err, &toolFailure) + require.ErrorIs(t, err, ErrToolFailure) + require.Equal(t, "command not found", toolFailure.Message) + require.JSONEq(t, `{"exit_code":127,"error":"command or working directory not found"}`, string(toolFailure.StructuredContent)) +} + +func TestClient_MCPArbitraryTransportErrorIsNotToolFailure(t *testing.T) { + // Given + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + _, _ = io.Copy(io.Discard, request.Body) + <-request.Context().Done() + })) + t.Cleanup(server.Close) + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: 20 * time.Millisecond, MaxResponseBytes: 1024}) + + // When + _, err := CallTool[fileReadArguments, fileReadResult](context.Background(), client, ToolCall[fileReadArguments]{Name: "server.exec", Arguments: fileReadArguments{ServerID: 7}}) + + // Then + var toolFailure *ToolFailure + require.Error(t, err) + require.NotErrorAs(t, err, &toolFailure) + require.NotErrorIs(t, err, ErrToolFailure) +} + +func TestClient_MCPMalformedStructuredToolFailureIsTyped(t *testing.T) { + // Given + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + var rpcRequest testJSONRPCRequest + require.NoError(t, json.NewDecoder(request.Body).Decode(&rpcRequest)) + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"jsonrpc":"2.0","id":` + jsonNumber(rpcRequest.ID) + `,"result":{"content":[{"type":"text","text":"invalid command"}],"structuredContent":"not-an-object","isError":true}}`)) + })) + t.Cleanup(server.Close) + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + + // When + _, err := CallTool[fileReadArguments, fileReadResult](context.Background(), client, ToolCall[fileReadArguments]{Name: "server.exec", Arguments: fileReadArguments{ServerID: 7}}) + + // Then + var toolFailure *ToolFailure + require.ErrorAs(t, err, &toolFailure) + require.ErrorIs(t, err, ErrToolFailure) + require.JSONEq(t, `"not-an-object"`, string(toolFailure.StructuredContent)) +} + +func TestClient_MCPRejectsOversize(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"jsonrpc":"2.0","id":1,"result":{"padding":"` + string(make([]byte, 256)) + `"}}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 64}) + _, err := client.Initialize(context.Background()) + require.ErrorIs(t, err, ErrResponseTooLarge) +} + +func TestClient_MCPDeadline(t *testing.T) { + // The handler must enter before the client deadline: scheduling the handler + // against a tiny timeout can make this test miss a real MCP request. + requestBodyDrained := make(chan struct{}) + requestCancelled := make(chan struct{}) + releaseHandler := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(_ http.ResponseWriter, request *http.Request) { + if _, err := io.Copy(io.Discard, request.Body); err != nil { + return + } + close(requestBodyDrained) + select { + case <-request.Context().Done(): + close(requestCancelled) + case <-releaseHandler: + } + })) + t.Cleanup(func() { + close(releaseHandler) + server.Close() + }) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + result := make(chan error, 1) + go func() { + _, err := client.Initialize(context.Background()) + result <- err + }() + waitContext, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + select { + case <-requestBodyDrained: + case <-waitContext.Done(): + t.Fatal("server did not receive MCP request") + } + var err error + select { + case err = <-result: + case <-waitContext.Done(): + t.Fatal("MCP request did not reach its deadline") + } + require.ErrorIs(t, err, context.DeadlineExceeded) + select { + case <-requestCancelled: + case <-waitContext.Done(): + t.Fatal("server request context was not cancelled") + } +} + +func TestClient_MCPRejectsResultAndErrorTogether(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + var rpcRequest testJSONRPCRequest + require.NoError(t, json.NewDecoder(request.Body).Decode(&rpcRequest)) + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"jsonrpc":"2.0","id":` + jsonNumber(rpcRequest.ID) + `,"result":{},"error":{"code":-32603,"message":"invalid"}}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + _, err := client.Initialize(context.Background()) + require.ErrorIs(t, err, ErrJSONRPC) +} + +func TestClient_MCPRejectsRedirectWithoutForwardingAuthorization(t *testing.T) { + var redirected atomic.Int32 + foreignServer := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + redirected.Add(1) + })) + t.Cleanup(foreignServer.Close) + + dashboardServer := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + http.Redirect(writer, &http.Request{}, foreignServer.URL+"/capture", http.StatusTemporaryRedirect) + })) + t.Cleanup(dashboardServer.Close) + + client := newTestClient(t, Config{BaseURL: dashboardServer.URL, BearerToken: "mcp-token", RequestTimeout: time.Second, MaxResponseBytes: 1024}) + _, err := client.Initialize(context.Background()) + require.ErrorIs(t, err, ErrRedirect) + require.Zero(t, redirected.Load()) +} + +func jsonNumber(value uint64) string { + encoded, err := json.Marshal(value) + if err != nil { + panic(err) + } + return string(encoded) +} diff --git a/integration/agentcompat/internal/client/sqlite_hold.go b/integration/agentcompat/internal/client/sqlite_hold.go new file mode 100644 index 00000000..c5c0e71b --- /dev/null +++ b/integration/agentcompat/internal/client/sqlite_hold.go @@ -0,0 +1,90 @@ +package client + +import ( + "context" + "encoding/base64" + "errors" + "net/http" +) + +var ErrInvalidSQLiteHoldReceipt = errors.New("client: invalid sqlite hold receipt") + +type SQLiteHoldState string + +const ( + SQLiteHoldStateArmed SQLiteHoldState = "armed" + SQLiteHoldStateSelected SQLiteHoldState = "selected" + SQLiteHoldStateFinalizing SQLiteHoldState = "finalizing" + SQLiteHoldStateReleased SQLiteHoldState = "released" + SQLiteHoldStateAborted SQLiteHoldState = "aborted" +) + +type SQLiteHoldReceipt struct { + ID string `json:"id"` + State SQLiteHoldState `json:"state,omitempty"` +} + +func (client *Client) ArmSQLiteHold(ctx context.Context) (SQLiteHoldReceipt, error) { + receipt, err := DoREST[struct{}, SQLiteHoldReceipt](ctx, client, RESTRequest[struct{}]{Method: http.MethodPost, Path: "/agentcompat/sqlite-hold/arm", Body: &struct{}{}}) + return validateSQLiteHoldReceipt(receipt, SQLiteHoldStateArmed, err) +} + +func (client *Client) WaitForSQLiteHold(ctx context.Context, receipt SQLiteHoldReceipt, target SQLiteHoldState) (SQLiteHoldReceipt, error) { + if err := validateSQLiteHoldRequest(receipt, target); err != nil { + return SQLiteHoldReceipt{}, err + } + request := SQLiteHoldReceipt{ID: receipt.ID, State: target} + result, err := DoREST[SQLiteHoldReceipt, SQLiteHoldReceipt](ctx, client, RESTRequest[SQLiteHoldReceipt]{Method: http.MethodPost, Path: "/agentcompat/sqlite-hold/wait", Body: &request}) + return validateSQLiteHoldReceipt(result, target, err) +} + +func (client *Client) SnapshotSQLiteHold(ctx context.Context, receipt SQLiteHoldReceipt) (SQLiteHoldReceipt, error) { + return client.sqliteHoldAction(ctx, receipt, "snapshot", "") +} + +func (client *Client) ReleaseSQLiteHold(ctx context.Context, receipt SQLiteHoldReceipt) (SQLiteHoldReceipt, error) { + return client.sqliteHoldAction(ctx, receipt, "release", SQLiteHoldStateReleased) +} + +func (client *Client) AbortSQLiteHold(ctx context.Context, receipt SQLiteHoldReceipt) (SQLiteHoldReceipt, error) { + return client.sqliteHoldAction(ctx, receipt, "abort", SQLiteHoldStateAborted) +} + +func (client *Client) sqliteHoldAction(ctx context.Context, receipt SQLiteHoldReceipt, action string, expected SQLiteHoldState) (SQLiteHoldReceipt, error) { + if err := validateSQLiteHoldRequest(receipt, ""); err != nil { + return SQLiteHoldReceipt{}, err + } + request := SQLiteHoldReceipt{ID: receipt.ID} + result, err := DoREST[SQLiteHoldReceipt, SQLiteHoldReceipt](ctx, client, RESTRequest[SQLiteHoldReceipt]{Method: http.MethodPost, Path: "/agentcompat/sqlite-hold/" + action, Body: &request}) + return validateSQLiteHoldReceipt(result, expected, err) +} + +func validateSQLiteHoldRequest(receipt SQLiteHoldReceipt, target SQLiteHoldState) error { + decoded, err := base64.RawURLEncoding.DecodeString(receipt.ID) + if err != nil || len(receipt.ID) != 43 || len(decoded) != 32 { + return ErrInvalidSQLiteHoldReceipt + } + if target != "" && target != SQLiteHoldStateSelected && target != SQLiteHoldStateFinalizing { + return ErrInvalidSQLiteHoldReceipt + } + return nil +} + +func validateSQLiteHoldReceipt(receipt SQLiteHoldReceipt, expected SQLiteHoldState, err error) (SQLiteHoldReceipt, error) { + if err != nil { + return SQLiteHoldReceipt{}, err + } + if err := validateSQLiteHoldRequest(receipt, ""); err != nil || !validSQLiteHoldState(receipt.State) || expected != "" && receipt.State != expected { + return SQLiteHoldReceipt{}, ErrInvalidSQLiteHoldReceipt + } + return receipt, nil +} + +func validSQLiteHoldState(state SQLiteHoldState) bool { + switch state { + case SQLiteHoldStateArmed, SQLiteHoldStateSelected, SQLiteHoldStateFinalizing, SQLiteHoldStateReleased, SQLiteHoldStateAborted: + return true + default: + return false + } +} diff --git a/integration/agentcompat/internal/client/sqlite_hold_test.go b/integration/agentcompat/internal/client/sqlite_hold_test.go new file mode 100644 index 00000000..f5d7705d --- /dev/null +++ b/integration/agentcompat/internal/client/sqlite_hold_test.go @@ -0,0 +1,78 @@ +package client + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestClientSQLiteHoldHelpersUseTypedRESTContracts(t *testing.T) { + // Given + receiptID := "ERERERERERERERERERERERERERERERERERERERERERE" + requests := make([]string, 0, 6) + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + requests = append(requests, request.URL.Path) + var payload map[string]json.RawMessage + require.NoError(t, json.NewDecoder(request.Body).Decode(&payload)) + state := SQLiteHoldStateArmed + switch request.URL.Path { + case "/agentcompat/sqlite-hold/arm": + require.Empty(t, payload) + case "/agentcompat/sqlite-hold/wait": + require.JSONEq(t, `"`+receiptID+`"`, string(payload["id"])) + require.NoError(t, json.Unmarshal(payload["state"], &state)) + require.Contains(t, []SQLiteHoldState{SQLiteHoldStateSelected, SQLiteHoldStateFinalizing}, state) + case "/agentcompat/sqlite-hold/snapshot": + require.Len(t, payload, 1) + state = SQLiteHoldStateSelected + case "/agentcompat/sqlite-hold/release": + require.Len(t, payload, 1) + state = SQLiteHoldStateReleased + case "/agentcompat/sqlite-hold/abort": + require.Len(t, payload, 1) + state = SQLiteHoldStateAborted + default: + writer.WriteHeader(http.StatusNotFound) + return + } + writer.Header().Set("Content-Type", "application/json") + _, err := writer.Write([]byte(`{"success":true,"data":{"id":"` + receiptID + `","state":"` + string(state) + `"}}`)) + require.NoError(t, err) + })) + t.Cleanup(server.Close) + httpClient := newTestClient(t, Config{BaseURL: server.URL}) + + // When + armed, err := httpClient.ArmSQLiteHold(context.Background()) + require.NoError(t, err) + selected, err := httpClient.WaitForSQLiteHold(context.Background(), armed, SQLiteHoldStateSelected) + require.NoError(t, err) + finalizing, err := httpClient.WaitForSQLiteHold(context.Background(), selected, SQLiteHoldStateFinalizing) + require.NoError(t, err) + snapshot, err := httpClient.SnapshotSQLiteHold(context.Background(), finalizing) + require.NoError(t, err) + released, err := httpClient.ReleaseSQLiteHold(context.Background(), snapshot) + require.NoError(t, err) + aborted, err := httpClient.AbortSQLiteHold(context.Background(), released) + + // Then + require.NoError(t, err) + require.Equal(t, SQLiteHoldStateArmed, armed.State) + require.Equal(t, SQLiteHoldStateSelected, selected.State) + require.Equal(t, SQLiteHoldStateFinalizing, finalizing.State) + require.Equal(t, SQLiteHoldStateSelected, snapshot.State) + require.Equal(t, SQLiteHoldStateReleased, released.State) + require.Equal(t, SQLiteHoldStateAborted, aborted.State) + require.Equal(t, []string{ + "/agentcompat/sqlite-hold/arm", + "/agentcompat/sqlite-hold/wait", + "/agentcompat/sqlite-hold/wait", + "/agentcompat/sqlite-hold/snapshot", + "/agentcompat/sqlite-hold/release", + "/agentcompat/sqlite-hold/abort", + }, requests) +} diff --git a/integration/agentcompat/internal/client/transfer.go b/integration/agentcompat/internal/client/transfer.go new file mode 100644 index 00000000..7ebbedce --- /dev/null +++ b/integration/agentcompat/internal/client/transfer.go @@ -0,0 +1,184 @@ +package client + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "time" +) + +type TransferURL struct { + URL string `json:"url"` + Method string `json:"method"` + ExpiresAt time.Time `json:"expires_at"` +} + +type DownloadURLRequest struct { + ServerID uint64 `json:"server_id"` + Path string `json:"path"` + TTLSeconds int `json:"ttl_seconds,omitempty"` +} + +type UploadURLRequest struct { + ServerID uint64 `json:"server_id"` + Path string `json:"path"` + TTLSeconds int `json:"ttl_seconds,omitempty"` + Mode string `json:"mode,omitempty"` + CreateDirs bool `json:"create_dirs,omitempty"` + IfMatchSHA256 string `json:"if_match_sha256,omitempty"` +} + +type UploadTransfer struct { + Body io.Reader + ContentLength int64 + SHA256 string +} + +type UploadResult struct { + Size int64 `json:"size"` + SHA256 string `json:"sha256"` +} + +type uploadResultPayload struct { + Size *int64 `json:"size"` + SHA256 *string `json:"sha256"` +} + +func RequestDownloadURL(ctx context.Context, client *Client, request DownloadURLRequest) (TransferURL, error) { + result, err := CallTool[DownloadURLRequest, TransferURL](ctx, client, ToolCall[DownloadURLRequest]{Name: "fs.download_url", Arguments: request}) + if err != nil { + return TransferURL{}, err + } + return client.validateTransferURL(result.StructuredContent, http.MethodGet) +} + +func RequestUploadURL(ctx context.Context, client *Client, request UploadURLRequest) (TransferURL, error) { + result, err := CallTool[UploadURLRequest, TransferURL](ctx, client, ToolCall[UploadURLRequest]{Name: "fs.upload_url", Arguments: request}) + if err != nil { + return TransferURL{}, err + } + return client.validateTransferURL(result.StructuredContent, http.MethodPost) +} + +func (client *Client) DownloadTransfer(ctx context.Context, transfer TransferURL, destination io.Writer) (int64, error) { + validated, err := client.validateTransferURL(transfer, http.MethodGet) + if err != nil { + return 0, err + } + requestContext, cancel := client.transferContext(ctx) + defer cancel() + request, err := http.NewRequestWithContext(requestContext, http.MethodGet, validated.URL, nil) + if err != nil { + return 0, errorsNewRedacted("create transfer request", err) + } + response, err := client.transferHTTPClient().Do(request) + if err != nil { + if requestContext.Err() != nil { + return 0, fmt.Errorf("download transfer: %w", requestContext.Err()) + } + return 0, errorsNewRedacted("download transfer", err) + } + defer response.Body.Close() + if response.StatusCode < 200 || response.StatusCode >= 300 { + return 0, &HTTPError{StatusCode: response.StatusCode} + } + written, err := io.Copy(destination, io.LimitReader(response.Body, client.maxTransferBytes)) + if err != nil { + return written, fmt.Errorf("copy download transfer: %w", err) + } + var overflow [1]byte + read, err := response.Body.Read(overflow[:]) + if err != nil && err != io.EOF { + return written, fmt.Errorf("probe download transfer size: %w", err) + } + if read > 0 { + return written, ErrTransferTooLarge + } + return written, nil +} + +func (client *Client) UploadTransfer(ctx context.Context, transfer TransferURL, upload UploadTransfer) (UploadResult, error) { + validated, err := client.validateTransferURL(transfer, http.MethodPost) + if err != nil { + return UploadResult{}, err + } + if upload.Body == nil || upload.ContentLength <= 0 { + return UploadResult{}, fmt.Errorf("upload content length: %w", ErrInvalidConfig) + } + if upload.ContentLength > client.maxTransferBytes { + return UploadResult{}, ErrTransferTooLarge + } + transferURL, err := url.Parse(validated.URL) + if err != nil { + return UploadResult{}, errorsNewRedacted("parse transfer URL", err) + } + if upload.SHA256 != "" { + query := transferURL.Query() + query.Set("sha256", upload.SHA256) + transferURL.RawQuery = query.Encode() + } + requestContext, cancel := client.transferContext(ctx) + defer cancel() + request, err := http.NewRequestWithContext(requestContext, http.MethodPost, transferURL.String(), io.LimitReader(upload.Body, upload.ContentLength)) + if err != nil { + return UploadResult{}, errorsNewRedacted("create transfer request", err) + } + request.ContentLength = upload.ContentLength + response, err := client.transferHTTPClient().Do(request) + if err != nil { + if requestContext.Err() != nil { + return UploadResult{}, fmt.Errorf("upload transfer: %w", requestContext.Err()) + } + return UploadResult{}, errorsNewRedacted("upload transfer", err) + } + defer response.Body.Close() + body, err := readBounded(response.Body, client.maxResponseBytes) + if err != nil { + return UploadResult{}, err + } + if response.StatusCode < 200 || response.StatusCode >= 300 { + return UploadResult{}, &HTTPError{StatusCode: response.StatusCode, Message: Redact(string(body))} + } + var payload uploadResultPayload + if err := json.NewDecoder(bytes.NewReader(body)).Decode(&payload); err != nil { + return UploadResult{}, fmt.Errorf("decode upload result: %w", err) + } + if payload.Size == nil || payload.SHA256 == nil || *payload.Size <= 0 || *payload.SHA256 == "" { + return UploadResult{}, fmt.Errorf("decode upload result: %w", ErrSemanticFailure) + } + return UploadResult{Size: *payload.Size, SHA256: *payload.SHA256}, nil +} + +func (client *Client) transferHTTPClient() *http.Client { + clone := *client.httpClient + clone.Jar = nil + clone.CheckRedirect = rejectRedirect + return &clone +} + +func (client *Client) validateTransferURL(transfer TransferURL, expectedMethod string) (TransferURL, error) { + parsed, err := url.Parse(transfer.URL) + if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { + return TransferURL{}, fmt.Errorf("transfer URL: %w", ErrInvalidConfig) + } + if parsed.User != nil || !client.sameOrigin(parsed) { + return TransferURL{}, fmt.Errorf("transfer origin: %w", ErrInvalidConfig) + } + if transfer.ExpiresAt.IsZero() || !time.Now().Before(transfer.ExpiresAt) { + return TransferURL{}, ErrTransferExpired + } + if !strings.EqualFold(transfer.Method, expectedMethod) { + return TransferURL{}, fmt.Errorf("transfer method: %w", ErrInvalidConfig) + } + transfer.Method = expectedMethod + return transfer, nil +} + +func (client *Client) sameOrigin(candidate *url.URL) bool { + return strings.EqualFold(candidate.Scheme, client.baseURL.Scheme) && strings.EqualFold(candidate.Host, client.baseURL.Host) +} diff --git a/integration/agentcompat/internal/client/transfer_rejection.go b/integration/agentcompat/internal/client/transfer_rejection.go new file mode 100644 index 00000000..0ad21bfe --- /dev/null +++ b/integration/agentcompat/internal/client/transfer_rejection.go @@ -0,0 +1,46 @@ +package client + +import ( + "context" + "fmt" + "io" + "net/http" +) + +type OversizeUploadProbe struct { + Body io.Reader + ContentLength int64 +} + +func (client *Client) ProbeOversizeUpload(ctx context.Context, transfer TransferURL, probe OversizeUploadProbe) error { + validated, err := client.validateTransferURL(transfer, http.MethodPost) + if err != nil { + return err + } + if probe.Body == nil || probe.ContentLength != client.maxTransferBytes+1 { + return fmt.Errorf("oversize upload probe: %w", ErrInvalidConfig) + } + requestContext, cancel := client.transferContext(ctx) + defer cancel() + request, err := http.NewRequestWithContext(requestContext, http.MethodPost, validated.URL, io.LimitReader(probe.Body, probe.ContentLength)) + if err != nil { + return errorsNewRedacted("create oversize transfer probe", err) + } + request.ContentLength = probe.ContentLength + response, err := client.transferHTTPClient().Do(request) + if err != nil { + if requestContext.Err() != nil { + return fmt.Errorf("oversize transfer probe: %w", requestContext.Err()) + } + return errorsNewRedacted("oversize transfer probe", err) + } + defer response.Body.Close() + body, err := readBounded(response.Body, client.maxResponseBytes) + if err != nil { + return err + } + if response.StatusCode >= 200 && response.StatusCode < 300 { + return fmt.Errorf("oversize transfer probe unexpectedly succeeded: %w", ErrSemanticFailure) + } + return &HTTPError{StatusCode: response.StatusCode, Message: Redact(string(body))} +} diff --git a/integration/agentcompat/internal/client/transfer_test.go b/integration/agentcompat/internal/client/transfer_test.go new file mode 100644 index 00000000..5224d229 --- /dev/null +++ b/integration/agentcompat/internal/client/transfer_test.go @@ -0,0 +1,249 @@ +package client + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestClient_TransferURLConsumptionOmitsAuthorization(t *testing.T) { + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/mcp": + require.Equal(t, "Bearer mcp-token", request.Header.Get("Authorization")) + var rpcRequest testJSONRPCRequest + require.NoError(t, json.NewDecoder(request.Body).Decode(&rpcRequest)) + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"jsonrpc":"2.0","id":` + jsonNumber(rpcRequest.ID) + `,"result":{"structuredContent":{"url":"` + server.URL + `/mcp/download/one-time-token","method":"GET","expires_at":"2030-01-02T03:04:05Z"}}}`)) + case "/mcp/download/one-time-token": + require.Empty(t, request.Header.Get("Authorization")) + require.Empty(t, request.Header.Get("Origin")) + _, _ = writer.Write([]byte("transfer payload")) + default: + http.NotFound(writer, request) + } + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{ + BaseURL: server.URL, + BearerToken: "mcp-token", + Origin: server.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + MaxTransferBytes: 1024, + }) + transfer, err := RequestDownloadURL(context.Background(), client, DownloadURLRequest{ServerID: 7, Path: "/tmp/report", TTLSeconds: 30}) + require.NoError(t, err) + require.Equal(t, http.MethodGet, transfer.Method) + + var destination bytes.Buffer + written, err := client.DownloadTransfer(context.Background(), transfer, &destination) + require.NoError(t, err) + require.Equal(t, int64(len("transfer payload")), written) + require.Equal(t, "transfer payload", destination.String()) +} + +func TestClient_TransferURLRejectsCrossOrigin(t *testing.T) { + foreignServer := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + t.Fatal("cross-origin transfer request was dispatched") + })) + t.Cleanup(foreignServer.Close) + + client := newTestClient(t, Config{ + BaseURL: "http://dashboard.example", + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + MaxTransferBytes: 1024, + }) + + var destination bytes.Buffer + _, err := client.DownloadTransfer(context.Background(), TransferURL{URL: foreignServer.URL + "/mcp/download/token", Method: http.MethodGet}, &destination) + require.Error(t, err) + require.True(t, errors.Is(err, ErrInvalidConfig)) +} + +func TestClient_TransferURLRejectsCrossOriginRedirect(t *testing.T) { + var foreignRequests atomic.Int32 + foreignServer := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + foreignRequests.Add(1) + })) + t.Cleanup(foreignServer.Close) + + dashboardServer := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + http.Redirect(writer, &http.Request{}, foreignServer.URL+"/stolen-token", http.StatusTemporaryRedirect) + })) + t.Cleanup(dashboardServer.Close) + + client := newTestClient(t, Config{ + BaseURL: dashboardServer.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + MaxTransferBytes: 1024, + }) + + var destination bytes.Buffer + _, err := client.DownloadTransfer(context.Background(), TransferURL{ + URL: dashboardServer.URL + "/mcp/download/token", + Method: http.MethodGet, + ExpiresAt: time.Now().Add(time.Minute), + }, &destination) + require.ErrorIs(t, err, ErrRedirect) + require.Zero(t, foreignRequests.Load()) +} + +func TestClient_TransferURLRejectsSameOriginRedirect(t *testing.T) { + var redirected atomic.Int32 + var dashboardServer *httptest.Server + dashboardServer = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + if request.URL.Path == "/capture" { + redirected.Add(1) + return + } + http.Redirect(writer, request, dashboardServer.URL+"/capture", http.StatusTemporaryRedirect) + })) + t.Cleanup(dashboardServer.Close) + + client := newTestClient(t, Config{BaseURL: dashboardServer.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024, MaxTransferBytes: 1024}) + var destination bytes.Buffer + _, err := client.DownloadTransfer(context.Background(), TransferURL{ + URL: dashboardServer.URL + "/mcp/download/token", + Method: http.MethodGet, + ExpiresAt: time.Now().Add(time.Minute), + }, &destination) + require.ErrorIs(t, err, ErrRedirect) + require.Zero(t, redirected.Load()) +} + +func TestClient_TransferURLRejectsExpiredCapability(t *testing.T) { + client := newTestClient(t, Config{BaseURL: "http://dashboard.example", RequestTimeout: time.Second, MaxResponseBytes: 1024, MaxTransferBytes: 1024}) + var destination bytes.Buffer + _, err := client.DownloadTransfer(context.Background(), TransferURL{ + URL: "http://dashboard.example/mcp/download/token", + Method: http.MethodGet, + ExpiresAt: time.Now().Add(-time.Second), + }, &destination) + require.ErrorIs(t, err, ErrTransferExpired) +} + +func TestClient_TransferURLRejectsMissingExpiryBeforeDispatch(t *testing.T) { + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + requests.Add(1) + writer.WriteHeader(http.StatusOK) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024, MaxTransferBytes: 1024}) + var destination bytes.Buffer + _, err := client.DownloadTransfer(context.Background(), TransferURL{ + URL: server.URL + "/mcp/download/token", + Method: http.MethodGet, + }, &destination) + require.ErrorIs(t, err, ErrTransferExpired) + require.Zero(t, requests.Load()) +} + +func TestClient_RequestUploadURLRejectsMissingExpiryBeforeDispatch(t *testing.T) { + var requests atomic.Int32 + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + requests.Add(1) + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"jsonrpc":"2.0","id":1,"result":{"structuredContent":{"url":"` + server.URL + `/mcp/upload/token","method":"POST"}}}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024, MaxTransferBytes: 1024}) + _, err := RequestUploadURL(context.Background(), client, UploadURLRequest{ServerID: 7, Path: "/tmp/report"}) + require.ErrorIs(t, err, ErrTransferExpired) + require.Equal(t, int32(1), requests.Load()) +} + +func TestClient_TransferRejectsOversizeWithoutWritingPastLimit(t *testing.T) { + dashboardServer := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + _, _ = writer.Write([]byte(strings.Repeat("x", 65))) + })) + t.Cleanup(dashboardServer.Close) + + client := newTestClient(t, Config{ + BaseURL: dashboardServer.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + MaxTransferBytes: 64, + }) + + var destination bytes.Buffer + written, err := client.DownloadTransfer(context.Background(), TransferURL{ + URL: dashboardServer.URL + "/mcp/download/token", + Method: http.MethodGet, + ExpiresAt: time.Now().Add(time.Minute), + }, &destination) + require.ErrorIs(t, err, ErrTransferTooLarge) + require.Equal(t, int64(64), written) + require.Len(t, destination.Bytes(), 64) +} + +func TestClient_UploadRejectsChunkedBodyWithoutDispatch(t *testing.T) { + var requests atomic.Int32 + dashboardServer := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + requests.Add(1) + })) + t.Cleanup(dashboardServer.Close) + + client := newTestClient(t, Config{ + BaseURL: dashboardServer.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + MaxTransferBytes: 64, + }) + _, err := client.UploadTransfer(context.Background(), TransferURL{ + URL: dashboardServer.URL + "/mcp/upload/token", + Method: http.MethodPost, + ExpiresAt: time.Now().Add(time.Minute), + }, UploadTransfer{Body: strings.NewReader("hidden body"), ContentLength: 0}) + require.ErrorIs(t, err, ErrInvalidConfig) + require.Zero(t, requests.Load()) +} + +func TestClient_UploadRejectsMisleadingSuccessEnvelope(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"success":true,"data":{"size":11,"sha256":"abc"}}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024, MaxTransferBytes: 1024}) + _, err := client.UploadTransfer(context.Background(), TransferURL{ + URL: server.URL + "/mcp/upload/token", + Method: http.MethodPost, + ExpiresAt: time.Now().Add(time.Minute), + }, UploadTransfer{Body: strings.NewReader("payload"), ContentLength: int64(len("payload"))}) + require.ErrorIs(t, err, ErrSemanticFailure) +} + +func TestClient_UploadRejectsMissingRequiredResultFields(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"size":7}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024, MaxTransferBytes: 1024}) + _, err := client.UploadTransfer(context.Background(), TransferURL{ + URL: server.URL + "/mcp/upload/token", + Method: http.MethodPost, + ExpiresAt: time.Now().Add(time.Minute), + }, UploadTransfer{Body: strings.NewReader("payload"), ContentLength: int64(len("payload"))}) + require.ErrorIs(t, err, ErrSemanticFailure) +} diff --git a/integration/agentcompat/internal/client/websocket.go b/integration/agentcompat/internal/client/websocket.go new file mode 100644 index 00000000..58219370 --- /dev/null +++ b/integration/agentcompat/internal/client/websocket.go @@ -0,0 +1,220 @@ +package client + +import ( + "context" + "errors" + "fmt" + "net" + "net/http" + "sync" + "time" + + "github.com/gorilla/websocket" +) + +type FrameType string + +const ( + FrameText FrameType = "text" + FrameBinary FrameType = "binary" +) + +var ErrUnsupportedFrame = errors.New("client: unsupported WebSocket frame") + +type Frame struct { + Type FrameType + Payload []byte +} + +type WebSocketConnection struct { + connection *websocket.Conn + timeout time.Duration + readLock sync.Mutex + writeLock sync.Mutex + closeOnce sync.Once + // closeDone publishes the first physical close result to every caller. + closeDone chan struct{} + closeError error + afterReadMessageForTest func() +} + +func (client *Client) DialWebSocket(ctx context.Context, path string) (*WebSocketConnection, error) { + requestURL, err := client.resolvePath(path) + if err != nil { + return nil, err + } + switch requestURL.Scheme { + case "http": + requestURL.Scheme = "ws" + case "https": + requestURL.Scheme = "wss" + default: + return nil, fmt.Errorf("WebSocket scheme: %w", ErrInvalidConfig) + } + header := make(http.Header) + if client.bearerToken != "" { + header.Set("Authorization", "Bearer "+client.bearerToken) + } + if client.origin != "" { + header.Set("Origin", client.origin) + } + requestContext, cancel := client.requestContext(ctx) + defer cancel() + connection, response, err := client.webSocketDialer.DialContext(requestContext, requestURL.String(), header) + if err != nil { + if response != nil && response.Body != nil { + defer response.Body.Close() + body, readErr := readBounded(response.Body, client.maxResponseBytes) + if readErr != nil { + return nil, fmt.Errorf("read WebSocket handshake failure: %w", readErr) + } + return nil, &WebSocketHandshakeError{StatusCode: response.StatusCode, Message: string(body)} + } + if requestContext.Err() != nil { + return nil, fmt.Errorf("dial WebSocket: %w", requestContext.Err()) + } + return nil, errorsNewRedacted("dial WebSocket", err) + } + if response != nil && response.Body != nil { + response.Body.Close() + } + connection.SetReadLimit(client.maxResponseBytes) + return &WebSocketConnection{connection: connection, timeout: client.requestTimeout, closeDone: make(chan struct{})}, nil +} + +func (connection *WebSocketConnection) ReadFrame(ctx context.Context) (Frame, error) { + readContext, cancel := context.WithTimeout(ctx, connection.timeout) + defer cancel() + return connection.readFrame(readContext) +} + +// ReadFrameUntil reads one frame using only the caller's cancellation and deadline. +func (connection *WebSocketConnection) ReadFrameUntil(ctx context.Context) (Frame, error) { + return connection.readFrame(ctx) +} + +func (connection *WebSocketConnection) readFrame(ctx context.Context) (Frame, error) { + connection.readLock.Lock() + defer connection.readLock.Unlock() + var cancellationState struct { + sync.Mutex + completed bool + } + stopCancellation := context.AfterFunc(ctx, func() { + cancellationState.Lock() + defer cancellationState.Unlock() + if !cancellationState.completed { + _ = connection.Close() + } + }) + defer func() { + cancellationState.Lock() + cancellationState.completed = true + cancellationState.Unlock() + stopCancellation() + }() + cancellationOccurred := func() bool { + cancellationState.Lock() + defer cancellationState.Unlock() + if ctx.Err() == nil { + return false + } + cancellationState.completed = true + return true + } + if deadline, ok := ctx.Deadline(); ok { + if err := connection.connection.SetReadDeadline(deadline); err != nil { + // A cancellation callback may close the socket while SetReadDeadline runs. + if cancellationOccurred() { + return Frame{}, fmt.Errorf("read WebSocket frame: %w", ctx.Err()) + } + return Frame{}, fmt.Errorf("set WebSocket read deadline: %w", err) + } + } else if err := connection.connection.SetReadDeadline(time.Time{}); err != nil { + // Gorilla retains prior deadlines until explicitly cleared. + if cancellationOccurred() { + return Frame{}, fmt.Errorf("read WebSocket frame: %w", ctx.Err()) + } + return Frame{}, fmt.Errorf("set WebSocket read deadline: %w", err) + } + messageType, payload, err := connection.connection.ReadMessage() + if connection.afterReadMessageForTest != nil { + connection.afterReadMessageForTest() + } + cancellationState.Lock() + cancellationWon := ctx.Err() != nil || !stopCancellation() + if cancellationWon { + _ = connection.Close() + } + cancellationState.completed = true + cancellationState.Unlock() + if cancellationWon { + return Frame{}, fmt.Errorf("read WebSocket frame: %w", ctx.Err()) + } + if err != nil { + if errors.Is(err, websocket.ErrReadLimit) { + return Frame{}, ErrResponseTooLarge + } + var closeError *websocket.CloseError + if errors.As(err, &closeError) { + return Frame{}, &WebSocketCloseError{Code: closeError.Code, Text: closeError.Text} + } + var networkError net.Error + if errors.As(err, &networkError) && networkError.Timeout() { + return Frame{}, fmt.Errorf("read WebSocket frame: %w", context.DeadlineExceeded) + } + return Frame{}, errorsNewRedacted("read WebSocket frame", err) + } + switch messageType { + case websocket.TextMessage: + return Frame{Type: FrameText, Payload: payload}, nil + case websocket.BinaryMessage: + return Frame{Type: FrameBinary, Payload: payload}, nil + default: + return Frame{}, ErrUnsupportedFrame + } +} + +func (connection *WebSocketConnection) WriteFrame(ctx context.Context, frame Frame) error { + connection.writeLock.Lock() + defer connection.writeLock.Unlock() + writeContext, cancel := context.WithTimeout(ctx, connection.timeout) + defer cancel() + stopCancellation := context.AfterFunc(writeContext, func() { _ = connection.Close() }) + defer stopCancellation() + deadline, _ := writeContext.Deadline() + if err := connection.connection.SetWriteDeadline(deadline); err != nil { + return fmt.Errorf("set WebSocket write deadline: %w", err) + } + var messageType int + switch frame.Type { + case FrameText: + messageType = websocket.TextMessage + case FrameBinary: + messageType = websocket.BinaryMessage + default: + return ErrUnsupportedFrame + } + if err := connection.connection.WriteMessage(messageType, frame.Payload); err != nil { + if writeContext.Err() != nil { + return fmt.Errorf("write WebSocket frame: %w", writeContext.Err()) + } + var networkError net.Error + if errors.As(err, &networkError) && networkError.Timeout() { + return fmt.Errorf("write WebSocket frame: %w", context.DeadlineExceeded) + } + return errorsNewRedacted("write WebSocket frame", err) + } + return nil +} + +func (connection *WebSocketConnection) Close() error { + connection.closeOnce.Do(func() { + if err := connection.connection.Close(); err != nil && !errors.Is(err, net.ErrClosed) { + connection.closeError = err + } + close(connection.closeDone) + }) + <-connection.closeDone + return connection.closeError +} diff --git a/integration/agentcompat/internal/client/websocket_close_test.go b/integration/agentcompat/internal/client/websocket_close_test.go new file mode 100644 index 00000000..c133e6c9 --- /dev/null +++ b/integration/agentcompat/internal/client/websocket_close_test.go @@ -0,0 +1,169 @@ +package client + +import ( + "context" + "errors" + "net" + "net/http" + "sync" + "testing" + "time" + + "github.com/gorilla/websocket" + "github.com/stretchr/testify/require" +) + +func TestClient_WebSocketClose_RetainsConcurrentPhysicalCloseResult(t *testing.T) { + // Given + closeErr := errors.New("physical close failed") + recordedConnection := newRetainedCloseConn(closeErr) + server := newWebSocketTestServer(t, websocket.Upgrader{}, func(_ *websocket.Conn, request *http.Request) { + <-request.Context().Done() + }) + connection := dialRetainedCloseWebSocket(t, server.URL, recordedConnection) + + // When + results := make(chan error, 8) + for range cap(results) { + go func() { results <- connection.Close() }() + } + + // Then + for range cap(results) { + require.Equal(t, closeErr, <-results) + } + require.Equal(t, 1, recordedConnection.closeCount()) + require.Equal(t, closeErr, connection.Close()) +} + +func TestClient_WebSocketClose_RetainsReadCancellationCloseResult(t *testing.T) { + // Given + closeErr := errors.New("read cancellation close failed") + recordedConnection := newRetainedCloseConn(closeErr) + serverReady := make(chan struct{}) + server := newWebSocketTestServer(t, websocket.Upgrader{}, func(_ *websocket.Conn, request *http.Request) { + close(serverReady) + <-request.Context().Done() + }) + connection := dialRetainedCloseWebSocket(t, server.URL, recordedConnection) + readContext, cancel := context.WithCancel(context.Background()) + readResult := make(chan error, 1) + go func() { + _, readErr := connection.ReadFrameUntil(readContext) + readResult <- readErr + }() + <-serverReady + + // When + cancel() + + // Then + require.ErrorIs(t, <-readResult, context.Canceled) + require.Equal(t, closeErr, connection.Close()) + require.Equal(t, 1, recordedConnection.closeCount()) +} + +func TestClient_WebSocketClose_RetainsWriteCancellationCloseResult(t *testing.T) { + // Given + closeErr := errors.New("write cancellation close failed") + recordedConnection := newRetainedCloseConn(closeErr) + serverReady := make(chan struct{}) + server := newWebSocketTestServer(t, websocket.Upgrader{}, func(_ *websocket.Conn, request *http.Request) { + close(serverReady) + <-request.Context().Done() + }) + connection := dialRetainedCloseWebSocket(t, server.URL, recordedConnection) + <-serverReady + recordedConnection.blockWrites() + writeContext, cancel := context.WithCancel(context.Background()) + writeResult := make(chan error, 1) + go func() { + writeResult <- connection.WriteFrame(writeContext, Frame{Type: FrameBinary, Payload: []byte("blocked")}) + }() + recordedConnection.awaitWrite(t) + + // When + cancel() + + // Then + require.ErrorIs(t, <-writeResult, context.Canceled) + require.Equal(t, closeErr, connection.Close()) + require.Equal(t, 1, recordedConnection.closeCount()) +} + +type retainedCloseConn struct { + net.Conn + mu sync.Mutex + closeErr error + closeCountValue int + writeBlocked bool + writeEntered chan struct{} + writeReleased chan struct{} + releaseWrite sync.Once +} + +func newRetainedCloseConn(closeErr error) *retainedCloseConn { + return &retainedCloseConn{closeErr: closeErr, writeEntered: make(chan struct{}, 1), writeReleased: make(chan struct{})} +} + +func (connection *retainedCloseConn) Write(payload []byte) (int, error) { + connection.mu.Lock() + blocked := connection.writeBlocked + connection.mu.Unlock() + if blocked { + connection.writeEntered <- struct{}{} + <-connection.writeReleased + } + return connection.Conn.Write(payload) +} + +func (connection *retainedCloseConn) Close() error { + connection.mu.Lock() + connection.closeCountValue++ + connection.mu.Unlock() + underlyingErr := connection.Conn.Close() + connection.releaseWrite.Do(func() { close(connection.writeReleased) }) + if connection.closeErr != nil { + return connection.closeErr + } + return underlyingErr +} + +func (connection *retainedCloseConn) blockWrites() { + connection.mu.Lock() + defer connection.mu.Unlock() + connection.writeBlocked = true +} + +func (connection *retainedCloseConn) awaitWrite(t *testing.T) { + t.Helper() + select { + case <-connection.writeEntered: + case <-time.After(time.Second): + t.Fatal("WebSocket client did not enter the blocked write") + } +} + +func (connection *retainedCloseConn) closeCount() int { + connection.mu.Lock() + defer connection.mu.Unlock() + return connection.closeCountValue +} + +func dialRetainedCloseWebSocket(t *testing.T, baseURL string, recordedConnection *retainedCloseConn) *WebSocketConnection { + t.Helper() + dialer := *websocket.DefaultDialer + dialer.NetDialContext = func(ctx context.Context, network, address string) (net.Conn, error) { + connection, err := (&net.Dialer{}).DialContext(ctx, network, address) + if err != nil { + return nil, err + } + recordedConnection.Conn = connection + return recordedConnection, nil + } + client := newTestClient(t, Config{BaseURL: baseURL, RequestTimeout: time.Second, MaxResponseBytes: 1024, WebSocketDialer: &dialer}) + connection, err := client.DialWebSocket(context.Background(), "/retained-close") + require.NoError(t, err) + t.Cleanup(func() { require.ErrorIs(t, connection.Close(), recordedConnection.closeErr) }) + return connection +} diff --git a/integration/agentcompat/internal/client/websocket_read_contract_test.go b/integration/agentcompat/internal/client/websocket_read_contract_test.go new file mode 100644 index 00000000..93e980bf --- /dev/null +++ b/integration/agentcompat/internal/client/websocket_read_contract_test.go @@ -0,0 +1,236 @@ +package client + +import ( + "context" + "net" + "net/http" + "sync" + "testing" + "time" + + "github.com/gorilla/websocket" + "github.com/stretchr/testify/require" +) + +func TestClient_ReadFrameUntil_ClearsReadDeadlineOnUnderlyingConnection(t *testing.T) { + // Given + firstFrame := make(chan struct{}) + allowSecondFrame := make(chan struct{}) + server := newWebSocketTestServer(t, websocket.Upgrader{}, func(connection *websocket.Conn, _ *http.Request) { + require.NoError(t, connection.WriteMessage(websocket.TextMessage, []byte("first"))) + close(firstFrame) + <-allowSecondFrame + require.NoError(t, connection.WriteMessage(websocket.TextMessage, []byte("second"))) + }) + recordedConnection := newRecordingDeadlineConn() + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024, WebSocketDialer: recordingWebSocketDialer(recordedConnection)}) + connection, err := client.DialWebSocket(context.Background(), "/deadline-clear") + require.NoError(t, err) + t.Cleanup(func() { _ = connection.Close() }) + shortContext, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + <-firstFrame + _, err = connection.ReadFrame(shortContext) + require.NoError(t, err) + firstDeadline := recordedConnection.awaitReadDeadline(t) + require.False(t, firstDeadline.IsZero()) + + // When + result := make(chan error, 1) + go func() { + _, readErr := connection.ReadFrameUntil(context.Background()) + result <- readErr + }() + require.True(t, recordedConnection.awaitReadDeadline(t).IsZero()) + close(allowSecondFrame) + + // Then + require.NoError(t, <-result) +} + +func TestClient_ReadFrameUntil_UsesCallerDeadlineInsteadOfDefaultTimeout(t *testing.T) { + // Given + allowFrame := make(chan struct{}) + server := newWebSocketTestServer(t, websocket.Upgrader{}, func(connection *websocket.Conn, _ *http.Request) { + <-allowFrame + require.NoError(t, connection.WriteMessage(websocket.TextMessage, []byte("held"))) + }) + recordedConnection := newRecordingDeadlineConn() + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024, WebSocketDialer: recordingWebSocketDialer(recordedConnection)}) + connection, err := client.DialWebSocket(context.Background(), "/caller-deadline") + require.NoError(t, err) + t.Cleanup(func() { _ = connection.Close() }) + wantDeadline := time.Now().Add(time.Hour) + callerContext, cancel := context.WithDeadline(context.Background(), wantDeadline) + defer cancel() + result := make(chan error, 1) + + // When + go func() { + _, readErr := connection.ReadFrameUntil(callerContext) + result <- readErr + }() + + // Then + require.Equal(t, wantDeadline, recordedConnection.awaitReadDeadline(t)) + close(allowFrame) + require.NoError(t, <-result) +} + +func TestClient_ReadFrameUntil_ReturnsCancellationWhenItWinsAfterRead(t *testing.T) { + // Given + frameSent := make(chan struct{}) + server := newWebSocketTestServer(t, websocket.Upgrader{}, func(connection *websocket.Conn, _ *http.Request) { + require.NoError(t, connection.WriteMessage(websocket.TextMessage, []byte("first"))) + close(frameSent) + _, _, _ = connection.ReadMessage() + }) + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + connection, err := client.DialWebSocket(context.Background(), "/cancel-wins") + require.NoError(t, err) + t.Cleanup(func() { _ = connection.Close() }) + callerContext, cancel := context.WithCancel(context.Background()) + defer cancel() + connection.afterReadMessageForTest = cancel + <-frameSent + + // When + _, err = connection.ReadFrameUntil(callerContext) + + // Then + require.ErrorIs(t, err, context.Canceled) +} + +func TestClient_ReadFrameUntil_ReturnsCancellationWhenDeadlineSetIsInterrupted(t *testing.T) { + // Given + server := newWebSocketTestServer(t, websocket.Upgrader{}, func(_ *websocket.Conn, request *http.Request) { + <-request.Context().Done() + }) + recordedConnection := newRecordingDeadlineConn() + recordedConnection.readDeadlineEntered = make(chan struct{}, 1) + recordedConnection.allowReadDeadline = make(chan struct{}) + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024, WebSocketDialer: recordingWebSocketDialer(recordedConnection)}) + connection, err := client.DialWebSocket(context.Background(), "/deadline-cancel") + require.NoError(t, err) + t.Cleanup(func() { _ = connection.Close() }) + callerContext, cancel := context.WithCancel(context.Background()) + defer cancel() + result := make(chan error, 1) + + // When + go func() { + _, readErr := connection.ReadFrameUntil(callerContext) + result <- readErr + }() + recordedConnection.awaitReadDeadlineEntered(t) + cancel() + recordedConnection.awaitClose(t) + close(recordedConnection.allowReadDeadline) + + // Then + require.ErrorIs(t, <-result, context.Canceled) +} + +func TestClient_ReadFrameUntil_KeepsConnectionOpenWhenCanceledAfterSuccess(t *testing.T) { + // Given + firstFrameSent := make(chan struct{}) + allowSecondFrame := make(chan struct{}) + server := newWebSocketTestServer(t, websocket.Upgrader{}, func(connection *websocket.Conn, _ *http.Request) { + require.NoError(t, connection.WriteMessage(websocket.TextMessage, []byte("first"))) + close(firstFrameSent) + <-allowSecondFrame + require.NoError(t, connection.WriteMessage(websocket.TextMessage, []byte("second"))) + }) + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + connection, err := client.DialWebSocket(context.Background(), "/success-cancel") + require.NoError(t, err) + t.Cleanup(func() { _ = connection.Close() }) + callerContext, cancel := context.WithCancel(context.Background()) + defer cancel() + <-firstFrameSent + + // When + firstFrame, err := connection.ReadFrameUntil(callerContext) + require.NoError(t, err) + cancel() + close(allowSecondFrame) + secondFrame, err := connection.ReadFrameUntil(context.Background()) + + // Then + require.Equal(t, []byte("first"), firstFrame.Payload) + require.NoError(t, err) + require.Equal(t, []byte("second"), secondFrame.Payload) +} + +type recordingDeadlineConn struct { + net.Conn + mu sync.Mutex + readDeadlines []time.Time + readDeadlineCalls chan time.Time + readDeadlineEntered chan struct{} + allowReadDeadline chan struct{} + closeCalls chan struct{} +} + +func (connection *recordingDeadlineConn) SetReadDeadline(deadline time.Time) error { + connection.mu.Lock() + connection.readDeadlines = append(connection.readDeadlines, deadline) + connection.mu.Unlock() + connection.readDeadlineCalls <- deadline + if connection.readDeadlineEntered != nil { + connection.readDeadlineEntered <- struct{}{} + <-connection.allowReadDeadline + } + return connection.Conn.SetReadDeadline(deadline) +} + +func (connection *recordingDeadlineConn) Close() error { + connection.closeCalls <- struct{}{} + return connection.Conn.Close() +} + +func newRecordingDeadlineConn() *recordingDeadlineConn { + return &recordingDeadlineConn{readDeadlineCalls: make(chan time.Time, 4), closeCalls: make(chan struct{}, 2)} +} + +func (connection *recordingDeadlineConn) awaitReadDeadline(t *testing.T) time.Time { + t.Helper() + select { + case deadline := <-connection.readDeadlineCalls: + return deadline + case <-time.After(time.Second): + t.Fatal("WebSocket client did not set a read deadline") + return time.Time{} + } +} + +func (connection *recordingDeadlineConn) awaitReadDeadlineEntered(t *testing.T) { + t.Helper() + select { + case <-connection.readDeadlineEntered: + case <-time.After(time.Second): + t.Fatal("WebSocket client did not enter SetReadDeadline") + } +} + +func (connection *recordingDeadlineConn) awaitClose(t *testing.T) { + t.Helper() + select { + case <-connection.closeCalls: + case <-time.After(time.Second): + t.Fatal("WebSocket client did not close after cancellation") + } +} + +func recordingWebSocketDialer(recordedConnection *recordingDeadlineConn) *websocket.Dialer { + dialer := *websocket.DefaultDialer + dialer.NetDialContext = func(ctx context.Context, network, address string) (net.Conn, error) { + connection, err := (&net.Dialer{}).DialContext(ctx, network, address) + if err != nil { + return nil, err + } + recordedConnection.Conn = connection + return recordedConnection, nil + } + return &dialer +} diff --git a/integration/agentcompat/internal/client/websocket_test.go b/integration/agentcompat/internal/client/websocket_test.go new file mode 100644 index 00000000..974c7563 --- /dev/null +++ b/integration/agentcompat/internal/client/websocket_test.go @@ -0,0 +1,230 @@ +package client + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gorilla/websocket" + "github.com/stretchr/testify/require" +) + +func TestClient_WebSocketFrameReassembly(t *testing.T) { + upgrader := websocket.Upgrader{WriteBufferSize: 4} + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + require.Equal(t, "Bearer ws-token", request.Header.Get("Authorization")) + require.NotEmpty(t, request.Header.Get("Origin")) + connection, err := upgrader.Upgrade(writer, request, nil) + require.NoError(t, err) + defer connection.Close() + + textWriter, err := connection.NextWriter(websocket.TextMessage) + require.NoError(t, err) + _, err = textWriter.Write([]byte("hello ")) + require.NoError(t, err) + _, err = textWriter.Write([]byte("world")) + require.NoError(t, err) + require.NoError(t, textWriter.Close()) + + binaryWriter, err := connection.NextWriter(websocket.BinaryMessage) + require.NoError(t, err) + _, err = binaryWriter.Write([]byte{1, 2}) + require.NoError(t, err) + _, err = binaryWriter.Write([]byte{3, 4}) + require.NoError(t, err) + require.NoError(t, binaryWriter.Close()) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{ + BaseURL: server.URL, + BearerToken: "ws-token", + Origin: server.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + }) + connection, err := client.DialWebSocket(context.Background(), "/stream") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, connection.Close()) }) + + textFrame, err := connection.ReadFrame(context.Background()) + require.NoError(t, err) + require.Equal(t, FrameText, textFrame.Type) + require.Equal(t, []byte("hello world"), textFrame.Payload) + + binaryFrame, err := connection.ReadFrame(context.Background()) + require.NoError(t, err) + require.Equal(t, FrameBinary, binaryFrame.Type) + require.Equal(t, []byte{1, 2, 3, 4}, binaryFrame.Payload) +} + +func TestClient_WebSocketRejectsOversize(t *testing.T) { + upgrader := websocket.Upgrader{} + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + connection, err := upgrader.Upgrade(writer, request, nil) + require.NoError(t, err) + defer connection.Close() + require.NoError(t, connection.WriteMessage(websocket.BinaryMessage, []byte(strings.Repeat("x", 65)))) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{ + BaseURL: server.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 64, + }) + connection, err := client.DialWebSocket(context.Background(), "/oversize") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, connection.Close()) }) + + _, err = connection.ReadFrame(context.Background()) + require.ErrorIs(t, err, ErrResponseTooLarge) +} + +func TestClient_DialWebSocket_ReturnsTypedHandshakeFailure(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + http.Error(writer, "permission denied", http.StatusForbidden) + })) + t.Cleanup(server.Close) + webSocketClient := newTestClient(t, Config{BaseURL: server.URL}) + + _, err := webSocketClient.DialWebSocket(context.Background(), "/terminal") + + var handshakeError *WebSocketHandshakeError + require.ErrorAs(t, err, &handshakeError) + require.Equal(t, http.StatusForbidden, handshakeError.StatusCode) + require.Contains(t, handshakeError.Message, "permission denied") +} + +func TestClient_WebSocketDeadline(t *testing.T) { + upgrader := websocket.Upgrader{} + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + connection, err := upgrader.Upgrade(writer, request, nil) + require.NoError(t, err) + defer connection.Close() + _, _, _ = connection.ReadMessage() + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{ + BaseURL: server.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + }) + connection, err := client.DialWebSocket(context.Background(), "/deadline") + require.NoError(t, err) + connection.timeout = 20 * time.Millisecond + + _, err = connection.ReadFrame(context.Background()) + require.True(t, errors.Is(err, context.DeadlineExceeded), err) + require.NoError(t, connection.Close()) +} + +func TestClient_RedactsAuthorization(t *testing.T) { + secret := "nzp_super-secret" + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"success":false,"error":"Authorization: Bearer ` + secret + `"}`)) + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{ + BaseURL: server.URL, + BearerToken: secret, + Origin: server.URL, + RequestTimeout: time.Second, + MaxResponseBytes: 1024, + }) + _, err := REST[struct{}, semanticResult](context.Background(), client, RESTRequest[struct{}]{Method: http.MethodGet, Path: "/redaction"}) + require.Error(t, err) + require.NotContains(t, err.Error(), secret) + require.Contains(t, err.Error(), "[REDACTED]") + + redacted := Redact("request failed with Authorization: Bearer " + secret) + require.False(t, strings.Contains(redacted, secret)) + require.Contains(t, redacted, "Authorization: Bearer [REDACTED]") + + quoted := Redact(`{"Authorization":"Bearer ` + secret + `"}`) + require.NotContains(t, quoted, secret) + require.Contains(t, quoted, "[REDACTED]") + + query := Redact("https://dashboard.example/mcp/download/path-token?access_token=" + secret + "&X-Amz-Signature=signature-secret") + require.NotContains(t, query, secret) + require.NotContains(t, query, "signature-secret") + + httpError := &HTTPError{StatusCode: http.StatusBadRequest, Message: Redact("Authorization: Bearer " + secret)} + require.NotContains(t, httpError.Message, secret) + rpcError := &RPCError{Code: -32603, Message: Redact("token=" + secret)} + require.NotContains(t, rpcError.Message, secret) +} + +func TestClient_RedactsCredentialClasses(t *testing.T) { + secret := "sensitive-value" + inputs := []string{ + "X-CSRF-Token: " + secret, + "password=" + secret, + "https://dashboard.example/path?access_token=" + secret, + "jwt_token: " + secret, + } + for _, input := range inputs { + redacted := Redact(input) + require.NotContains(t, redacted, secret) + require.Contains(t, redacted, "[REDACTED]") + } +} + +func TestClient_WebSocketReadDeadline(t *testing.T) { + upgrader := websocket.Upgrader{} + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + connection, err := upgrader.Upgrade(writer, request, nil) + require.NoError(t, err) + defer connection.Close() + <-request.Context().Done() + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: 20 * time.Millisecond, MaxResponseBytes: 1024}) + connection, err := client.DialWebSocket(context.Background(), "/stream") + require.NoError(t, err) + t.Cleanup(func() { _ = connection.Close() }) + + _, err = connection.ReadFrame(context.Background()) + require.ErrorIs(t, err, context.DeadlineExceeded) +} + +func TestClient_WebSocketWriteHonorsParentCancellation(t *testing.T) { + upgrader := websocket.Upgrader{} + serverReady := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + connection, err := upgrader.Upgrade(writer, request, nil) + require.NoError(t, err) + defer connection.Close() + close(serverReady) + <-request.Context().Done() + })) + t.Cleanup(server.Close) + + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: 5 * time.Second, MaxResponseBytes: 1024}) + connection, err := client.DialWebSocket(context.Background(), "/blocked-write") + require.NoError(t, err) + t.Cleanup(func() { _ = connection.Close() }) + + writeContext, cancel := context.WithCancel(context.Background()) + defer cancel() + result := make(chan error, 1) + <-serverReady + go func() { + result <- connection.WriteFrame(writeContext, Frame{Type: FrameBinary, Payload: make([]byte, 128<<20)}) + }() + cancel() + select { + case err = <-result: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("WebSocket write did not stop after parent cancellation") + } +} diff --git a/integration/agentcompat/internal/client/websocket_until_test.go b/integration/agentcompat/internal/client/websocket_until_test.go new file mode 100644 index 00000000..8febf9ad --- /dev/null +++ b/integration/agentcompat/internal/client/websocket_until_test.go @@ -0,0 +1,103 @@ +package client + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/gorilla/websocket" + "github.com/stretchr/testify/require" +) + +func TestClient_ReadFrameUntil_UsesCallerDeadlineInsteadOfRequestTimeout(t *testing.T) { + // Given + upgrader := websocket.Upgrader{} + serverReady := make(chan struct{}) + allowFrame := make(chan struct{}) + server := newWebSocketTestServer(t, upgrader, func(connection *websocket.Conn, _ *http.Request) { + close(serverReady) + <-allowFrame + require.NoError(t, connection.WriteMessage(websocket.BinaryMessage, []byte("held-session"))) + }) + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: 100 * time.Millisecond, MaxResponseBytes: 1024}) + connection, err := client.DialWebSocket(context.Background(), "/held") + require.NoError(t, err) + t.Cleanup(func() { _ = connection.Close() }) + callerContext, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + result := make(chan struct { + frame Frame + err error + }, 1) + + // When + go func() { + frame, readErr := connection.ReadFrameUntil(callerContext) + result <- struct { + frame Frame + err error + }{frame: frame, err: readErr} + }() + <-serverReady + requestTimeout, cancelRequestTimeout := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancelRequestTimeout() + <-requestTimeout.Done() + close(allowFrame) + + // Then + select { + case readResult := <-result: + require.NoError(t, readResult.err) + require.Equal(t, FrameBinary, readResult.frame.Type) + require.Equal(t, []byte("held-session"), readResult.frame.Payload) + case <-time.After(time.Second): + t.Fatal("ReadFrameUntil did not receive the channel-released frame") + } +} + +func TestClient_ReadFrameUntil_ReturnsParentCancellationAndUnblocksRead(t *testing.T) { + // Given + upgrader := websocket.Upgrader{} + serverReady := make(chan struct{}) + server := newWebSocketTestServer(t, upgrader, func(_ *websocket.Conn, request *http.Request) { + close(serverReady) + <-request.Context().Done() + }) + client := newTestClient(t, Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024}) + connection, err := client.DialWebSocket(context.Background(), "/cancel") + require.NoError(t, err) + t.Cleanup(func() { _ = connection.Close() }) + callerContext, cancel := context.WithCancel(context.Background()) + defer cancel() + result := make(chan error, 1) + + // When + go func() { + _, readErr := connection.ReadFrameUntil(callerContext) + result <- readErr + }() + <-serverReady + cancel() + + // Then + select { + case readErr := <-result: + require.ErrorIs(t, readErr, context.Canceled) + case <-time.After(time.Second): + t.Fatal("ReadFrameUntil did not stop after parent cancellation") + } +} + +func newWebSocketTestServer(t *testing.T, upgrader websocket.Upgrader, serve func(*websocket.Conn, *http.Request)) *httptest.Server { + t.Helper() + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + connection, err := upgrader.Upgrade(writer, request, nil) + require.NoError(t, err) + defer connection.Close() + serve(connection, request) + })) + t.Cleanup(server.Close) + return server +}