mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 17:50:12 +00:00
test(agentcompat): add protocol clients
Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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]`)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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())
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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})
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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"`
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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))}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user