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:
localai-org-maint-botandlocalai-org-maint-bot authored and GitHub committed 2026-10-08 16:41:33 +02:00
1 parent 4a1089be35
commit 0102751b31
5 files changed
+511 -4

No files matched your search

+3 -4
View File
@@ -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{
+4
View File
@@ -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`)
+10
View File
@@ -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") }
+193
View File
@@ -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()
}
+301
View File
@@ -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())
})
})