返回 DeepSeek-Reasonix
transport_sse.go
根目录 / internal / plugin / transport_sse.go
1 package plugin
2
3 import (
4 "bufio"
5 "bytes"
6 "context"
7 "encoding/json"
8 "errors"
9 "fmt"
10 "io"
11 "net/http"
12 "net/url"
13 "strings"
14 "sync"
15 "time"
16
17 "reasonix/internal/tool"
18 )
19
20 // sseTransport implements MCP's legacy HTTP+SSE transport. The client keeps a
21 // long-lived GET stream open; the server announces a POST endpoint through an
22 // `event: endpoint` frame, and all JSON-RPC messages travel to that endpoint
23 // while responses and server messages return on the GET stream.
24 type sseTransport struct {
25 name string
26 getURL *url.URL
27 headers map[string]string
28 client *http.Client
29 roots []mcpRoot
30 progress progressRouter
31 replies chan inboundMessage
32 replyTimeout time.Duration
33
34 ctx context.Context
35 cancel context.CancelFunc
36
37 callMu sync.Mutex
38 mu sync.Mutex
39 nextID int
40 pending map[int]chan rpcResponse
41 readErr error
42 endpoint *url.URL
43 endpointErr error
44 endpointReady chan struct{}
45 endpointOnce sync.Once
46 closeOnce sync.Once
47 }
48
49 // sseReplyQueueBound keeps a server-request flood from creating an unbounded
50 // number of blocked HTTP POST goroutines. Overflow is intentionally dropped so
51 // the GET reader can continue routing client responses and notifications.
52 const sseReplyQueueBound = 16
53
54 func newSSETransport(ctx context.Context, s Spec) (*sseTransport, error) {
55 if strings.TrimSpace(s.URL) == "" {
56 return nil, fmt.Errorf("sse plugin %q: url is required", s.Name)
57 }
58 getURL, err := url.Parse(s.URL)
59 if err != nil || getURL.Scheme == "" || getURL.Host == "" {
60 return nil, fmt.Errorf("sse plugin %q: invalid url %q", s.Name, s.URL)
61 }
62 headers := make(map[string]string, len(s.Headers))
63 for key, value := range s.Headers {
64 headers[key] = value
65 }
66 lifeCtx, cancel := context.WithCancel(ctx)
67 t := &sseTransport{
68 name: s.Name,
69 getURL: getURL,
70 headers: headers,
71 roots: mcpRoots(s.WorkspaceRoot),
72 replies: make(chan inboundMessage, sseReplyQueueBound),
73 replyTimeout: s.CallTimeout,
74 ctx: lifeCtx,
75 cancel: cancel,
76 pending: map[int]chan rpcResponse{},
77 endpointReady: make(chan struct{}),
78 }
79 if t.replyTimeout <= 0 {
80 t.replyTimeout = s.DefaultCallTimeout
81 }
82 if t.replyTimeout <= 0 {
83 t.replyTimeout = defaultCallTimeout
84 }
85 t.client = &http.Client{CheckRedirect: func(req *http.Request, via []*http.Request) error {
86 if len(via) == 0 || sameHTTPOrigin(via[0].URL, req.URL) {
87 return nil
88 }
89 return http.ErrUseLastResponse
90 }}
91 go t.replyLoop()
92 go t.readLoop()
93 return t, nil
94 }
95
96 func (t *sseTransport) registerProgress(token string, sink tool.ProgressFunc) func() {
97 return t.progress.registerProgress(token, sink)
98 }
99
100 func (t *sseTransport) readLoop() {
101 defer close(t.replies)
102 req, err := http.NewRequestWithContext(t.ctx, http.MethodGet, t.getURL.String(), nil)
103 if err != nil {
104 t.fail(err)
105 return
106 }
107 req.Header.Set("Accept", "text/event-stream")
108 for key, value := range t.headers {
109 req.Header.Set(key, value)
110 }
111 resp, err := t.client.Do(req)
112 if err != nil {
113 t.fail(err)
114 return
115 }
116 defer resp.Body.Close()
117 if resp.StatusCode/100 != 2 {
118 body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
119 t.fail(fmt.Errorf("GET %s: http %d: %s", t.getURL, resp.StatusCode, strings.TrimSpace(string(body))))
120 return
121 }
122 if !strings.HasPrefix(strings.ToLower(resp.Header.Get("Content-Type")), "text/event-stream") {
123 t.fail(fmt.Errorf("GET %s: expected text/event-stream, got %q", t.getURL, resp.Header.Get("Content-Type")))
124 return
125 }
126
127 baseURL := resp.Request.URL
128 scanner := bufio.NewScanner(resp.Body)
129 scanner.Buffer(make([]byte, 0, 64*1024), maxHTTPBody)
130 eventName := "message"
131 var data strings.Builder
132 dispatch := func() {
133 if data.Len() == 0 {
134 eventName = "message"
135 return
136 }
137 payload := data.String()
138 data.Reset()
139 t.handleEvent(eventName, payload, baseURL)
140 eventName = "message"
141 }
142 for scanner.Scan() {
143 line := scanner.Text()
144 if line == "" {
145 dispatch()
146 continue
147 }
148 if strings.HasPrefix(line, ":") {
149 continue
150 }
151 if value, ok := strings.CutPrefix(line, "event:"); ok {
152 eventName = strings.TrimSpace(value)
153 continue
154 }
155 if value, ok := strings.CutPrefix(line, "data:"); ok {
156 if data.Len() > 0 {
157 data.WriteByte('\n')
158 }
159 data.WriteString(strings.TrimPrefix(value, " "))
160 }
161 }
162 dispatch()
163 if err := scanner.Err(); err != nil {
164 t.fail(fmt.Errorf("read SSE: %w", err))
165 return
166 }
167 t.fail(io.EOF)
168 }
169
170 func (t *sseTransport) handleEvent(eventName, payload string, baseURL *url.URL) {
171 switch eventName {
172 case "endpoint":
173 endpoint, err := url.Parse(strings.TrimSpace(payload))
174 if err == nil {
175 endpoint = baseURL.ResolveReference(endpoint)
176 if !sameHTTPOrigin(baseURL, endpoint) {
177 err = fmt.Errorf("server announced cross-origin endpoint %s", endpoint)
178 }
179 }
180 t.setEndpoint(endpoint, err)
181 case "message", "":
182 t.handleMessage([]byte(payload))
183 }
184 }
185
186 func (t *sseTransport) setEndpoint(endpoint *url.URL, err error) {
187 set := false
188 t.endpointOnce.Do(func() {
189 t.mu.Lock()
190 t.endpoint = endpoint
191 t.endpointErr = err
192 t.mu.Unlock()
193 close(t.endpointReady)
194 set = true
195 })
196 if set && err != nil {
197 t.failPending(err)
198 }
199 }
200
201 func (t *sseTransport) handleMessage(payload []byte) {
202 message, ok := decodeInboundMessage(payload)
203 if !ok {
204 return
205 }
206 if message.Method != "" {
207 if isNotificationID(message.ID) {
208 if message.Method == "notifications/progress" {
209 t.progress.dispatchProgress(message.Params)
210 }
211 return
212 }
213 select {
214 case t.replies <- message:
215 default:
216 // Let the remote server time out excess requests. Blocking here could
217 // prevent an unrelated client response from ever reaching its caller.
218 }
219 return
220 }
221 var response rpcResponse
222 if err := json.Unmarshal(payload, &response); err != nil {
223 return
224 }
225 t.mu.Lock()
226 ch := t.pending[response.ID]
227 delete(t.pending, response.ID)
228 t.mu.Unlock()
229 if ch != nil {
230 ch <- response
231 }
232 }
233
234 // replyLoop serializes server-request responses through one bounded worker.
235 // Each POST has a finite deadline inherited from the server's call-timeout
236 // configuration; after a transport error the worker keeps draining so the SSE
237 // reader never waits on the queue.
238 func (t *sseTransport) replyLoop() {
239 dead := false
240 for message := range t.replies {
241 if dead {
242 continue
243 }
244 body, err := json.Marshal(serverRequestReply(message.ID, message.Method, t.roots))
245 if err == nil {
246 ctx, cancel := context.WithTimeout(t.ctx, t.replyTimeout)
247 err = t.post(ctx, body)
248 cancel()
249 }
250 if err != nil && t.ctx.Err() == nil {
251 t.fail(fmt.Errorf("reply to %s: %w", message.Method, err))
252 dead = true
253 }
254 }
255 }
256
257 func (t *sseTransport) call(ctx context.Context, method string, params any) (json.RawMessage, error) {
258 t.callMu.Lock()
259 defer t.callMu.Unlock()
260
261 if err := t.waitEndpoint(ctx); err != nil {
262 return nil, fmt.Errorf("plugin %q: %s: %w", t.name, method, err)
263 }
264 t.mu.Lock()
265 if t.readErr != nil {
266 err := t.readErr
267 t.mu.Unlock()
268 return nil, fmt.Errorf("plugin %q: %s: %w", t.name, method, err)
269 }
270 t.nextID++
271 id := t.nextID
272 ch := make(chan rpcResponse, 1)
273 t.pending[id] = ch
274 t.mu.Unlock()
275 defer func() {
276 t.mu.Lock()
277 delete(t.pending, id)
278 t.mu.Unlock()
279 }()
280
281 body, err := json.Marshal(rpcRequest{JSONRPC: "2.0", ID: id, Method: method, Params: params})
282 if err != nil {
283 return nil, err
284 }
285 if err := t.post(ctx, body); err != nil {
286 return nil, fmt.Errorf("plugin %q: %s: %w", t.name, method, err)
287 }
288 select {
289 case <-ctx.Done():
290 return nil, ctx.Err()
291 case response, ok := <-ch:
292 if !ok {
293 t.mu.Lock()
294 err := t.readErr
295 t.mu.Unlock()
296 return nil, fmt.Errorf("plugin %q: %s: %w", t.name, method, err)
297 }
298 if response.Error != nil {
299 return nil, fmt.Errorf("plugin %q: %w", t.name, response.Error)
300 }
301 return response.Result, nil
302 }
303 }
304
305 func (t *sseTransport) notify(ctx context.Context, method string, params any) error {
306 if err := t.waitEndpoint(ctx); err != nil {
307 return fmt.Errorf("plugin %q: %s: %w", t.name, method, err)
308 }
309 body, err := json.Marshal(rpcRequest{JSONRPC: "2.0", Method: method, Params: params})
310 if err != nil {
311 return err
312 }
313 if err := t.post(ctx, body); err != nil {
314 return fmt.Errorf("plugin %q: %s: %w", t.name, method, err)
315 }
316 return nil
317 }
318
319 func (t *sseTransport) waitEndpoint(ctx context.Context) error {
320 select {
321 case <-ctx.Done():
322 return ctx.Err()
323 case <-t.endpointReady:
324 t.mu.Lock()
325 defer t.mu.Unlock()
326 if t.endpointErr != nil {
327 return t.endpointErr
328 }
329 if t.endpoint == nil {
330 return errSSEEndpointMissing
331 }
332 return nil
333 }
334 }
335
336 var errSSEEndpointMissing = errors.New("SSE stream ended before announcing an endpoint")
337
338 func (t *sseTransport) post(ctx context.Context, body []byte) error {
339 t.mu.Lock()
340 endpoint := t.endpoint
341 t.mu.Unlock()
342 if endpoint == nil {
343 return errSSEEndpointMissing
344 }
345 req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint.String(), bytes.NewReader(body))
346 if err != nil {
347 return err
348 }
349 req.Header.Set("Content-Type", "application/json")
350 req.Header.Set("Accept", "application/json, text/event-stream")
351 for key, value := range t.headers {
352 req.Header.Set(key, value)
353 }
354 resp, err := t.client.Do(req)
355 if err != nil {
356 return err
357 }
358 defer resp.Body.Close()
359 responseBody, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
360 if resp.StatusCode/100 != 2 {
361 return fmt.Errorf("http %d: %s", resp.StatusCode, strings.TrimSpace(string(responseBody)))
362 }
363 return nil
364 }
365
366 func (t *sseTransport) fail(err error) {
367 t.endpointOnce.Do(func() {
368 t.mu.Lock()
369 t.endpointErr = err
370 t.mu.Unlock()
371 close(t.endpointReady)
372 })
373 t.failPending(err)
374 }
375
376 func (t *sseTransport) failPending(err error) {
377 t.mu.Lock()
378 if t.readErr == nil {
379 t.readErr = err
380 }
381 for id, ch := range t.pending {
382 close(ch)
383 delete(t.pending, id)
384 }
385 t.mu.Unlock()
386 }
387
388 func (t *sseTransport) close() {
389 t.closeOnce.Do(func() {
390 t.cancel()
391 t.client.CloseIdleConnections()
392 })
393 }
394
394 lines GO