mirror of
https://github.com/mudler/LocalAI.git
synced 2026-10-09 22:54:42 -04:00
fix(mcp): connect to legacy SSE servers (#12304)
* fix(mcp): connect to legacy SSE servers Remote model MCP connections only use Streamable HTTP, so legacy SSE servers fail initialization. Retry with SSE when the initial POST returns 400, 404, or 405. Share one discovery timeout across attempts and cancel failed connections without truncating successful sessions. Preserve HTTP policy and reject foreign SSE message endpoints before attaching credentials. Add SDK integration tests and document automatic transport selection. Assisted-by: Codex:gpt-6 Signed-off-by: Ettore Di Giacinto <mudler@localai.io> * fix(mcp): avoid narrowing HTTP status codes Store the initialization status in a 64-bit atomic value to avoid the integer overflow conversion reported by gosec. Assisted-by: Codex:GPT-6 gosec Signed-off-by: Ettore Di Giacinto <mudler@localai.io> --------- Signed-off-by: Ettore Di Giacinto <mudler@localai.io> Co-authored-by: localai-org-maint-bot <306269227+localai-org-maint-bot@users.noreply.github.com>
This commit is contained in:
1 parent
4a1089be35
commit
0102751b31
5 files changed
+511
-4
No files matched your search
@@ -19,6 +19,7 @@ import (
|
||||
|
||||
"github.com/mudler/LocalAI/pkg/functions"
|
||||
"github.com/mudler/LocalAI/pkg/httpclient"
|
||||
"github.com/mudler/LocalAI/pkg/mcptransport"
|
||||
"github.com/mudler/LocalAI/pkg/signals"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
@@ -234,8 +235,7 @@ func SessionsFromMCPConfig(
|
||||
httpclient.WithTransport(newBearerTokenRoundTripper(server.Token, httpclient.HardenedTransport())),
|
||||
)
|
||||
|
||||
transport := &mcp.StreamableClientTransport{Endpoint: server.URL, HTTPClient: httpClient}
|
||||
mcpSession, err := connectMCP(ctx, transport, config.DefaultMCPDiscoveryTimeout)
|
||||
mcpSession, err := mcptransport.Connect(ctx, client, server.URL, httpClient, config.DefaultMCPDiscoveryTimeout)
|
||||
if err != nil {
|
||||
xlog.Error("Failed to connect to MCP server", "error", err, "url", server.URL)
|
||||
continue
|
||||
@@ -344,8 +344,7 @@ func NamedSessionsFromMCPConfig(
|
||||
httpclient.WithTransport(newBearerTokenRoundTripper(server.Token, httpclient.HardenedTransport())),
|
||||
)
|
||||
|
||||
transport := &mcp.StreamableClientTransport{Endpoint: server.URL, HTTPClient: httpClient}
|
||||
mcpSession, err := connectMCP(ctx, transport, config.DefaultMCPDiscoveryTimeout)
|
||||
mcpSession, err := mcptransport.Connect(ctx, client, server.URL, httpClient, config.DefaultMCPDiscoveryTimeout)
|
||||
if err != nil {
|
||||
xlog.Error("Failed to connect to MCP server", "error", err, "name", serverName, "url", server.URL)
|
||||
allSessions = append(allSessions, NamedSession{
|
||||
|
||||
@@ -95,6 +95,10 @@ Configure HTTP-based MCP servers:
|
||||
- **`url`**: The MCP server endpoint URL
|
||||
- **`token`**: Bearer token for authentication (optional)
|
||||
|
||||
LocalAI automatically selects the transport for remote model MCP servers. It tries Streamable HTTP first, including servers that do not assign session IDs. If the initial POST returns HTTP 400, 404, or 405, LocalAI retries with legacy SSE. Both attempts share the discovery timeout and use the configured bearer token. Authentication failures and redirects do not trigger fallback.
|
||||
|
||||
Use the endpoint URL published by your server: usually `/mcp` for Streamable HTTP or `/sse` for legacy SSE. Legacy SSE servers must advertise a message endpoint on the same origin (scheme, host, and port). No transport setting is required.
|
||||
|
||||
Remote model MCP connections originate from the LocalAI process. If LocalAI runs in Docker, the URL must therefore resolve and be reachable **from the LocalAI container**, not only from the host browser. For another service in the same Compose project, use its Compose service name and container port. Host-only DNS names, VPN DNS, and private routes must also be made available inside the container.
|
||||
|
||||
#### STDIO Servers (`stdio`)
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package mcptransport
|
||||
|
||||
import (
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestMCPTransport(t *testing.T) { RegisterFailHandler(Fail); RunSpecs(t, "Remote MCP transport") }
|
||||
@@ -0,0 +1,193 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
// Package mcptransport connects remote MCP servers using Streamable HTTP or legacy SSE.
|
||||
package mcptransport
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
"github.com/mudler/LocalAI/pkg/httpclient"
|
||||
)
|
||||
|
||||
// Connect prefers Streamable HTTP and retries with legacy SSE only when the
|
||||
// initial POST is rejected with 400, 404, or 405. The timeout covers both
|
||||
// handshakes; an established session lives until it is closed or ctx is canceled.
|
||||
func Connect(ctx context.Context, client *mcp.Client, endpoint string, httpClient *http.Client, timeout time.Duration) (*mcp.ClientSession, error) {
|
||||
deadline, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
if err := deadline.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if httpClient == nil {
|
||||
httpClient = httpclient.New()
|
||||
}
|
||||
// Copy the client so its credentials, redirect policy and other settings are
|
||||
// retained without mutating a client that other sessions may share.
|
||||
streamClient := *httpClient
|
||||
base := streamClient.Transport
|
||||
if base == nil {
|
||||
base = httpclient.HardenedTransport()
|
||||
}
|
||||
status := &initializeStatus{base: base}
|
||||
streamClient.Transport = status
|
||||
session, err := connect(ctx, deadline, client, &mcp.StreamableClientTransport{Endpoint: endpoint, HTTPClient: &streamClient}, &streamClient)
|
||||
if err == nil {
|
||||
return session, nil
|
||||
}
|
||||
if deadline.Err() != nil {
|
||||
return nil, deadline.Err()
|
||||
}
|
||||
switch status.code.Load() {
|
||||
case http.StatusBadRequest, http.StatusNotFound, http.StatusMethodNotAllowed:
|
||||
origin, parseErr := url.Parse(endpoint)
|
||||
if parseErr != nil {
|
||||
return nil, parseErr
|
||||
}
|
||||
sseClient := *httpClient
|
||||
// SSE servers advertise a separate POST endpoint. Reject foreign origins
|
||||
// before a credential-injecting transport can see the request.
|
||||
sseClient.Transport = &sameOriginTransport{base: base, origin: origin}
|
||||
return connect(ctx, deadline, client, &mcp.SSEClientTransport{Endpoint: endpoint, HTTPClient: &sseClient}, &sseClient)
|
||||
default:
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// The SDK sends initialize as the first POST. Keep its status, rather than
|
||||
// allowing the initialized notification or a background GET to overwrite it.
|
||||
type initializeStatus struct {
|
||||
base http.RoundTripper
|
||||
seen atomic.Bool
|
||||
code atomic.Int64
|
||||
}
|
||||
|
||||
func (t *initializeStatus) RoundTrip(r *http.Request) (*http.Response, error) {
|
||||
initial := r.Method == http.MethodPost && t.seen.CompareAndSwap(false, true)
|
||||
response, err := t.base.RoundTrip(r)
|
||||
if initial && err == nil {
|
||||
t.code.Store(int64(response.StatusCode))
|
||||
}
|
||||
return response, err
|
||||
}
|
||||
|
||||
type sameOriginTransport struct {
|
||||
base http.RoundTripper
|
||||
origin *url.URL
|
||||
}
|
||||
|
||||
func (t *sameOriginTransport) RoundTrip(r *http.Request) (*http.Response, error) {
|
||||
if !strings.EqualFold(r.URL.Scheme, t.origin.Scheme) || !strings.EqualFold(r.URL.Host, t.origin.Host) {
|
||||
return nil, errors.New("MCP SSE endpoint must use the configured server's origin")
|
||||
}
|
||||
return t.base.RoundTrip(r)
|
||||
}
|
||||
|
||||
// Keep the SDK's concrete connection intact: its private sessionUpdated method
|
||||
// is needed to set Streamable HTTP protocol headers and start its event stream.
|
||||
type attemptTransport struct {
|
||||
mcp.Transport
|
||||
cleanup func()
|
||||
}
|
||||
|
||||
func (t *attemptTransport) Connect(ctx context.Context) (mcp.Connection, error) {
|
||||
connection, err := t.Transport.Connect(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stop := context.AfterFunc(ctx, func() { _ = connection.Close() })
|
||||
t.cleanup = func() { stop(); _ = connection.Close() }
|
||||
return connection, nil
|
||||
}
|
||||
|
||||
func connect(ctx, deadline context.Context, client *mcp.Client, transport mcp.Transport, httpClient *http.Client) (*mcp.ClientSession, error) {
|
||||
if err := deadline.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
attemptCtx, cancel := context.WithCancel(ctx)
|
||||
httpClient.Transport = &attemptRequests{base: httpClient.Transport, ctx: attemptCtx}
|
||||
attempt := &attemptTransport{Transport: transport}
|
||||
type result struct {
|
||||
session *mcp.ClientSession
|
||||
err error
|
||||
}
|
||||
done := make(chan result)
|
||||
go func() {
|
||||
session, err := client.Connect(attemptCtx, attempt, nil)
|
||||
cleanup := func() {
|
||||
cancel()
|
||||
if attempt.cleanup != nil {
|
||||
attempt.cleanup()
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
cleanup()
|
||||
}
|
||||
select {
|
||||
case done <- result{session, err}:
|
||||
if session != nil {
|
||||
// Release cancellation hooks when callers close a successful session.
|
||||
_ = session.Wait()
|
||||
cleanup()
|
||||
}
|
||||
case <-deadline.Done():
|
||||
// An unbuffered handoff prevents a late successful session being orphaned.
|
||||
cleanup()
|
||||
if session != nil {
|
||||
_ = session.Close()
|
||||
}
|
||||
}
|
||||
}()
|
||||
select {
|
||||
case result := <-done:
|
||||
if err := deadline.Err(); err != nil {
|
||||
cancel()
|
||||
return nil, err
|
||||
}
|
||||
return result.session, result.err
|
||||
case <-deadline.Done():
|
||||
cancel()
|
||||
return nil, deadline.Err()
|
||||
}
|
||||
}
|
||||
|
||||
// The SDK detaches cleanup requests and cancellation notifications from the
|
||||
// handshake context. Bind those requests back to the attempt lifetime too, so
|
||||
// an unreachable server cannot keep a failed handshake goroutine alive.
|
||||
type attemptRequests struct {
|
||||
base http.RoundTripper
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
func (t *attemptRequests) RoundTrip(r *http.Request) (*http.Response, error) {
|
||||
if err := t.ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ctx, cancel := context.WithCancel(r.Context())
|
||||
stop := context.AfterFunc(t.ctx, cancel)
|
||||
response, err := t.base.RoundTrip(r.Clone(ctx))
|
||||
cleanup := func() { stop(); cancel() }
|
||||
if err != nil {
|
||||
cleanup()
|
||||
return response, err
|
||||
}
|
||||
response.Body = &attemptBody{ReadCloser: response.Body, cleanup: cleanup}
|
||||
return response, nil
|
||||
}
|
||||
|
||||
type attemptBody struct {
|
||||
io.ReadCloser
|
||||
cleanup func()
|
||||
}
|
||||
|
||||
func (b *attemptBody) Close() error {
|
||||
defer b.cleanup()
|
||||
return b.ReadCloser.Close()
|
||||
}
|
||||
@@ -0,0 +1,301 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package mcptransport
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
"github.com/mudler/LocalAI/pkg/httpclient"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
type roundTripFunc func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) }
|
||||
|
||||
func testServer() *mcp.Server {
|
||||
server := mcp.NewServer(&mcp.Implementation{Name: "test", Version: "1"}, nil)
|
||||
mcp.AddTool(server, &mcp.Tool{Name: "echo", Description: "Echo text"}, func(_ context.Context, _ *mcp.CallToolRequest, args struct {
|
||||
Text string `json:"text"`
|
||||
}) (*mcp.CallToolResult, any, error) {
|
||||
return &mcp.CallToolResult{Content: []mcp.Content{&mcp.TextContent{Text: args.Text}}}, nil, nil
|
||||
})
|
||||
return server
|
||||
}
|
||||
func testClient() *mcp.Client {
|
||||
return mcp.NewClient(&mcp.Implementation{Name: "test", Version: "1"}, nil)
|
||||
}
|
||||
func checkTool(session *mcp.ClientSession) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
tools, err := session.ListTools(ctx, nil)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(tools.Tools).To(HaveLen(1))
|
||||
result, err := session.CallTool(ctx, &mcp.CallToolParams{Name: "echo", Arguments: map[string]any{"text": "hello"}})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(result.Content).To(Equal([]mcp.Content{&mcp.TextContent{Text: "hello"}}))
|
||||
}
|
||||
|
||||
var _ = Describe("Remote MCP transport selection", func() {
|
||||
DescribeTable("connects to Streamable HTTP without a probe session", func(stateless bool) {
|
||||
var initializations atomic.Int32
|
||||
handler := mcp.NewStreamableHTTPHandler(func(*http.Request) *mcp.Server { return testServer() }, &mcp.StreamableHTTPOptions{Stateless: stateless})
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodPost && r.Header.Get("Mcp-Protocol-Version") == "" {
|
||||
initializations.Add(1)
|
||||
}
|
||||
handler.ServeHTTP(w, r)
|
||||
}))
|
||||
defer server.Close()
|
||||
session, err := Connect(context.Background(), testClient(), server.URL, httpclient.New(), time.Second)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
defer func() { _ = session.Close() }()
|
||||
checkTool(session)
|
||||
Expect(initializations.Load()).To(Equal(int32(1)))
|
||||
}, Entry("stateful", false), Entry("stateless", true))
|
||||
|
||||
DescribeTable("falls back for a rejected initialize POST and preserves bearer credentials", func(status int) {
|
||||
var gets, posts, rejected atomic.Int32
|
||||
handler := mcp.NewSSEHandler(func(*http.Request) *mcp.Server { return testServer() }, nil)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("Authorization") != "Bearer secret" {
|
||||
rejected.Add(1)
|
||||
http.Error(w, "unauthorized", 401)
|
||||
return
|
||||
}
|
||||
if r.Method == http.MethodPost && r.URL.RawQuery == "" {
|
||||
http.Error(w, "use SSE", status)
|
||||
return
|
||||
}
|
||||
if r.Method == http.MethodGet {
|
||||
gets.Add(1)
|
||||
// The SDK advertises this same-origin path for subsequent messages.
|
||||
r.URL.Path = "/messages"
|
||||
} else {
|
||||
if r.URL.Path != "/messages" {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
posts.Add(1)
|
||||
}
|
||||
handler.ServeHTTP(w, r)
|
||||
}))
|
||||
defer server.Close()
|
||||
base := httpclient.HardenedTransport()
|
||||
hc := httpclient.New(httpclient.WithTransport(roundTripFunc(func(r *http.Request) (*http.Response, error) {
|
||||
clone := r.Clone(r.Context())
|
||||
clone.Header.Set("Authorization", "Bearer secret")
|
||||
return base.RoundTrip(clone)
|
||||
})))
|
||||
session, err := Connect(context.Background(), testClient(), server.URL, hc, time.Second)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
defer func() { _ = session.Close() }()
|
||||
checkTool(session)
|
||||
Expect(gets.Load()).To(Equal(int32(1)))
|
||||
Expect(posts.Load()).To(BeNumerically(">=", 4))
|
||||
Expect(rejected.Load()).To(BeZero())
|
||||
}, Entry("400", 400), Entry("404", 404), Entry("405", 405))
|
||||
|
||||
DescribeTable("does not fall back on other HTTP errors or follow redirects", func(status int) {
|
||||
var gets, redirected atomic.Int32
|
||||
target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { redirected.Add(1) }))
|
||||
defer target.Close()
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodGet {
|
||||
gets.Add(1)
|
||||
}
|
||||
w.Header().Set("Location", target.URL)
|
||||
http.Error(w, "failure", status)
|
||||
}))
|
||||
defer server.Close()
|
||||
session, err := Connect(context.Background(), testClient(), server.URL, httpclient.New(), time.Second)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(session).To(BeNil())
|
||||
Expect(gets.Load()).To(BeZero())
|
||||
Expect(redirected.Load()).To(BeZero())
|
||||
}, Entry("401", 401), Entry("403", 403), Entry("500", 500), Entry("302", 302), Entry("307", 307))
|
||||
|
||||
It("does not mistake a rejected initialized notification for rejected initialization", func() {
|
||||
var gets atomic.Int32
|
||||
handler := mcp.NewStreamableHTTPHandler(func(*http.Request) *mcp.Server { return testServer() }, &mcp.StreamableHTTPOptions{Stateless: true})
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodGet {
|
||||
gets.Add(1)
|
||||
}
|
||||
if r.Method == http.MethodPost && r.Header.Get("Mcp-Protocol-Version") != "" {
|
||||
http.Error(w, "rejected notification", 405)
|
||||
return
|
||||
}
|
||||
handler.ServeHTTP(w, r)
|
||||
}))
|
||||
defer server.Close()
|
||||
session, err := Connect(context.Background(), testClient(), server.URL, httpclient.New(), time.Second)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(session).To(BeNil())
|
||||
Expect(gets.Load()).To(Equal(int32(1)))
|
||||
})
|
||||
|
||||
DescribeTable("keeps a connected session alive beyond the discovery timeout", func(legacy bool) {
|
||||
var handler http.Handler = mcp.NewStreamableHTTPHandler(func(*http.Request) *mcp.Server { return testServer() }, nil)
|
||||
if legacy {
|
||||
handler = mcp.NewSSEHandler(func(*http.Request) *mcp.Server { return testServer() }, nil)
|
||||
}
|
||||
server := httptest.NewServer(handler)
|
||||
defer server.Close()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
session, err := Connect(ctx, testClient(), server.URL, httpclient.New(), 100*time.Millisecond)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
defer func() { _ = session.Close() }()
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
checkTool(session)
|
||||
cancel()
|
||||
done := make(chan struct{})
|
||||
go func() { _ = session.Wait(); close(done) }()
|
||||
Eventually(done).WithTimeout(time.Second).Should(BeClosed())
|
||||
}, Entry("Streamable HTTP", false), Entry("legacy SSE", true))
|
||||
|
||||
DescribeTable("bounds stalled setup and cancels its HTTP requests", func(legacy bool, cancelParent bool) {
|
||||
entered := make(chan struct{}, 1)
|
||||
canceled := make(chan struct{}, 8)
|
||||
release := make(chan struct{})
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if legacy && r.Method == http.MethodPost {
|
||||
http.Error(w, "use SSE", 405)
|
||||
return
|
||||
}
|
||||
// Consume initialize so the server can observe the peer disconnecting.
|
||||
_, _ = io.Copy(io.Discard, r.Body)
|
||||
select {
|
||||
case entered <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
select {
|
||||
case <-r.Context().Done():
|
||||
canceled <- struct{}{}
|
||||
case <-release:
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
defer close(release)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
if cancelParent {
|
||||
go func() { <-entered; cancel() }()
|
||||
}
|
||||
start := time.Now()
|
||||
timeout := 100 * time.Millisecond
|
||||
if cancelParent {
|
||||
timeout = 5 * time.Second
|
||||
}
|
||||
session, err := Connect(ctx, testClient(), server.URL, httpclient.New(), timeout)
|
||||
Expect(session).To(BeNil())
|
||||
if cancelParent {
|
||||
Expect(err).To(MatchError(context.Canceled))
|
||||
} else {
|
||||
Expect(err).To(MatchError(context.DeadlineExceeded))
|
||||
}
|
||||
Expect(time.Since(start)).To(BeNumerically("<", time.Second))
|
||||
Eventually(canceled).WithTimeout(time.Second).Should(Receive())
|
||||
}, Entry("Streamable timeout", false, false), Entry("SSE timeout", true, false), Entry("Streamable parent cancellation", false, true), Entry("SSE parent cancellation", true, true))
|
||||
|
||||
It("shares one timeout across negotiation attempts", func() {
|
||||
var gets atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodPost {
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
http.Error(w, "SSE", 405)
|
||||
return
|
||||
}
|
||||
gets.Add(1)
|
||||
<-r.Context().Done()
|
||||
}))
|
||||
defer server.Close()
|
||||
start := time.Now()
|
||||
_, err := Connect(context.Background(), testClient(), server.URL, httpclient.New(), 250*time.Millisecond)
|
||||
Expect(err).To(MatchError(context.DeadlineExceeded))
|
||||
Expect(time.Since(start)).To(BeNumerically("<", 350*time.Millisecond))
|
||||
Expect(gets.Load()).To(Equal(int32(1)))
|
||||
})
|
||||
|
||||
It("does not fall back on a malformed successful initialize response", func() {
|
||||
var gets atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodGet {
|
||||
gets.Add(1)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"jsonrpc": "2.0", "id": 1, "error": map[string]any{"code": -32603, "message": "bad initialization"}})
|
||||
}))
|
||||
defer server.Close()
|
||||
session, err := Connect(context.Background(), testClient(), server.URL, httpclient.New(), time.Second)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(session).To(BeNil())
|
||||
Expect(gets.Load()).To(BeZero())
|
||||
})
|
||||
It("rejects a foreign SSE message endpoint before attaching credentials", func() {
|
||||
var leaked atomic.Int32
|
||||
foreign := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { leaked.Add(1); http.Error(w, "unexpected", 500) }))
|
||||
defer foreign.Close()
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
http.Error(w, "SSE", 405)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = fmt.Fprintf(w, "event: endpoint\ndata: %s/messages\n\n", foreign.URL)
|
||||
w.(http.Flusher).Flush()
|
||||
<-r.Context().Done()
|
||||
}))
|
||||
defer server.Close()
|
||||
base := httpclient.HardenedTransport()
|
||||
hc := httpclient.New(httpclient.WithTransport(roundTripFunc(func(r *http.Request) (*http.Response, error) {
|
||||
r.Header.Set("Authorization", "Bearer secret")
|
||||
return base.RoundTrip(r)
|
||||
})))
|
||||
session, err := Connect(context.Background(), testClient(), server.URL, hc, time.Second)
|
||||
Expect(session).To(BeNil())
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(leaked.Load()).To(BeZero())
|
||||
})
|
||||
|
||||
It("closes a session that completes after the caller times out", func() {
|
||||
initialized := make(chan struct{})
|
||||
release := make(chan struct{})
|
||||
exited := make(chan struct{})
|
||||
handler := mcp.NewStreamableHTTPHandler(func(*http.Request) *mcp.Server { return testServer() }, nil)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodGet {
|
||||
defer close(exited)
|
||||
}
|
||||
handler.ServeHTTP(w, r)
|
||||
}))
|
||||
defer server.Close()
|
||||
base := httpclient.HardenedTransport()
|
||||
hc := httpclient.New(httpclient.WithTransport(roundTripFunc(func(r *http.Request) (*http.Response, error) {
|
||||
response, err := base.RoundTrip(r)
|
||||
if r.Method == http.MethodPost && r.Header.Get("Mcp-Protocol-Version") != "" && err == nil {
|
||||
close(initialized)
|
||||
// Deliberately return late even after cancellation, as a custom transport may.
|
||||
<-release
|
||||
}
|
||||
return response, err
|
||||
})))
|
||||
client := testClient()
|
||||
session, err := Connect(context.Background(), client, server.URL, hc, 100*time.Millisecond)
|
||||
close(release)
|
||||
Expect(session).To(BeNil())
|
||||
Expect(err).To(MatchError(context.DeadlineExceeded))
|
||||
Expect(initialized).To(BeClosed())
|
||||
Eventually(exited).WithTimeout(time.Second).Should(BeClosed())
|
||||
})
|
||||
|
||||
})
|
||||
Reference in new issue
Block a user