From 50bbf26bee4c48ab3181d4a27b4a4d65e509da63 Mon Sep 17 00:00:00 2001 From: Jairus Christensen Date: Fri, 21 Aug 2026 11:14:23 -0600 Subject: [PATCH 1/3] Honor configured vMCP backend timeouts Apply operational timeout defaults and per-workload overrides to Modern and Legacy backend operations, including SSE operation contexts and session initialization. Defer the MCP POST write deadline during backend computation, then restore a bounded deadline when a non-streaming response begins. Signed-off-by: Jairus Christensen --- pkg/transport/middleware/write_timeout.go | 69 +++++++- .../middleware/write_timeout_test.go | 45 +++++- pkg/vmcp/cli/serve.go | 27 +++- pkg/vmcp/cli/serve_test.go | 47 ++++++ pkg/vmcp/client/client.go | 71 ++++++++- pkg/vmcp/client/client_test.go | 147 ++++++++++++++++++ pkg/vmcp/server/server.go | 16 +- pkg/vmcp/server/testutil_test.go | 10 ++ .../server/write_timeout_integration_test.go | 33 ++++ .../session/connector_integration_test.go | 80 ++++++++++ pkg/vmcp/session/default_session_test.go | 36 +++++ pkg/vmcp/session/factory.go | 43 ++++- .../session/internal/backend/mcp_session.go | 73 +++++++-- .../internal/backend/mcp_session_test.go | 8 +- 14 files changed, 661 insertions(+), 44 deletions(-) diff --git a/pkg/transport/middleware/write_timeout.go b/pkg/transport/middleware/write_timeout.go index 85a46fb251..20cbfe2620 100644 --- a/pkg/transport/middleware/write_timeout.go +++ b/pkg/transport/middleware/write_timeout.go @@ -7,22 +7,72 @@ import ( "log/slog" "net/http" "strings" + "sync" "time" ) -// WriteTimeout clears the write deadline for qualifying SSE connections -// (GET + Accept: text/event-stream + matching path) so http.Server.WriteTimeout -// does not kill long-lived streams (golang/go#16100). All other requests are -// left untouched. -func WriteTimeout(endpointPath string) func(http.Handler) http.Handler { +const defaultResponseWriteTimeout = 30 * time.Second + +type postResponseWriter struct { + http.ResponseWriter + writeTimeout time.Duration + armOnce sync.Once +} + +func (w *postResponseWriter) Unwrap() http.ResponseWriter { return w.ResponseWriter } + +func (w *postResponseWriter) armWriteDeadline() { + w.armOnce.Do(func() { + // A streamable-HTTP POST may legitimately negotiate an SSE response. + // Like a qualifying GET stream, its lifetime is governed by request + // cancellation rather than a socket write deadline. + if strings.Contains(w.Header().Get("Content-Type"), "text/event-stream") { + return + } + if err := http.NewResponseController(w.ResponseWriter).SetWriteDeadline(time.Now().Add(w.writeTimeout)); err != nil { + slog.Warn("failed to arm MCP response write deadline", "error", err) + } + }) +} + +func (w *postResponseWriter) WriteHeader(statusCode int) { + w.armWriteDeadline() + w.ResponseWriter.WriteHeader(statusCode) +} + +func (w *postResponseWriter) Write(p []byte) (int, error) { + w.armWriteDeadline() + return w.ResponseWriter.Write(p) +} + +func (w *postResponseWriter) Flush() { + w.armWriteDeadline() + if err := http.NewResponseController(w.ResponseWriter).Flush(); err != nil { + slog.Debug("failed to flush MCP response", "error", err) + } +} + +// WriteTimeout clears the server-level write deadline while an MCP POST is +// computing, then arms a fresh deadline when a non-streaming response begins. +// Qualifying SSE requests remain unbounded and are canceled through request +// contexts. Other requests are left untouched. The optional duration controls +// the response-write allowance and defaults to 30 seconds. +func WriteTimeout(endpointPath string, responseWriteTimeout ...time.Duration) func(http.Handler) http.Handler { + writeTimeout := defaultResponseWriteTimeout + if len(responseWriteTimeout) > 0 && responseWriteTimeout[0] > 0 { + writeTimeout = responseWriteTimeout[0] + } + return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.Method == http.MethodGet && + isMCPPost := r.Method == http.MethodPost && r.URL.Path == endpointPath + isMCPStream := r.Method == http.MethodGet && strings.Contains(r.Header.Get("Accept"), "text/event-stream") && - r.URL.Path == endpointPath { + r.URL.Path == endpointPath + if isMCPPost || isMCPStream { rc := http.NewResponseController(w) if err := rc.SetWriteDeadline(time.Time{}); err != nil { - slog.Warn("failed to clear write deadline for SSE connection; stream may be killed by server WriteTimeout", + slog.Warn("failed to clear write deadline for MCP request; request may be killed by server WriteTimeout", "error", err, "method", r.Method, "path", r.URL.Path, @@ -30,6 +80,9 @@ func WriteTimeout(endpointPath string) func(http.Handler) http.Handler { ) } } + if isMCPPost { + w = &postResponseWriter{ResponseWriter: w, writeTimeout: writeTimeout} + } next.ServeHTTP(w, r) }) } diff --git a/pkg/transport/middleware/write_timeout_test.go b/pkg/transport/middleware/write_timeout_test.go index 7cd9c9857c..7fa0b7160b 100644 --- a/pkg/transport/middleware/write_timeout_test.go +++ b/pkg/transport/middleware/write_timeout_test.go @@ -27,11 +27,13 @@ type deadlineTrackingResponseWriter struct { *httptest.ResponseRecorder deadlineSet bool deadline time.Time + deadlines []time.Time } func (d *deadlineTrackingResponseWriter) SetWriteDeadline(t time.Time) error { d.deadlineSet = true d.deadline = t + d.deadlines = append(d.deadlines, t) return nil } @@ -97,9 +99,10 @@ func TestWriteTimeout_GETOnWrongPathLeavesDeadlineUntouched(t *testing.T) { assert.Equal(t, http.StatusOK, w.Code) } -// TestWriteTimeout_POSTLeavesDeadlineUntouched verifies that POST requests are not -// touched by the middleware — their deadline comes from http.Server.WriteTimeout. -func TestWriteTimeout_POSTLeavesDeadlineUntouched(t *testing.T) { +// TestWriteTimeout_MCPPOSTDefersDeadlineUntilResponse verifies that MCP POST +// requests are unbounded while computing and receive a fresh write deadline +// when their non-streaming response starts. +func TestWriteTimeout_MCPPOSTDefersDeadlineUntilResponse(t *testing.T) { t.Parallel() w := newDeadlineTracker() @@ -107,7 +110,41 @@ func TestWriteTimeout_POSTLeavesDeadlineUntouched(t *testing.T) { mw(noopHandler).ServeHTTP(w, r) - assert.False(t, w.deadlineSet, "POST deadline is managed by http.Server.WriteTimeout, not the middleware") + require.Len(t, w.deadlines, 2) + assert.True(t, w.deadlines[0].IsZero(), "the handler computation deadline must be cleared") + assert.False(t, w.deadlines[1].IsZero(), "response writing must receive a bounded deadline") + assert.True(t, w.deadlines[1].After(time.Now()), "response write deadline must be in the future") + assert.Equal(t, http.StatusOK, w.Code) +} + +func TestWriteTimeout_MCPPostSSEResponseRemainsUnbounded(t *testing.T) { + t.Parallel() + + w := newDeadlineTracker() + r := httptest.NewRequest(http.MethodPost, testEndpointPath, nil) + handler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(http.StatusOK) + }) + + mw(handler).ServeHTTP(w, r) + + require.Len(t, w.deadlines, 1) + assert.True(t, w.deadlines[0].IsZero(), "SSE response must retain the cleared deadline") + assert.Equal(t, http.StatusOK, w.Code) +} + +// TestWriteTimeout_POSTOnWrongPathLeavesDeadlineUntouched verifies that only +// the configured MCP endpoint gets the relaxed POST deadline. +func TestWriteTimeout_POSTOnWrongPathLeavesDeadlineUntouched(t *testing.T) { + t.Parallel() + + w := newDeadlineTracker() + r := httptest.NewRequest(http.MethodPost, "/health", nil) + + mw(noopHandler).ServeHTTP(w, r) + + assert.False(t, w.deadlineSet, "POST on non-MCP path must retain the server WriteTimeout") assert.Equal(t, http.StatusOK, w.Code) } diff --git a/pkg/vmcp/cli/serve.go b/pkg/vmcp/cli/serve.go index 0f7efa9bcb..f89d216498 100644 --- a/pkg/vmcp/cli/serve.go +++ b/pkg/vmcp/cli/serve.go @@ -354,6 +354,10 @@ func Serve(ctx context.Context, cfg ServeConfig) error { if revisions, ok := backendClient.(vmcp.RevisionReporter); ok { sessionFactoryOpts = append(sessionFactoryOpts, vmcpsession.WithRevisionLookup(revisions.CachedRevision)) } + sessionFactoryOpts = append( + sessionFactoryOpts, + vmcpsession.WithRequestTimeoutResolver(backendRequestTimeoutResolver(vmcpCfg)), + ) sessionFactory := vmcpsession.NewSessionFactory(outgoingRegistry, sessionFactoryOpts...) // When the optimizer is enabled, its meta-tools are pass-through tools. @@ -520,6 +524,24 @@ func getStatusReportingInterval(cfg *config.Config) time.Duration { return 0 } +// backendRequestTimeoutResolver resolves the documented operational timeout +// for a backend workload. Configuration loaded from YAML has defaults applied +// and is immutable after startup; the defensive fallback also supports quick +// mode and direct embedders that omit Operational. +func backendRequestTimeoutResolver(cfg *config.Config) func(workloadID string) time.Duration { + timeouts := config.DefaultOperationalConfig().Timeouts + if cfg != nil && cfg.Operational != nil && cfg.Operational.Timeouts != nil { + timeouts = cfg.Operational.Timeouts + } + + return func(workloadID string) time.Duration { + if timeout, ok := timeouts.PerWorkload[workloadID]; ok && timeout > 0 { + return time.Duration(timeout) + } + return time.Duration(timeouts.Default) + } +} + // loadAndValidateConfig loads and validates the vMCP configuration file. func loadAndValidateConfig(configPath string) (*config.Config, error) { slog.Info(fmt.Sprintf("Loading configuration from: %s", configPath)) @@ -624,7 +646,10 @@ func discoverBackends( return nil, nil, nil, fmt.Errorf("failed to create outgoing authentication registry: %w", err) } - backendClient, err := vmcpclient.NewHTTPBackendClient(outgoingRegistry) + backendClient, err := vmcpclient.NewHTTPBackendClient( + outgoingRegistry, + vmcpclient.WithRequestTimeoutResolver(backendRequestTimeoutResolver(cfg)), + ) if err != nil { return nil, nil, nil, fmt.Errorf("failed to create backend client: %w", err) } diff --git a/pkg/vmcp/cli/serve_test.go b/pkg/vmcp/cli/serve_test.go index 88e05312cd..4054452718 100644 --- a/pkg/vmcp/cli/serve_test.go +++ b/pkg/vmcp/cli/serve_test.go @@ -8,6 +8,7 @@ import ( "os" "path/filepath" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -23,6 +24,52 @@ import ( vmcpmocks "github.com/stacklok/toolhive/pkg/vmcp/mocks" ) +func TestBackendRequestTimeoutResolver(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + cfg *config.Config + workloadID string + want time.Duration + }{ + { + name: "missing operational config uses default", + cfg: &config.Config{}, + workloadID: "elastic", + want: 30 * time.Second, + }, + { + name: "configured default applies to unmatched workload", + cfg: &config.Config{Operational: &config.OperationalConfig{ + Timeouts: &config.TimeoutConfig{Default: config.Duration(90 * time.Second)}, + }}, + workloadID: "other", + want: 90 * time.Second, + }, + { + name: "per workload timeout overrides configured default", + cfg: &config.Config{Operational: &config.OperationalConfig{ + Timeouts: &config.TimeoutConfig{ + Default: config.Duration(90 * time.Second), + PerWorkload: map[string]config.Duration{ + "elastic": config.Duration(240 * time.Second), + }, + }, + }}, + workloadID: "elastic", + want: 240 * time.Second, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + assert.Equal(t, tt.want, backendRequestTimeoutResolver(tt.cfg)(tt.workloadID)) + }) + } +} + // TestLoadAndValidateConfig covers all config-loading paths. func TestLoadAndValidateConfig(t *testing.T) { t.Parallel() diff --git a/pkg/vmcp/client/client.go b/pkg/vmcp/client/client.go index 54797566c4..4a97b09d65 100644 --- a/pkg/vmcp/client/client.go +++ b/pkg/vmcp/client/client.go @@ -65,6 +65,10 @@ const ( // A tools/list response with 1000 tools would be limited to 100MB total. maxResponseSize = 100 * 1024 * 1024 // 100 MB + // defaultBackendRequestTimeout bounds individual backend operations when no + // workload-specific timeout is configured. + defaultBackendRequestTimeout = 30 * time.Second + // defaultRefutationTTL is the default lifetime of a modernHintRefuted // entry. Chosen to be long enough that a persistently hint-lying backend // (#6154) pays a negligible probe cost (one server/discover per backend @@ -129,6 +133,22 @@ func WithRefutationTTL(d time.Duration) Option { } } +// WithRequestTimeoutResolver configures the timeout used for each backend +// operation. The resolver receives the backend workload ID and may return a +// workload-specific duration. A nil resolver, or a non-positive result, uses +// the 30-second default. +// +// The resolver may be called concurrently and must therefore be safe for +// concurrent use. SSE connection lifetimes remain unbounded, but individual +// operations on those connections are bounded through request contexts. +func WithRequestTimeoutResolver(resolver func(workloadID string) time.Duration) Option { + return func(h *httpBackendClient) { + if resolver != nil { + h.requestTimeoutResolver = resolver + } + } +} + // httpBackendClient implements vmcp.BackendClient using stacklok/toolhive-core/mcpcompat HTTP client. // It supports streamable-HTTP and SSE transports for backend MCP servers. type httpBackendClient struct { @@ -209,6 +229,11 @@ type httpBackendClient struct { // confirming probe in legacyInit. Defaults to defaultRefutationTTL; // configurable via WithRefutationTTL (tests). refutationTTL time.Duration + + // requestTimeoutResolver returns the wall-clock timeout for an individual + // backend operation by workload ID. It is immutable after construction and + // may be read concurrently. + requestTimeoutResolver func(workloadID string) time.Duration } // NewHTTPBackendClient creates a new HTTP-based backend client. @@ -240,6 +265,24 @@ func NewHTTPBackendClient(registry vmcpauth.OutgoingAuthRegistry, opts ...Option return c, nil } +// requestTimeout returns the configured request timeout for workloadID, +// falling back to the historical 30-second default when no positive override +// is available. +func (h *httpBackendClient) requestTimeout(workloadID string) time.Duration { + if h.requestTimeoutResolver != nil { + if timeout := h.requestTimeoutResolver(workloadID); timeout > 0 { + return timeout + } + } + return defaultBackendRequestTimeout +} + +func (h *httpBackendClient) requestContext( + ctx context.Context, target *vmcp.BackendTarget, +) (context.Context, context.CancelFunc) { + return context.WithTimeout(ctx, h.requestTimeout(target.WorkloadID)) +} + // backendDialer returns a net.Dialer with the standard backend timeouts and an // optional Control hook. Centralising the timeout constants here ensures the // fallback-construction branch and the dial-control replacement branch always @@ -517,7 +560,7 @@ func (h *httpBackendClient) resolveAuthStrategy(target *vmcp.BackendTarget) (vmc // sampling handlers so a backend's mid-call server->client traffic reaches the // downstream client. Non-forwarding calls get the plain client (no standalone GET // stream), which is byte-for-byte the pre-forwarding construction. -func (*httpBackendClient) newStreamableHTTPClient( +func (h *httpBackendClient) newStreamableHTTPClient( ctx context.Context, target *vmcp.BackendTarget, baseTransport http.RoundTripper, forwarding bool, fwd *boundForwarders, ) (*client.Client, error) { @@ -537,9 +580,10 @@ func (*httpBackendClient) newStreamableHTTPClient( } return resp, nil }) - httpClient := newBackendHTTPClient(sizeLimitedTransport, 30*time.Second) + requestTimeout := h.requestTimeout(target.WorkloadID) + httpClient := newBackendHTTPClient(sizeLimitedTransport, requestTimeout) transportOpts := []transport.StreamableHTTPCOption{ - transport.WithHTTPTimeout(30 * time.Second), + transport.WithHTTPTimeout(requestTimeout), transport.WithHTTPBasicClient(httpClient), } if fwd != nil && forwarding { @@ -1052,15 +1096,15 @@ func (h *httpBackendClient) CachedRevision(workloadID string) (mcpparser.Revisio // identity, header-forward, trace, TLS/SSRF — see buildBackendRoundTripper) in an // *http.Client for the raw Modern shim. This is a LIVE production path: the // discover probe must carry the same security controls as every other backend -// call, so it must NOT use a bare http.Client. A 30s timeout matches the -// streamable-HTTP client; the response body is bounded inside modernCall +// call, so it must NOT use a bare http.Client. Its workload-aware timeout +// matches the streamable-HTTP client; the response body is bounded inside modernCall // (io.LimitReader), so no size-limit transport wrapper is needed here. func (h *httpBackendClient) buildModernHTTPClient(ctx context.Context, target *vmcp.BackendTarget) (*http.Client, error) { rt, err := h.buildBackendRoundTripper(ctx, target) if err != nil { return nil, err } - return newBackendHTTPClient(rt, 30*time.Second), nil + return newBackendHTTPClient(rt, h.requestTimeout(target.WorkloadID)), nil } // newBackendHTTPClient is the single choke point for every backend *http.Client @@ -1399,6 +1443,9 @@ func newCapabilityListFromMCP( // first-probe blip that pinned Legacy) self-corrects: the failing path triggers a // re-probe and one retry under the corrected revision. func (h *httpBackendClient) ListCapabilities(ctx context.Context, target *vmcp.BackendTarget) (*vmcp.CapabilityList, error) { + ctx, cancel := h.requestContext(ctx, target) + defer cancel() + slog.Debug("querying capabilities from backend", "backend", target.WorkloadName, "url", target.BaseURL) var out *vmcp.CapabilityList err := h.dispatch(ctx, target, func(ctx context.Context, rev mcpparser.Revision) error { @@ -1714,6 +1761,9 @@ func (h *httpBackendClient) CallTool( meta map[string]any, paramHeaders map[string]string, ) (*vmcp.ToolCallResult, error) { + ctx, cancel := h.requestContext(ctx, target) + defer cancel() + slog.Debug("calling tool on backend", "tool", toolName, "backend", target.WorkloadName) var out *vmcp.ToolCallResult err := h.dispatch(ctx, target, func(ctx context.Context, rev mcpparser.Revision) error { @@ -1921,6 +1971,9 @@ func toolResultFromMCP(result *mcp.CallToolResult, toolName, backendID string) * func (h *httpBackendClient) ReadResource( ctx context.Context, target *vmcp.BackendTarget, uri string, ) (*vmcp.ResourceReadResult, error) { + ctx, cancel := h.requestContext(ctx, target) + defer cancel() + slog.Debug("reading resource from backend", "resource", uri, "backend", target.WorkloadName) var out *vmcp.ResourceReadResult err := h.dispatch(ctx, target, func(ctx context.Context, rev mcpparser.Revision) error { @@ -2038,6 +2091,9 @@ func (h *httpBackendClient) GetPrompt( name string, arguments map[string]any, ) (*vmcp.PromptGetResult, error) { + ctx, cancel := h.requestContext(ctx, target) + defer cancel() + slog.Debug("getting prompt from backend", "prompt", name, "backend", target.WorkloadName) var out *vmcp.PromptGetResult err := h.dispatch(ctx, target, func(ctx context.Context, rev mcpparser.Revision) error { @@ -2158,6 +2214,9 @@ func (h *httpBackendClient) Complete( argName, argValue string, contextArgs map[string]string, ) (*vmcp.CompletionResult, error) { + ctx, cancel := h.requestContext(ctx, target) + defer cancel() + slog.Debug("requesting completion from backend", "ref_type", ref.Type, "backend", target.WorkloadName) var out *vmcp.CompletionResult err := h.dispatch(ctx, target, func(ctx context.Context, rev mcpparser.Revision) error { diff --git a/pkg/vmcp/client/client_test.go b/pkg/vmcp/client/client_test.go index efc7c50c3c..48e5c8d13c 100644 --- a/pkg/vmcp/client/client_test.go +++ b/pkg/vmcp/client/client_test.go @@ -754,6 +754,153 @@ func TestNewHTTPBackendClient_NilRegistry(t *testing.T) { }) } +func TestHTTPBackendClient_RequestTimeout(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + resolver func(string) time.Duration + want time.Duration + }{ + {name: "nil resolver uses default", want: defaultBackendRequestTimeout}, + { + name: "positive workload timeout is used", + resolver: func(string) time.Duration { + return 240 * time.Second + }, + want: 240 * time.Second, + }, + { + name: "zero timeout uses default", + resolver: func(string) time.Duration { return 0 }, + want: defaultBackendRequestTimeout, + }, + { + name: "negative timeout uses default", + resolver: func(string) time.Duration { return -time.Second }, + want: defaultBackendRequestTimeout, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + backendClient := &httpBackendClient{requestTimeoutResolver: tt.resolver} + assert.Equal(t, tt.want, backendClient.requestTimeout("elastic")) + }) + } +} + +func TestHTTPBackendClient_ModernHTTPClientUsesWorkloadTimeout(t *testing.T) { + t.Parallel() + + registry := auth.NewDefaultOutgoingAuthRegistry() + require.NoError(t, registry.RegisterStrategy( + authtypes.StrategyTypeUnauthenticated, &strategies.UnauthenticatedStrategy{}, + )) + backendClient, err := NewHTTPBackendClient( + registry, + WithRequestTimeoutResolver(func(workloadID string) time.Duration { + if workloadID == "elastic" { + return 240 * time.Second + } + return defaultBackendRequestTimeout + }), + ) + require.NoError(t, err) + + clientImpl := backendClient.(*httpBackendClient) + httpClient, err := clientImpl.buildModernHTTPClient(t.Context(), &vmcp.BackendTarget{ + WorkloadID: "elastic", + BaseURL: "http://127.0.0.1:1", + }) + require.NoError(t, err) + assert.Equal(t, 240*time.Second, httpClient.Timeout) +} + +// TestHTTPBackendClient_RequestTimeoutIntegration proves that the resolved +// workload timeout reaches the real streamable-HTTP tools/call path used by the +// vMCP core. The paired timeout/success cases prevent the historical fixed +// 30-second timeout from satisfying the test accidentally. +func TestHTTPBackendClient_RequestTimeoutIntegration(t *testing.T) { + t.Parallel() + + const toolDelay = 300 * time.Millisecond + + tests := []struct { + name string + resolver func(string) time.Duration + wantTimeoutErr bool + }{ + { + name: "default timeout is enforced", + resolver: func(string) time.Duration { return 100 * time.Millisecond }, + wantTimeoutErr: true, + }, + { + name: "workload override permits slow response", + resolver: func(workloadID string) time.Duration { + if workloadID == "elastic" { + return 2 * time.Second + } + return 100 * time.Millisecond + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + mcpSrv := mcpserver.NewMCPServer("timeout-test", "1.0.0") + mcpSrv.AddTool( + mcp.NewTool("slow_echo", mcp.WithString("input", mcp.Required())), + func(_ context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { + time.Sleep(toolDelay) + args, _ := req.Params.Arguments.(map[string]any) + input, _ := args["input"].(string) + return &mcp.CallToolResult{Content: []mcp.Content{mcp.NewTextContent(input)}}, nil + }, + ) + ts := httptest.NewServer(mcpserver.NewStreamableHTTPServer(mcpSrv)) + t.Cleanup(ts.Close) + + registry := auth.NewDefaultOutgoingAuthRegistry() + require.NoError(t, registry.RegisterStrategy( + authtypes.StrategyTypeUnauthenticated, &strategies.UnauthenticatedStrategy{}, + )) + backendClient, err := NewHTTPBackendClient( + registry, WithRequestTimeoutResolver(tt.resolver), + ) + require.NoError(t, err) + + clientImpl := backendClient.(*httpBackendClient) + clientImpl.revisions.Store("elastic", mcpparser.RevisionLegacy) + target := &vmcp.BackendTarget{ + WorkloadID: "elastic", + WorkloadName: "elastic", + BaseURL: ts.URL, + TransportType: "streamable-http", + } + + result, err := backendClient.CallTool( + t.Context(), target, "slow_echo", map[string]any{"input": "slow response"}, nil, nil, + ) + if tt.wantTimeoutErr { + require.Error(t, err) + assert.Nil(t, result) + return + } + + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, result.Content, 1) + assert.Equal(t, "slow response", result.Content[0].Text) + }) + } +} + // TestTracePropagatingRoundTripper tests the trace context propagation RoundTripper. func TestTracePropagatingRoundTripper(t *testing.T) { t.Parallel() diff --git a/pkg/vmcp/server/server.go b/pkg/vmcp/server/server.go index ba19076d34..9ef9942f29 100644 --- a/pkg/vmcp/server/server.go +++ b/pkg/vmcp/server/server.go @@ -59,8 +59,9 @@ const ( // defaultWriteTimeout is the server-level write deadline set on http.Server.WriteTimeout. // It protects all routes (health, metrics, well-known, etc.) from slow-write clients. - // For qualifying SSE (GET) connections, transportmiddleware.WriteTimeout clears this - // per-request via http.ResponseController.SetWriteDeadline(time.Time{}) (golang/go#16100). + // For MCP POST requests, transportmiddleware.WriteTimeout clears this during + // handler computation and re-arms it when a non-streaming response starts. + // Qualifying SSE connections remain unbounded (golang/go#16100). defaultWriteTimeout = 30 * time.Second // defaultIdleTimeout is the maximum amount of time to wait for the next request when keep-alive's are enabled. @@ -702,12 +703,11 @@ func (s *Server) Handler(_ context.Context) (http.Handler, error) { // Apply Accept header validation (rejects GET requests without Accept: text/event-stream) mcpHandler = headerValidatingMiddleware(mcpHandler) - // Clear the write deadline for qualifying SSE connections (GET + - // Accept: text/event-stream + MCP endpoint path) so the server-level - // WriteTimeout does not kill long-lived SSE streams (see golang/go#16100). - // Non-qualifying requests are left untouched; http.Server.WriteTimeout - // (defaultWriteTimeout) remains in effect for them. - mcpHandler = transportmiddleware.WriteTimeout(s.config.EndpointPath)(mcpHandler) + // Suspend the server write deadline while MCP POST handlers perform bounded + // backend operations, then re-arm it when a non-streaming response starts. + // Qualifying SSE connections remain long-lived (see golang/go#16100). Other + // requests retain the server-level defaultWriteTimeout. + mcpHandler = transportmiddleware.WriteTimeout(s.config.EndpointPath, defaultWriteTimeout)(mcpHandler) // Cap request body size before the MCP parser (and all inner middleware) // buffers it via io.ReadAll, rejecting oversized bodies with 413. This is diff --git a/pkg/vmcp/server/testutil_test.go b/pkg/vmcp/server/testutil_test.go index cb8fce7950..3d953585ea 100644 --- a/pkg/vmcp/server/testutil_test.go +++ b/pkg/vmcp/server/testutil_test.go @@ -10,6 +10,7 @@ import ( "net/http" "net/http/httptest" "testing" + "time" mcpmcp "github.com/stacklok/toolhive-core/mcpcompat/mcp" mcpserver "github.com/stacklok/toolhive-core/mcpcompat/server" @@ -26,6 +27,14 @@ import ( // Returns the full URL to the backend's /mcp endpoint. func startRealMCPBackend(t *testing.T) string { t.Helper() + return startRealMCPBackendWithToolDelay(t, 0) +} + +// startRealMCPBackendWithToolDelay is startRealMCPBackend with a configurable +// delay before the echo tool responds. It is used to exercise vMCP request +// deadlines over a real streamable-HTTP backend connection. +func startRealMCPBackendWithToolDelay(t *testing.T, toolDelay time.Duration) string { + t.Helper() mcpSrv := mcpserver.NewMCPServer("real-backend", "1.0.0") mcpSrv.AddTool( @@ -34,6 +43,7 @@ func startRealMCPBackend(t *testing.T) string { mcpmcp.WithString("input", mcpmcp.Required()), ), func(_ context.Context, req mcpmcp.CallToolRequest) (*mcpmcp.CallToolResult, error) { + time.Sleep(toolDelay) args, _ := req.Params.Arguments.(map[string]any) input, _ := args["input"].(string) return &mcpmcp.CallToolResult{ diff --git a/pkg/vmcp/server/write_timeout_integration_test.go b/pkg/vmcp/server/write_timeout_integration_test.go index 1ca695760c..757ef2e12c 100644 --- a/pkg/vmcp/server/write_timeout_integration_test.go +++ b/pkg/vmcp/server/write_timeout_integration_test.go @@ -16,6 +16,39 @@ import ( "github.com/stretchr/testify/require" ) +// TestIntegration_MCPPostSurvivesWriteTimeout verifies the full vMCP handler +// keeps a tools/call POST alive while a real backend takes longer than the +// server-level WriteTimeout to respond. Backend request timeouts remain the +// authoritative bound for MCP operations. +func TestIntegration_MCPPostSurvivesWriteTimeout(t *testing.T) { + t.Parallel() + + const shortTimeout = 100 * time.Millisecond + const toolDelay = 3 * shortTimeout + + backendURL := startRealMCPBackendWithToolDelay(t, toolDelay) + handler := newRealTestHandler(t, backendURL) + + ts := httptest.NewUnstartedServer(handler) + ts.Config.WriteTimeout = shortTimeout + ts.Start() + t.Cleanup(ts.Close) + + client := NewMCPTestClient(t, ts.URL) + client.InitializeSession() + waitForEchoTool(t, ts.URL, client.SessionID()) + + started := time.Now() + resp := client.CallTool("echo", map[string]any{"input": "slow response"}) + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.Contains(t, string(body), "slow response") + assert.GreaterOrEqual(t, time.Since(started), toolDelay-50*time.Millisecond) +} + // TestIntegration_SSEGetConnectionSurvivesWriteTimeout verifies that the full // vMCP server — with writeTimeoutMiddleware wired in — keeps a qualifying SSE // GET connection alive past the server-level WriteTimeout. diff --git a/pkg/vmcp/session/connector_integration_test.go b/pkg/vmcp/session/connector_integration_test.go index cd93f21c07..5fcf706c0c 100644 --- a/pkg/vmcp/session/connector_integration_test.go +++ b/pkg/vmcp/session/connector_integration_test.go @@ -9,6 +9,7 @@ import ( "net/http/httptest" "sync" "testing" + "time" "github.com/google/uuid" "github.com/stretchr/testify/assert" @@ -35,6 +36,11 @@ import ( // - prompt "greet": returns a greeting message func startInProcessMCPServer(t *testing.T) string { t.Helper() + return startInProcessMCPServerWithToolDelay(t, 0) +} + +func startInProcessMCPServerWithToolDelay(t *testing.T, toolDelay time.Duration) string { + t.Helper() mcpSrv := mcpserver.NewMCPServer("integration-test-backend", "1.0.0") @@ -44,6 +50,7 @@ func startInProcessMCPServer(t *testing.T) string { mcpmcp.WithString("input", mcpmcp.Required()), ), func(_ context.Context, req mcpmcp.CallToolRequest) (*mcpmcp.CallToolResult, error) { + time.Sleep(toolDelay) args, _ := req.Params.Arguments.(map[string]any) input, _ := args["input"].(string) return &mcpmcp.CallToolResult{ @@ -154,6 +161,79 @@ func TestSessionFactory_Integration_CallTool(t *testing.T) { assert.Equal(t, "hello world", result.Content[0].Text) } +// TestSessionFactory_Integration_RequestTimeoutResolver proves that the timeout +// selected for a workload reaches the persistent streamable-HTTP client's tool +// call. The paired cases are intentional: a success-only case would also pass +// if the historical hard-coded 30-second timeout were still in use. +func TestSessionFactory_Integration_RequestTimeoutResolver(t *testing.T) { + t.Parallel() + + const toolDelay = 300 * time.Millisecond + + tests := []struct { + name string + resolver func(string) time.Duration + wantTimeoutErr bool + }{ + { + name: "default timeout is enforced", + resolver: func(string) time.Duration { + return 100 * time.Millisecond + }, + wantTimeoutErr: true, + }, + { + name: "workload override permits slow response", + resolver: func(workloadID string) time.Duration { + if workloadID == "slow-backend" { + return 2 * time.Second + } + return 100 * time.Millisecond + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + baseURL := startInProcessMCPServerWithToolDelay(t, toolDelay) + backend := &vmcp.Backend{ + ID: "slow-backend", + Name: "slow-backend", + BaseURL: baseURL, + TransportType: "streamable-http", + } + + factory := NewSessionFactory( + newUnauthenticatedRegistry(t), + WithRequestTimeoutResolver(tt.resolver), + ) + sess, err := factory.MakeSessionWithID( + t.Context(), uuid.New().String(), nil, []*vmcp.Backend{backend}, nil, + ) + require.NoError(t, err) + // The deliberately tiny timeout in the failure case also applies to + // the best-effort MCP session DELETE performed by Close. + t.Cleanup(func() { _ = sess.Close() }) + + result, err := sess.CallTool( + t.Context(), nil, "echo", map[string]any{"input": "slow response"}, nil, + ) + if tt.wantTimeoutErr { + require.Error(t, err) + assert.Nil(t, result) + return + } + + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, result.Content, 1) + assert.Equal(t, "slow response", result.Content[0].Text) + }) + } +} + func TestSessionFactory_Integration_ReadResource(t *testing.T) { t.Parallel() diff --git a/pkg/vmcp/session/default_session_test.go b/pkg/vmcp/session/default_session_test.go index 5980cfa5b9..b7fd532eb1 100644 --- a/pkg/vmcp/session/default_session_test.go +++ b/pkg/vmcp/session/default_session_test.go @@ -805,6 +805,42 @@ func TestNewSessionFactory_BackendInitTimeout(t *testing.T) { require.NoError(t, sess.Close()) } +func TestNewSessionFactory_WorkloadTimeoutExtendsBackendInit(t *testing.T) { + t.Parallel() + + backend := &vmcp.Backend{ID: "slow", Name: "slow", BaseURL: "http://x:9", TransportType: "streamable-http"} + connector := func(ctx context.Context, _ *vmcp.BackendTarget, _ *auth.Identity, _ string, _ internalbk.ListChangedSink) (internalbk.Session, *vmcp.CapabilityList, error) { + timer := time.NewTimer(150 * time.Millisecond) + defer timer.Stop() + select { + case <-ctx.Done(): + return nil, nil, ctx.Err() + case <-timer.C: + return &mockConnectedBackend{}, &vmcp.CapabilityList{ + Tools: []vmcp.Tool{{Name: "ready"}}, + }, nil + } + } + + factory := newSessionFactoryWithConnector( + connector, + WithBackendInitTimeout(50*time.Millisecond), + WithRequestTimeoutResolver(func(workloadID string) time.Duration { + if workloadID == "slow" { + return time.Second + } + return 50 * time.Millisecond + }), + ) + sess, err := factory.MakeSessionWithID( + context.Background(), uuid.New().String(), nil, []*vmcp.Backend{backend}, nil, + ) + require.NoError(t, err) + require.NotNil(t, sess) + assert.Len(t, sess.Tools(), 1) + require.NoError(t, sess.Close()) +} + func TestNewSessionFactory_ParallelInit(t *testing.T) { t.Parallel() diff --git a/pkg/vmcp/session/factory.go b/pkg/vmcp/session/factory.go index 3a64ea3bd7..cf6604336f 100644 --- a/pkg/vmcp/session/factory.go +++ b/pkg/vmcp/session/factory.go @@ -131,10 +131,11 @@ type backendConnector func( // defaultMultiSessionFactory is the production MultiSessionFactory implementation. type defaultMultiSessionFactory struct { - connector backendConnector - maxConcurrency int - backendInitTimeout time.Duration - revisionLookup func(workloadID string) (mcpparser.Revision, bool) + connector backendConnector + maxConcurrency int + backendInitTimeout time.Duration + revisionLookup func(workloadID string) (mcpparser.Revision, bool) + requestTimeoutResolver func(workloadID string) time.Duration } // MultiSessionFactoryOption configures a defaultMultiSessionFactory. @@ -160,6 +161,23 @@ func WithBackendInitTimeout(d time.Duration) MultiSessionFactoryOption { } } +// WithRequestTimeoutResolver configures the timeout used for individual +// backend operations. The resolver receives a backend workload ID and may +// return a workload-specific duration. A nil resolver, or a non-positive +// result, preserves the historical 30-second default. +// +// The resolver may be called concurrently and must therefore be safe for +// concurrent use. A workload timeout longer than WithBackendInitTimeout also +// extends that workload's initialization deadline; the shorter configured +// value never reduces an explicit initialization allowance. +func WithRequestTimeoutResolver(resolver func(workloadID string) time.Duration) MultiSessionFactoryOption { + return func(f *defaultMultiSessionFactory) { + if resolver != nil { + f.requestTimeoutResolver = resolver + } + } +} + // WithRevisionLookup supplies a function that reports a backend's cached MCP // revision (Legacy vs. Modern) by workload ID, so initOneBackend can skip // connecting to a backend known to speak the stateless Modern (2026-07-28) @@ -189,13 +207,18 @@ func WithRevisionLookup(lookup func(workloadID string) (mcpparser.Revision, bool // NewSessionFactory creates a MultiSessionFactory that connects to backends // over HTTP using the given outgoing auth registry. func NewSessionFactory(registry vmcpauth.OutgoingAuthRegistry, opts ...MultiSessionFactoryOption) MultiSessionFactory { - return newSessionFactoryWithConnector(backend.NewHTTPConnector(registry), opts...) + f := newSessionFactoryWithConnector(nil, opts...) + f.connector = backend.NewHTTPConnector( + registry, + backend.WithRequestTimeoutResolver(f.requestTimeoutResolver), + ) + return f } // newSessionFactoryWithConnector creates a MultiSessionFactory backed by an // arbitrary connector. Used by tests to inject a fake connector without // requiring real HTTP backends. -func newSessionFactoryWithConnector(connector backendConnector, opts ...MultiSessionFactoryOption) MultiSessionFactory { +func newSessionFactoryWithConnector(connector backendConnector, opts ...MultiSessionFactoryOption) *defaultMultiSessionFactory { f := &defaultMultiSessionFactory{ connector: connector, maxConcurrency: defaultMaxBackendInitConcurrency, @@ -286,7 +309,13 @@ func (f *defaultMultiSessionFactory) initOneBackend( return nil, true } - bCtx, cancel := context.WithTimeout(ctx, f.backendInitTimeout) + initTimeout := f.backendInitTimeout + if f.requestTimeoutResolver != nil { + if requestTimeout := f.requestTimeoutResolver(target.WorkloadID); requestTimeout > initTimeout { + initTimeout = requestTimeout + } + } + bCtx, cancel := context.WithTimeout(ctx, initTimeout) defer cancel() conn, caps, err := f.connector(bCtx, target, identity, sessionHint, sink) diff --git a/pkg/vmcp/session/internal/backend/mcp_session.go b/pkg/vmcp/session/internal/backend/mcp_session.go index c8affcdb01..0f3ab72885 100644 --- a/pkg/vmcp/session/internal/backend/mcp_session.go +++ b/pkg/vmcp/session/internal/backend/mcp_session.go @@ -36,11 +36,43 @@ const ( maxBackendResponseSize = 100 * 1024 * 1024 // 100 MB // defaultBackendRequestTimeout is the wall-clock deadline for individual - // streamable-HTTP requests. Applied at both the http.Client and SDK layers - // (defense-in-depth). Not used for SSE, whose stream lifetime is unbounded. + // backend operations. Streamable-HTTP also applies it at the http.Client and + // SDK layers; SSE keeps its connection unbounded but bounds operation contexts. defaultBackendRequestTimeout = 30 * time.Second ) +// HTTPConnectorOption configures the persistent HTTP backend connector. +type HTTPConnectorOption func(*httpConnectorConfig) + +type httpConnectorConfig struct { + requestTimeoutResolver func(workloadID string) time.Duration +} + +// WithRequestTimeoutResolver configures the timeout used for each backend +// operation. The resolver receives the backend workload ID and may return a +// workload-specific duration. A nil resolver, or a non-positive result, uses +// the 30-second default. +// +// The resolver may be called concurrently and must therefore be safe for +// concurrent use. SSE connection lifetimes remain unbounded, but individual +// operations on those connections are bounded through request contexts. +func WithRequestTimeoutResolver(resolver func(workloadID string) time.Duration) HTTPConnectorOption { + return func(cfg *httpConnectorConfig) { + if resolver != nil { + cfg.requestTimeoutResolver = resolver + } + } +} + +func (c *httpConnectorConfig) requestTimeout(workloadID string) time.Duration { + if c.requestTimeoutResolver != nil { + if timeout := c.requestTimeoutResolver(workloadID); timeout > 0 { + return timeout + } + } + return defaultBackendRequestTimeout +} + // ChangeKind identifies which capability class a backend reported changed via // ListChangedSink. Using a typed constant (rather than a bare string) means a // typo is a compile error at the producer and consumer instead of a silent @@ -183,6 +215,7 @@ type mcpSession struct { client *mcpclient.Client target *vmcp.BackendTarget // bound at creation; used for capability name translation backendSessionID string // backend-assigned session ID (may be empty) + requestTimeout time.Duration // bounds each backend operation, including over SSE } // SessionID returns the backend-assigned session ID. @@ -198,6 +231,9 @@ func (c *mcpSession) CallTool( arguments map[string]any, meta map[string]any, ) (*vmcp.ToolCallResult, error) { + ctx, cancel := context.WithTimeout(ctx, c.requestTimeout) + defer cancel() + backendName := c.target.GetBackendCapabilityName(toolName) if backendName != toolName { slog.Debug("Translating tool name", "clientName", toolName, "backendName", backendName) @@ -246,6 +282,9 @@ func (c *mcpSession) ReadResource( ctx context.Context, uri string, ) (*vmcp.ResourceReadResult, error) { + ctx, cancel := context.WithTimeout(ctx, c.requestTimeout) + defer cancel() + backendURI := c.target.GetBackendCapabilityName(uri) if backendURI != uri { slog.Debug("Translating resource URI", "clientURI", uri, "backendURI", backendURI) @@ -276,6 +315,9 @@ func (c *mcpSession) GetPrompt( name string, arguments map[string]any, ) (*vmcp.PromptGetResult, error) { + ctx, cancel := context.WithTimeout(ctx, c.requestTimeout) + defer cancel() + backendName := c.target.GetBackendCapabilityName(name) if backendName != name { slog.Debug("Translating prompt name", "clientName", name, "backendName", backendName) @@ -320,7 +362,7 @@ func (c *mcpSession) GetPrompt( // createMCPClient for what that does and does not enable (nil-sink callers are // completely unaffected: no OnNotification handler is registered and no // standalone GET stream is opened). -func NewHTTPConnector(registry vmcpauth.OutgoingAuthRegistry) func( +func NewHTTPConnector(registry vmcpauth.OutgoingAuthRegistry, opts ...HTTPConnectorOption) func( ctx context.Context, target *vmcp.BackendTarget, identity *auth.Identity, @@ -328,6 +370,10 @@ func NewHTTPConnector(registry vmcpauth.OutgoingAuthRegistry) func( sink ListChangedSink, ) (Session, *vmcp.CapabilityList, error) { provider := secrets.NewEnvironmentProvider() + connectorConfig := &httpConnectorConfig{} + for _, opt := range opts { + opt(connectorConfig) + } return func( ctx context.Context, target *vmcp.BackendTarget, @@ -335,12 +381,18 @@ func NewHTTPConnector(registry vmcpauth.OutgoingAuthRegistry) func( sessionHint string, sink ListChangedSink, ) (Session, *vmcp.CapabilityList, error) { - c, err := createMCPClient(ctx, target, identity, registry, sessionHint, provider, sink) + requestTimeout := connectorConfig.requestTimeout(target.WorkloadID) + requestCtx, cancel := context.WithTimeout(ctx, requestTimeout) + defer cancel() + + c, err := createMCPClient( + requestCtx, target, identity, registry, sessionHint, provider, sink, requestTimeout, + ) if err != nil { return nil, nil, fmt.Errorf("failed to create MCP client for backend %s: %w", target.WorkloadID, err) } - caps, err := initAndQueryCapabilities(ctx, c, target) + caps, err := initAndQueryCapabilities(requestCtx, c, target) if err != nil { _ = c.Close() return nil, nil, fmt.Errorf("failed to initialise backend %s: %w", target.WorkloadID, err) @@ -356,7 +408,9 @@ func NewHTTPConnector(registry vmcpauth.OutgoingAuthRegistry) func( backendSessionID = sh.GetSessionId() } - return &mcpSession{client: c, target: target, backendSessionID: backendSessionID}, caps, nil + return &mcpSession{ + client: c, target: target, backendSessionID: backendSessionID, requestTimeout: requestTimeout, + }, caps, nil } } @@ -390,6 +444,7 @@ func createMCPClient( sessionHint string, provider secrets.Provider, sink ListChangedSink, + requestTimeout time.Duration, ) (*mcpclient.Client, error) { // Resolve and validate the auth strategy once at client creation time. strategyName := authtypes.StrategyTypeUnauthenticated @@ -453,7 +508,7 @@ func createMCPClient( // WithHTTPTimeout additionally wraps each SDK request in a // context.WithTimeout so the mcpcompat transport surfaces a descriptive // error before the stdlib deadline fires. Both are set to - // defaultBackendRequestTimeout: defense-in-depth. + // requestTimeout: defense-in-depth. sizeLimited := httpRoundTripperFunc(func(req *http.Request) (*http.Response, error) { resp, err := base.RoundTrip(req) if err != nil { @@ -473,11 +528,11 @@ func createMCPClient( // auth/identity/header-forward credentials at the attacker host. httpClient := &http.Client{ Transport: sizeLimited, - Timeout: defaultBackendRequestTimeout, + Timeout: requestTimeout, CheckRedirect: networking.SameHostRedirectPolicy(), } streamableOpts := []mcptransport.StreamableHTTPCOption{ - mcptransport.WithHTTPTimeout(defaultBackendRequestTimeout), + mcptransport.WithHTTPTimeout(requestTimeout), mcptransport.WithHTTPBasicClient(httpClient), } if sessionHint != "" { diff --git a/pkg/vmcp/session/internal/backend/mcp_session_test.go b/pkg/vmcp/session/internal/backend/mcp_session_test.go index 9d4a3a01bd..d7fc602b58 100644 --- a/pkg/vmcp/session/internal/backend/mcp_session_test.go +++ b/pkg/vmcp/session/internal/backend/mcp_session_test.go @@ -112,7 +112,10 @@ func TestCreateMCPClient_UnsupportedTransport(t *testing.T) { TransportType: transport, } - _, err := createMCPClient(context.Background(), target, nil, newTestRegistry(t), "", secrets.NewEnvironmentProvider(), nil) + _, err := createMCPClient( + context.Background(), target, nil, newTestRegistry(t), "", secrets.NewEnvironmentProvider(), nil, + defaultBackendRequestTimeout, + ) require.Error(t, err) assert.ErrorIs(t, err, vmcp.ErrUnsupportedTransport, "transport %q should return ErrUnsupportedTransport", transport) @@ -372,6 +375,7 @@ func TestCreateMCPClient_ContinuousListeningGatedOnSink(t *testing.T) { c, err := createMCPClient( context.Background(), target, nil, newTestRegistry(t), "", secrets.NewEnvironmentProvider(), tc.sink, + defaultBackendRequestTimeout, ) require.NoError(t, err) t.Cleanup(func() { _ = c.Close() }) @@ -439,6 +443,7 @@ func TestCreateMCPClient_ListChangedSink_FiresOnBackendNotification(t *testing.T c, err := createMCPClient( context.Background(), target, nil, newTestRegistry(t), "", secrets.NewEnvironmentProvider(), sink, + defaultBackendRequestTimeout, ) require.NoError(t, err) t.Cleanup(func() { _ = c.Close() }) @@ -510,6 +515,7 @@ func TestCreateMCPClient_ListChangedSink_DoesNotStallInFlightCall(t *testing.T) c, err := createMCPClient( context.Background(), target, nil, newTestRegistry(t), "", secrets.NewEnvironmentProvider(), sink, + defaultBackendRequestTimeout, ) require.NoError(t, err) t.Cleanup(func() { _ = c.Close() }) From 4a6522048926d258d7b9ac440db4347630d83b68 Mon Sep 17 00:00:00 2001 From: Jairus Christensen Date: Fri, 21 Aug 2026 14:02:28 -0600 Subject: [PATCH 2/3] Preserve backend initialization timeout Use the factory's widened initialization context when connecting persistent backends, while retaining the configured timeout for ordinary session operations. Cover the shorter request-timeout case with the real HTTP connector. Signed-off-by: Jairus Christensen --- .../session/connector_integration_test.go | 53 ++++++++++++++++++- .../session/internal/backend/mcp_session.go | 12 +++-- 2 files changed, 60 insertions(+), 5 deletions(-) diff --git a/pkg/vmcp/session/connector_integration_test.go b/pkg/vmcp/session/connector_integration_test.go index 5fcf706c0c..20482597c1 100644 --- a/pkg/vmcp/session/connector_integration_test.go +++ b/pkg/vmcp/session/connector_integration_test.go @@ -41,6 +41,11 @@ func startInProcessMCPServer(t *testing.T) string { func startInProcessMCPServerWithToolDelay(t *testing.T, toolDelay time.Duration) string { t.Helper() + return startInProcessMCPServerWithDelays(t, 0, toolDelay) +} + +func startInProcessMCPServerWithDelays(t *testing.T, initDelay, toolDelay time.Duration) string { + t.Helper() mcpSrv := mcpserver.NewMCPServer("integration-test-backend", "1.0.0") @@ -87,7 +92,13 @@ func startInProcessMCPServerWithToolDelay(t *testing.T, toolDelay time.Duration) streamableSrv := mcpserver.NewStreamableHTTPServer(mcpSrv) mux := http.NewServeMux() - mux.Handle("/mcp", streamableSrv) + var initOnce sync.Once + mux.Handle("/mcp", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPost { + initOnce.Do(func() { time.Sleep(initDelay) }) + } + streamableSrv.ServeHTTP(w, r) + })) ts := httptest.NewServer(mux) t.Cleanup(ts.Close) @@ -234,6 +245,46 @@ func TestSessionFactory_Integration_RequestTimeoutResolver(t *testing.T) { } } +func TestSessionFactory_Integration_ShortRequestTimeoutDoesNotShrinkInit(t *testing.T) { + t.Parallel() + + const ( + requestTimeout = 50 * time.Millisecond + initTimeout = 2 * time.Second + backendDelay = 200 * time.Millisecond + ) + + baseURL := startInProcessMCPServerWithDelays(t, backendDelay, backendDelay) + backend := &vmcp.Backend{ + ID: "slow-init-backend", + Name: "slow-init-backend", + BaseURL: baseURL, + TransportType: "streamable-http", + } + + factory := NewSessionFactory( + newUnauthenticatedRegistry(t), + WithBackendInitTimeout(initTimeout), + WithRequestTimeoutResolver(func(string) time.Duration { return requestTimeout }), + ) + + started := time.Now() + sess, err := factory.MakeSessionWithID( + t.Context(), uuid.New().String(), nil, []*vmcp.Backend{backend}, nil, + ) + require.NoError(t, err) + require.GreaterOrEqual(t, time.Since(started), backendDelay) + require.Len(t, sess.Tools(), 1) + t.Cleanup(func() { _ = sess.Close() }) + + result, err := sess.CallTool( + t.Context(), nil, "echo", map[string]any{"input": "slow response"}, nil, + ) + require.Error(t, err) + assert.ErrorContains(t, err, "context deadline exceeded") + assert.Nil(t, result) +} + func TestSessionFactory_Integration_ReadResource(t *testing.T) { t.Parallel() diff --git a/pkg/vmcp/session/internal/backend/mcp_session.go b/pkg/vmcp/session/internal/backend/mcp_session.go index 0f3ab72885..34421c80ce 100644 --- a/pkg/vmcp/session/internal/backend/mcp_session.go +++ b/pkg/vmcp/session/internal/backend/mcp_session.go @@ -382,17 +382,21 @@ func NewHTTPConnector(registry vmcpauth.OutgoingAuthRegistry, opts ...HTTPConnec sink ListChangedSink, ) (Session, *vmcp.CapabilityList, error) { requestTimeout := connectorConfig.requestTimeout(target.WorkloadID) - requestCtx, cancel := context.WithTimeout(ctx, requestTimeout) - defer cancel() + transportTimeout := requestTimeout + if deadline, ok := ctx.Deadline(); ok { + if remaining := time.Until(deadline); remaining > transportTimeout { + transportTimeout = remaining + } + } c, err := createMCPClient( - requestCtx, target, identity, registry, sessionHint, provider, sink, requestTimeout, + ctx, target, identity, registry, sessionHint, provider, sink, transportTimeout, ) if err != nil { return nil, nil, fmt.Errorf("failed to create MCP client for backend %s: %w", target.WorkloadID, err) } - caps, err := initAndQueryCapabilities(requestCtx, c, target) + caps, err := initAndQueryCapabilities(ctx, c, target) if err != nil { _ = c.Close() return nil, nil, fmt.Errorf("failed to initialise backend %s: %w", target.WorkloadID, err) From 609316391dfda0d57963fa9edd710e12fca7bb2c Mon Sep 17 00:00:00 2001 From: Jairus Christensen Date: Fri, 21 Aug 2026 14:16:21 -0600 Subject: [PATCH 3/3] Test default backend initialization allowance Signed-off-by: Jairus Christensen --- pkg/vmcp/session/connector_integration_test.go | 4 ---- 1 file changed, 4 deletions(-) diff --git a/pkg/vmcp/session/connector_integration_test.go b/pkg/vmcp/session/connector_integration_test.go index 20482597c1..42821fafdd 100644 --- a/pkg/vmcp/session/connector_integration_test.go +++ b/pkg/vmcp/session/connector_integration_test.go @@ -224,8 +224,6 @@ func TestSessionFactory_Integration_RequestTimeoutResolver(t *testing.T) { t.Context(), uuid.New().String(), nil, []*vmcp.Backend{backend}, nil, ) require.NoError(t, err) - // The deliberately tiny timeout in the failure case also applies to - // the best-effort MCP session DELETE performed by Close. t.Cleanup(func() { _ = sess.Close() }) result, err := sess.CallTool( @@ -250,7 +248,6 @@ func TestSessionFactory_Integration_ShortRequestTimeoutDoesNotShrinkInit(t *test const ( requestTimeout = 50 * time.Millisecond - initTimeout = 2 * time.Second backendDelay = 200 * time.Millisecond ) @@ -264,7 +261,6 @@ func TestSessionFactory_Integration_ShortRequestTimeoutDoesNotShrinkInit(t *test factory := NewSessionFactory( newUnauthenticatedRegistry(t), - WithBackendInitTimeout(initTimeout), WithRequestTimeoutResolver(func(string) time.Duration { return requestTimeout }), )