diff --git a/core/http/endpoints/mcp/tools.go b/core/http/endpoints/mcp/tools.go index 02ea46037..37a1460ec 100644 --- a/core/http/endpoints/mcp/tools.go +++ b/core/http/endpoints/mcp/tools.go @@ -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{ diff --git a/docs/content/features/mcp.md b/docs/content/features/mcp.md index d898434a3..1aeef1c03 100644 --- a/docs/content/features/mcp.md +++ b/docs/content/features/mcp.md @@ -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`) diff --git a/pkg/mcptransport/suite_test.go b/pkg/mcptransport/suite_test.go new file mode 100644 index 000000000..7009ec88a --- /dev/null +++ b/pkg/mcptransport/suite_test.go @@ -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") } diff --git a/pkg/mcptransport/transport.go b/pkg/mcptransport/transport.go new file mode 100644 index 000000000..aa8209b42 --- /dev/null +++ b/pkg/mcptransport/transport.go @@ -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() +} diff --git a/pkg/mcptransport/transport_test.go b/pkg/mcptransport/transport_test.go new file mode 100644 index 000000000..dff23334d --- /dev/null +++ b/pkg/mcptransport/transport_test.go @@ -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()) + }) + +})