Skip to content

Commit 25178cd

Browse files
committed
fix: 修复gpt->claude同步请求返回sse的bug
1 parent a461538 commit 25178cd

1 file changed

Lines changed: 187 additions & 58 deletions

File tree

backend/internal/service/openai_gateway_messages.go

Lines changed: 187 additions & 58 deletions
Original file line numberDiff line numberDiff line change
@@ -39,14 +39,19 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
3939
return nil, fmt.Errorf("parse anthropic request: %w", err)
4040
}
4141
originalModel := anthropicReq.Model
42-
isStream := anthropicReq.Stream
42+
clientStream := anthropicReq.Stream // client's original stream preference
4343

4444
// 2. Convert Anthropic → Responses
4545
responsesReq, err := apicompat.AnthropicToResponses(&anthropicReq)
4646
if err != nil {
4747
return nil, fmt.Errorf("convert anthropic to responses: %w", err)
4848
}
4949

50+
// Upstream always uses streaming (upstream may not support sync mode).
51+
// The client's original preference determines the response format.
52+
responsesReq.Stream = true
53+
isStream := true
54+
5055
// 2b. Handle BetaFastMode → service_tier: "priority"
5156
if containsBetaToken(c.GetHeader("anthropic-beta"), claude.BetaFastMode) {
5257
responsesReq.ServiceTier = "priority"
@@ -169,12 +174,14 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
169174
}
170175

171176
// 9. Handle normal response
177+
// Upstream is always streaming; choose response format based on client preference.
172178
var result *OpenAIForwardResult
173179
var handleErr error
174-
if isStream {
180+
if clientStream {
175181
result, handleErr = s.handleAnthropicStreamingResponse(resp, c, originalModel, mappedModel, startTime)
176182
} else {
177-
result, handleErr = s.handleAnthropicNonStreamingResponse(resp, c, originalModel, mappedModel, startTime)
183+
// Client wants JSON: buffer the streaming response and assemble a JSON reply.
184+
result, handleErr = s.handleAnthropicBufferedStreamingResponse(resp, c, originalModel, mappedModel, startTime)
178185
}
179186

180187
// Propagate ServiceTier and ReasoningEffort to result for billing
@@ -256,9 +263,13 @@ func (s *OpenAIGatewayService) handleAnthropicErrorResponse(
256263
return nil, fmt.Errorf("upstream error: %d %s", resp.StatusCode, upstreamMsg)
257264
}
258265

259-
// handleAnthropicNonStreamingResponse reads a Responses API JSON response,
260-
// converts it to Anthropic Messages format, and writes it to the client.
261-
func (s *OpenAIGatewayService) handleAnthropicNonStreamingResponse(
266+
// handleAnthropicBufferedStreamingResponse reads all Responses SSE events from
267+
// the upstream streaming response, finds the terminal event (response.completed
268+
// / response.incomplete / response.failed), converts the complete response to
269+
// Anthropic Messages JSON format, and writes it to the client.
270+
// This is used when the client requested stream=false but the upstream is always
271+
// streaming.
272+
func (s *OpenAIGatewayService) handleAnthropicBufferedStreamingResponse(
262273
resp *http.Response,
263274
c *gin.Context,
264275
originalModel string,
@@ -267,29 +278,61 @@ func (s *OpenAIGatewayService) handleAnthropicNonStreamingResponse(
267278
) (*OpenAIForwardResult, error) {
268279
requestID := resp.Header.Get("x-request-id")
269280

270-
respBody, err := io.ReadAll(resp.Body)
271-
if err != nil {
272-
return nil, fmt.Errorf("read upstream response: %w", err)
273-
}
281+
scanner := bufio.NewScanner(resp.Body)
282+
scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024)
274283

275-
var responsesResp apicompat.ResponsesResponse
276-
if err := json.Unmarshal(respBody, &responsesResp); err != nil {
277-
return nil, fmt.Errorf("parse responses response: %w", err)
278-
}
284+
var finalResponse *apicompat.ResponsesResponse
285+
var usage OpenAIUsage
279286

280-
anthropicResp := apicompat.ResponsesToAnthropic(&responsesResp, originalModel)
287+
for scanner.Scan() {
288+
line := scanner.Text()
281289

282-
var usage OpenAIUsage
283-
if responsesResp.Usage != nil {
284-
usage = OpenAIUsage{
285-
InputTokens: responsesResp.Usage.InputTokens,
286-
OutputTokens: responsesResp.Usage.OutputTokens,
290+
if !strings.HasPrefix(line, "data: ") || line == "data: [DONE]" {
291+
continue
287292
}
288-
if responsesResp.Usage.InputTokensDetails != nil {
289-
usage.CacheReadInputTokens = responsesResp.Usage.InputTokensDetails.CachedTokens
293+
payload := line[6:]
294+
295+
var event apicompat.ResponsesStreamEvent
296+
if err := json.Unmarshal([]byte(payload), &event); err != nil {
297+
logger.L().Warn("openai messages buffered: failed to parse event",
298+
zap.Error(err),
299+
zap.String("request_id", requestID),
300+
)
301+
continue
302+
}
303+
304+
// Terminal events carry the complete ResponsesResponse with output + usage.
305+
if (event.Type == "response.completed" || event.Type == "response.incomplete" || event.Type == "response.failed") &&
306+
event.Response != nil {
307+
finalResponse = event.Response
308+
if event.Response.Usage != nil {
309+
usage = OpenAIUsage{
310+
InputTokens: event.Response.Usage.InputTokens,
311+
OutputTokens: event.Response.Usage.OutputTokens,
312+
}
313+
if event.Response.Usage.InputTokensDetails != nil {
314+
usage.CacheReadInputTokens = event.Response.Usage.InputTokensDetails.CachedTokens
315+
}
316+
}
290317
}
291318
}
292319

320+
if err := scanner.Err(); err != nil {
321+
if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
322+
logger.L().Warn("openai messages buffered: read error",
323+
zap.Error(err),
324+
zap.String("request_id", requestID),
325+
)
326+
}
327+
}
328+
329+
if finalResponse == nil {
330+
writeAnthropicError(c, http.StatusBadGateway, "api_error", "Upstream stream ended without a terminal response event")
331+
return nil, fmt.Errorf("upstream stream ended without terminal event")
332+
}
333+
334+
anthropicResp := apicompat.ResponsesToAnthropic(finalResponse, originalModel)
335+
293336
if s.responseHeaderFilter != nil {
294337
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
295338
}
@@ -307,6 +350,9 @@ func (s *OpenAIGatewayService) handleAnthropicNonStreamingResponse(
307350

308351
// handleAnthropicStreamingResponse reads Responses SSE events from upstream,
309352
// converts each to Anthropic SSE events, and writes them to the client.
353+
// When StreamKeepaliveInterval is configured, it uses a goroutine + channel
354+
// pattern to send Anthropic ping events during periods of upstream silence,
355+
// preventing proxy/client timeout disconnections.
310356
func (s *OpenAIGatewayService) handleAnthropicStreamingResponse(
311357
resp *http.Response,
312358
c *gin.Context,
@@ -322,6 +368,7 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse(
322368
c.Writer.Header().Set("Content-Type", "text/event-stream")
323369
c.Writer.Header().Set("Cache-Control", "no-cache")
324370
c.Writer.Header().Set("Connection", "keep-alive")
371+
c.Writer.Header().Set("X-Accel-Buffering", "no")
325372
c.Writer.WriteHeader(http.StatusOK)
326373

327374
state := apicompat.NewResponsesEventToAnthropicState()
@@ -333,28 +380,35 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse(
333380
scanner := bufio.NewScanner(resp.Body)
334381
scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024)
335382

336-
for scanner.Scan() {
337-
line := scanner.Text()
338-
339-
if !strings.HasPrefix(line, "data: ") || line == "data: [DONE]" {
340-
continue
383+
// resultWithUsage builds the final result snapshot.
384+
resultWithUsage := func() *OpenAIForwardResult {
385+
return &OpenAIForwardResult{
386+
RequestID: requestID,
387+
Usage: usage,
388+
Model: originalModel,
389+
BillingModel: mappedModel,
390+
Stream: true,
391+
Duration: time.Since(startTime),
392+
FirstTokenMs: firstTokenMs,
341393
}
342-
payload := line[6:]
394+
}
343395

396+
// processDataLine handles a single "data: ..." SSE line from upstream.
397+
// Returns (clientDisconnected bool).
398+
processDataLine := func(payload string) bool {
344399
if firstChunk {
345400
firstChunk = false
346401
ms := int(time.Since(startTime).Milliseconds())
347402
firstTokenMs = &ms
348403
}
349404

350-
// Parse the Responses SSE event
351405
var event apicompat.ResponsesStreamEvent
352406
if err := json.Unmarshal([]byte(payload), &event); err != nil {
353407
logger.L().Warn("openai messages stream: failed to parse event",
354408
zap.Error(err),
355409
zap.String("request_id", requestID),
356410
)
357-
continue
411+
return false
358412
}
359413

360414
// Extract usage from completion events
@@ -381,56 +435,131 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse(
381435
continue
382436
}
383437
if _, err := fmt.Fprint(c.Writer, sse); err != nil {
384-
// Client disconnected — return collected usage
385438
logger.L().Info("openai messages stream: client disconnected",
386439
zap.String("request_id", requestID),
387440
)
388-
return &OpenAIForwardResult{
389-
RequestID: requestID,
390-
Usage: usage,
391-
Model: originalModel,
392-
BillingModel: mappedModel,
393-
Stream: true,
394-
Duration: time.Since(startTime),
395-
FirstTokenMs: firstTokenMs,
396-
}, nil
441+
return true
397442
}
398443
}
399444
if len(events) > 0 {
400445
c.Writer.Flush()
401446
}
447+
return false
402448
}
403449

404-
if err := scanner.Err(); err != nil {
405-
if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
450+
// finalizeStream sends any remaining Anthropic events and returns the result.
451+
finalizeStream := func() (*OpenAIForwardResult, error) {
452+
if finalEvents := apicompat.FinalizeResponsesAnthropicStream(state); len(finalEvents) > 0 {
453+
for _, evt := range finalEvents {
454+
sse, err := apicompat.ResponsesAnthropicEventToSSE(evt)
455+
if err != nil {
456+
continue
457+
}
458+
fmt.Fprint(c.Writer, sse) //nolint:errcheck
459+
}
460+
c.Writer.Flush()
461+
}
462+
return resultWithUsage(), nil
463+
}
464+
465+
// handleScanErr logs scanner errors if meaningful.
466+
handleScanErr := func(err error) {
467+
if err != nil && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
406468
logger.L().Warn("openai messages stream: read error",
407469
zap.Error(err),
408470
zap.String("request_id", requestID),
409471
)
410472
}
411473
}
412474

413-
// Ensure the Anthropic stream is properly terminated
414-
if finalEvents := apicompat.FinalizeResponsesAnthropicStream(state); len(finalEvents) > 0 {
415-
for _, evt := range finalEvents {
416-
sse, err := apicompat.ResponsesAnthropicEventToSSE(evt)
417-
if err != nil {
475+
// ── Determine keepalive interval ──
476+
keepaliveInterval := time.Duration(0)
477+
if s.cfg != nil && s.cfg.Gateway.StreamKeepaliveInterval > 0 {
478+
keepaliveInterval = time.Duration(s.cfg.Gateway.StreamKeepaliveInterval) * time.Second
479+
}
480+
481+
// ── No keepalive: fast synchronous path (no goroutine overhead) ──
482+
if keepaliveInterval <= 0 {
483+
for scanner.Scan() {
484+
line := scanner.Text()
485+
if !strings.HasPrefix(line, "data: ") || line == "data: [DONE]" {
418486
continue
419487
}
420-
fmt.Fprint(c.Writer, sse) //nolint:errcheck
488+
if processDataLine(line[6:]) {
489+
return resultWithUsage(), nil
490+
}
421491
}
422-
c.Writer.Flush()
492+
handleScanErr(scanner.Err())
493+
return finalizeStream()
423494
}
424495

425-
return &OpenAIForwardResult{
426-
RequestID: requestID,
427-
Usage: usage,
428-
Model: originalModel,
429-
BillingModel: mappedModel,
430-
Stream: true,
431-
Duration: time.Since(startTime),
432-
FirstTokenMs: firstTokenMs,
433-
}, nil
496+
// ── With keepalive: goroutine + channel + select ──
497+
type scanEvent struct {
498+
line string
499+
err error
500+
}
501+
events := make(chan scanEvent, 16)
502+
done := make(chan struct{})
503+
sendEvent := func(ev scanEvent) bool {
504+
select {
505+
case events <- ev:
506+
return true
507+
case <-done:
508+
return false
509+
}
510+
}
511+
go func() {
512+
defer close(events)
513+
for scanner.Scan() {
514+
if !sendEvent(scanEvent{line: scanner.Text()}) {
515+
return
516+
}
517+
}
518+
if err := scanner.Err(); err != nil {
519+
_ = sendEvent(scanEvent{err: err})
520+
}
521+
}()
522+
defer close(done)
523+
524+
keepaliveTicker := time.NewTicker(keepaliveInterval)
525+
defer keepaliveTicker.Stop()
526+
lastDataAt := time.Now()
527+
528+
for {
529+
select {
530+
case ev, ok := <-events:
531+
if !ok {
532+
// Upstream closed
533+
return finalizeStream()
534+
}
535+
if ev.err != nil {
536+
handleScanErr(ev.err)
537+
return finalizeStream()
538+
}
539+
lastDataAt = time.Now()
540+
line := ev.line
541+
if !strings.HasPrefix(line, "data: ") || line == "data: [DONE]" {
542+
continue
543+
}
544+
if processDataLine(line[6:]) {
545+
return resultWithUsage(), nil
546+
}
547+
548+
case <-keepaliveTicker.C:
549+
if time.Since(lastDataAt) < keepaliveInterval {
550+
continue
551+
}
552+
// Send Anthropic-format ping event
553+
if _, err := fmt.Fprint(c.Writer, "event: ping\ndata: {\"type\":\"ping\"}\n\n"); err != nil {
554+
// Client disconnected
555+
logger.L().Info("openai messages stream: client disconnected during keepalive",
556+
zap.String("request_id", requestID),
557+
)
558+
return resultWithUsage(), nil
559+
}
560+
c.Writer.Flush()
561+
}
562+
}
434563
}
435564

436565
// writeAnthropicError writes an error response in Anthropic Messages API format.

0 commit comments

Comments
 (0)