返回 DeepSeek-Reasonix
transport_http.go
根目录 / internal / plugin / transport_http.go
1 package plugin
2
3 import (
4 "bytes"
5 "context"
6 "errors"
7 "fmt"
8 "io"
9 "maps"
10 "net/http"
11 "net/url"
12 "strings"
13 "time"
14
15 mcpsdk "github.com/modelcontextprotocol/go-sdk/mcp"
16 )
17
18 const mcpSubscriptionsListenMethod = "subscriptions/listen"
19
20 // asyncStreamableHTTPSubscriptions keeps the optional SEP-2575 notification
21 // stream from becoming part of the mandatory startup critical path. Some MCP
22 // HTTP bridges buffer a streaming Web Response before writing HTTP response
23 // headers, so the SDK's synchronous transport write would otherwise block
24 // Client.Connect even though server/discover already completed successfully.
25 //
26 // The underlying call still uses the SDK-owned connection and context. A
27 // compliant server therefore keeps delivering notifications normally, while a
28 // buffering server is cancelled with the session without blocking tools/list.
29 func asyncStreamableHTTPSubscriptions(next mcpsdk.MethodHandler) mcpsdk.MethodHandler {
30 return func(ctx context.Context, method string, req mcpsdk.Request) (mcpsdk.Result, error) {
31 if method != mcpSubscriptionsListenMethod {
32 return next(ctx, method, req)
33 }
34 if err := ctx.Err(); err != nil {
35 return nil, err
36 }
37 go func() {
38 _, _ = next(ctx, method, req)
39 }()
40 return &mcpsdk.SubscriptionsListenResult{}, nil
41 }
42 }
43
44 func newHTTPTransport(s Spec) (*sdkSessionTransport, error) {
45 if strings.TrimSpace(s.Type) == "" {
46 s.Type = "http"
47 }
48 // Transient OAuth/probe connections declare no optional capabilities.
49 return newSDKSessionTransport(context.Background(), s, HostProfileCore)
50 }
51
52 func validateMCPURL(name, transport, raw string) error {
53 if strings.TrimSpace(raw) == "" {
54 return fmt.Errorf("%s plugin %q: url is required", transport, name)
55 }
56 u, err := url.Parse(raw)
57 if err != nil || u == nil || u.Scheme == "" || u.Host == "" {
58 return fmt.Errorf("%s plugin %q: invalid url", transport, name)
59 }
60 switch strings.ToLower(u.Scheme) {
61 case "http", "https":
62 return nil
63 default:
64 return fmt.Errorf("%s plugin %q: url must use http or https", transport, name)
65 }
66 }
67
68 func newMCPHTTPClient(lifetime context.Context, s Spec) (*http.Client, error) {
69 origin, err := url.Parse(strings.TrimSpace(s.URL))
70 if err != nil || origin == nil || origin.Host == "" {
71 return nil, fmt.Errorf("invalid MCP endpoint")
72 }
73 headers := make(map[string]string, len(s.Headers))
74 maps.Copy(headers, s.Headers)
75 base := http.DefaultTransport.(*http.Transport).Clone()
76 client := &http.Client{
77 Transport: &sameOriginMCPRoundTripper{
78 origin: origin,
79 headers: headers,
80 base: base,
81 lifetime: lifetime,
82 },
83 }
84 client.CheckRedirect = func(req *http.Request, _ []*http.Request) error {
85 if sameHTTPOrigin(origin, req.URL) {
86 return nil
87 }
88 return http.ErrUseLastResponse
89 }
90 return client, nil
91 }
92
93 type sameOriginMCPRoundTripper struct {
94 origin *url.URL
95 headers map[string]string
96 base http.RoundTripper
97 lifetime context.Context
98 }
99
100 func (rt *sameOriginMCPRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
101 if req == nil || !sameHTTPOrigin(rt.origin, req.URL) {
102 return nil, errors.New("MCP request changed origin; configured headers were not sent")
103 }
104 requestCtx := req.Context()
105 cancelRequest := func() {}
106 stopLifetime := func() bool { return true }
107 // Keep protocol cleanup independent from the session lifetime: Close first
108 // cancels active GET/POST requests, then the SDK sends this bounded DELETE.
109 if req.Method != http.MethodDelete {
110 var cancel context.CancelFunc
111 requestCtx, cancel = context.WithCancel(req.Context())
112 cancelRequest = cancel
113 if rt.lifetime != nil {
114 stopLifetime = context.AfterFunc(rt.lifetime, cancelRequest)
115 }
116 }
117 cancelLifetimeRequest := func() {
118 stopLifetime()
119 cancelRequest()
120 }
121 request := req.Clone(requestCtx)
122 request.Header = req.Header.Clone()
123 for key, value := range rt.headers {
124 request.Header.Set(key, value)
125 }
126
127 base := rt.base
128 if base == nil {
129 base = http.DefaultTransport
130 }
131 if request.Method != http.MethodDelete {
132 response, err := base.RoundTrip(request)
133 return responseWithCancel(response, err, cancelLifetimeRequest)
134 }
135
136 deleteCtx, cancelDelete := context.WithTimeout(request.Context(), 2*time.Second)
137 request = request.Clone(deleteCtx)
138 response, err := base.RoundTrip(request)
139 return responseWithCancel(response, err, func() {
140 cancelDelete()
141 cancelLifetimeRequest()
142 })
143 }
144
145 func responseWithCancel(response *http.Response, err error, cancel func()) (*http.Response, error) {
146 if err != nil {
147 cancel()
148 return nil, err
149 }
150 if response.Body == nil {
151 cancel()
152 return response, nil
153 }
154 response.Body = &cancelOnCloseBody{ReadCloser: response.Body, cancel: cancel}
155 return response, nil
156 }
157
158 func (rt *sameOriginMCPRoundTripper) CloseIdleConnections() {
159 if closer, ok := rt.base.(interface{ CloseIdleConnections() }); ok {
160 closer.CloseIdleConnections()
161 }
162 }
163
164 type cancelOnCloseBody struct {
165 io.ReadCloser
166 cancel func()
167 }
168
169 func (b *cancelOnCloseBody) Close() error {
170 err := b.ReadCloser.Close()
171 b.cancel()
172 return err
173 }
174
175 func sameHTTPOrigin(a, b *url.URL) bool {
176 if a == nil || b == nil || !strings.EqualFold(a.Scheme, b.Scheme) || !strings.EqualFold(a.Hostname(), b.Hostname()) {
177 return false
178 }
179 effectivePort := func(u *url.URL) string {
180 if port := u.Port(); port != "" {
181 return port
182 }
183 switch strings.ToLower(u.Scheme) {
184 case "http":
185 return "80"
186 case "https":
187 return "443"
188 default:
189 return ""
190 }
191 }
192 return effectivePort(a) == effectivePort(b)
193 }
194
195 func (t *sdkSessionTransport) newEndpoint(ctx context.Context) (sdkEndpoint, error) {
196 if t.endpointFactory != nil {
197 return t.endpointFactory(ctx)
198 }
199 switch canonicalMCPRuntimeTransport(t.spec.Type) {
200 case "stdio":
201 process, err := newStdioTransport(ctx, t.spec)
202 if err != nil {
203 return sdkEndpoint{}, err
204 }
205 return sdkEndpoint{
206 transport: &mcpsdk.IOTransport{Reader: process.stdout, Writer: process.stdin},
207 close: process.close,
208 startupStderr: process.startupStderr,
209 }, nil
210 case "streamable-http":
211 client, err := newMCPHTTPClient(ctx, t.spec)
212 if err != nil {
213 return sdkEndpoint{}, err
214 }
215 return sdkEndpoint{
216 transport: &mcpsdk.StreamableClientTransport{
217 Endpoint: t.spec.URL,
218 HTTPClient: client,
219 MaxRetries: 5,
220 OAuthHandler: t.oauth,
221 },
222 close: client.CloseIdleConnections,
223 }, nil
224 case "sse":
225 client, err := newMCPHTTPClient(ctx, t.spec)
226 if err != nil {
227 return sdkEndpoint{}, err
228 }
229 return sdkEndpoint{
230 transport: &mcpsdk.SSEClientTransport{Endpoint: t.spec.URL, HTTPClient: client},
231 close: client.CloseIdleConnections,
232 }, nil
233 default:
234 return sdkEndpoint{}, fmt.Errorf("unknown MCP transport %q", t.spec.Type)
235 }
236 }
237
238 // do is retained as a narrow HTTP security test hook. MCP protocol traffic goes
239 // through the SDK transport above.
240 func (t *sdkSessionTransport) do(ctx context.Context, body []byte) (*http.Response, error) {
241 client, err := newMCPHTTPClient(ctx, t.spec)
242 if err != nil {
243 return nil, err
244 }
245 req, err := http.NewRequestWithContext(ctx, http.MethodPost, t.spec.URL, bytes.NewReader(body))
246 if err != nil {
247 return nil, err
248 }
249 req.Header.Set("Content-Type", "application/json")
250 req.Header.Set("Accept", "application/json, text/event-stream")
251 return client.Do(req)
252 }
253
253 lines GO