Prevent backend stream goroutine leaks
This commit is contained in:
parent
49e43248b0
commit
701e9bcaad
|
|
@ -293,7 +293,7 @@ outer:
|
|||
}
|
||||
|
||||
// streamAnthropicSSE reads the Anthropic SSE stream and sends StreamEvents to ch.
|
||||
func streamAnthropicSSE(body io.Reader, ch chan<- StreamEvent) {
|
||||
func streamAnthropicSSE(ctx context.Context, body io.Reader, ch chan<- StreamEvent) {
|
||||
scanner := newStreamScanner(body)
|
||||
|
||||
type toolAccum struct {
|
||||
|
|
@ -353,11 +353,15 @@ func streamAnthropicSSE(body io.Reader, ch chan<- StreamEvent) {
|
|||
switch v.Delta.Type {
|
||||
case "text_delta":
|
||||
if v.Delta.Text != "" {
|
||||
ch <- StreamEvent{Content: v.Delta.Text}
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{Content: v.Delta.Text}) {
|
||||
return
|
||||
}
|
||||
}
|
||||
case "thinking_delta":
|
||||
if v.Delta.Thinking != "" {
|
||||
ch <- StreamEvent{Reasoning: v.Delta.Thinking}
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{Reasoning: v.Delta.Thinking}) {
|
||||
return
|
||||
}
|
||||
}
|
||||
case "input_json_delta":
|
||||
if t := tools[v.Index]; t != nil {
|
||||
|
|
@ -396,7 +400,9 @@ func streamAnthropicSSE(body io.Reader, ch chan<- StreamEvent) {
|
|||
Arguments: json.RawMessage(t.args.String()),
|
||||
})
|
||||
}
|
||||
ch <- ev
|
||||
if !sendStreamEvent(ctx, ch, ev) {
|
||||
return
|
||||
}
|
||||
done = true
|
||||
|
||||
case "error":
|
||||
|
|
@ -406,7 +412,9 @@ func streamAnthropicSSE(body io.Reader, ch chan<- StreamEvent) {
|
|||
} `json:"error"`
|
||||
}
|
||||
json.Unmarshal(raw, &v) //nolint:errcheck
|
||||
ch <- StreamEvent{Done: true, StopReason: "error: " + v.Error.Message}
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{Done: true, StopReason: "error: " + v.Error.Message}) {
|
||||
return
|
||||
}
|
||||
done = true
|
||||
}
|
||||
}
|
||||
|
|
@ -436,7 +444,7 @@ func streamAnthropicSSE(body io.Reader, ch chan<- StreamEvent) {
|
|||
}
|
||||
|
||||
if err := scanner.Err(); err != nil && !done {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream read: %v", err)}
|
||||
sendStreamEvent(ctx, ch, StreamEvent{Done: true, StopReason: fmt.Sprintf("stream read: %v", err)})
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
package backend
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
|
@ -11,7 +12,7 @@ func collectAnthropic(sse string) []StreamEvent {
|
|||
ch := make(chan StreamEvent, 32)
|
||||
go func() {
|
||||
defer close(ch)
|
||||
streamAnthropicSSE(strings.NewReader(sse), ch)
|
||||
streamAnthropicSSE(context.Background(), strings.NewReader(sse), ch)
|
||||
}()
|
||||
var evs []StreamEvent
|
||||
for ev := range ch {
|
||||
|
|
@ -156,7 +157,7 @@ func TestAnthropicSSE_ScannerError(t *testing.T) {
|
|||
ch := make(chan StreamEvent, 32)
|
||||
go func() {
|
||||
defer close(ch)
|
||||
streamAnthropicSSE(r, ch)
|
||||
streamAnthropicSSE(context.Background(), r, ch)
|
||||
}()
|
||||
var evs []StreamEvent
|
||||
for ev := range ch {
|
||||
|
|
|
|||
|
|
@ -173,7 +173,7 @@ type Backend interface {
|
|||
// label is used in error messages (e.g. "openai", "ollama").
|
||||
// Returns *RateLimitError for 429, *TransientError for 5xx and network errors,
|
||||
// and a plain error for other non-200 responses.
|
||||
func streamRequest(client *http.Client, req *http.Request, label string, parseFn func(io.Reader, chan<- StreamEvent)) (<-chan StreamEvent, error) {
|
||||
func streamRequest(client *http.Client, req *http.Request, label string, parseFn func(context.Context, io.Reader, chan<- StreamEvent)) (<-chan StreamEvent, error) {
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, &TransientError{Message: err.Error()}
|
||||
|
|
@ -204,12 +204,20 @@ func streamRequest(client *http.Client, req *http.Request, label string, parseFn
|
|||
go func() {
|
||||
defer close(ch)
|
||||
defer resp.Body.Close()
|
||||
parseFn(resp.Body, ch)
|
||||
parseFn(req.Context(), resp.Body, ch)
|
||||
}()
|
||||
return ch, nil
|
||||
}
|
||||
|
||||
// isToolUnsupported returns true when a 400/422 response body indicates the
|
||||
func sendStreamEvent(ctx context.Context, ch chan<- StreamEvent, ev StreamEvent) bool {
|
||||
select {
|
||||
case ch <- ev:
|
||||
return true
|
||||
case <-ctx.Done():
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// model does not support tool/function calling.
|
||||
func isToolUnsupported(body string) bool {
|
||||
markers := []string{
|
||||
|
|
|
|||
|
|
@ -179,12 +179,12 @@ func (b *KiroBackend) runStream(ctx context.Context, req *kiroGenerateRequest, c
|
|||
for {
|
||||
token, err := b.authSource.AccessToken(ctx)
|
||||
if err != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("auth: %v", err)}
|
||||
sendStreamEvent(ctx, ch, StreamEvent{Done: true, StopReason: fmt.Sprintf("auth: %v", err)})
|
||||
return
|
||||
}
|
||||
endpoint, err := b.resolveEndpoint(ctx)
|
||||
if err != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("endpoint: %v", err)}
|
||||
sendStreamEvent(ctx, ch, StreamEvent{Done: true, StopReason: fmt.Sprintf("endpoint: %v", err)})
|
||||
return
|
||||
}
|
||||
|
||||
|
|
@ -196,9 +196,7 @@ func (b *KiroBackend) runStream(ctx context.Context, req *kiroGenerateRequest, c
|
|||
if stopped || event.Content == "" {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case ch <- StreamEvent{Content: event.Content}:
|
||||
case <-ctx.Done():
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{Content: event.Content}) {
|
||||
stopped = true
|
||||
}
|
||||
},
|
||||
|
|
@ -206,9 +204,7 @@ func (b *KiroBackend) runStream(ctx context.Context, req *kiroGenerateRequest, c
|
|||
if stopped || (event.Text == "" && event.RedactedContent == "" && event.Signature == "") {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case ch <- StreamEvent{Reasoning: event.Text}:
|
||||
case <-ctx.Done():
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{Reasoning: event.Text}) {
|
||||
stopped = true
|
||||
}
|
||||
},
|
||||
|
|
@ -226,7 +222,6 @@ func (b *KiroBackend) runStream(ctx context.Context, req *kiroGenerateRequest, c
|
|||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
ch <- StreamEvent{Done: true, StopReason: ctx.Err().Error()}
|
||||
return
|
||||
case <-time.After(wait):
|
||||
}
|
||||
|
|
@ -234,13 +229,13 @@ func (b *KiroBackend) runStream(ctx context.Context, req *kiroGenerateRequest, c
|
|||
}
|
||||
if throttleAttempts == 0 && b.shouldRefresh(streamErr) {
|
||||
if refreshErr := b.authSource.Refresh(ctx); refreshErr != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("token refresh: %v", refreshErr)}
|
||||
sendStreamEvent(ctx, ch, StreamEvent{Done: true, StopReason: fmt.Sprintf("token refresh: %v", refreshErr)})
|
||||
return
|
||||
}
|
||||
throttleAttempts++ // use throttleAttempts as general attempt counter for auth retry
|
||||
continue
|
||||
}
|
||||
ch <- StreamEvent{Done: true, StopReason: streamErr.Error()}
|
||||
sendStreamEvent(ctx, ch, StreamEvent{Done: true, StopReason: streamErr.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
|
|
@ -255,7 +250,9 @@ func (b *KiroBackend) runStream(ctx context.Context, req *kiroGenerateRequest, c
|
|||
})
|
||||
}
|
||||
}
|
||||
ch <- final
|
||||
if !sendStreamEvent(ctx, ch, final) {
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -223,7 +223,7 @@ func (b *OllamaBackend) ChatStream(ctx context.Context, messages []Message, tool
|
|||
|
||||
// streamOllamaNDJSON reads Ollama's newline-delimited JSON stream from r
|
||||
// and sends StreamEvents to ch.
|
||||
func streamOllamaNDJSON(r io.Reader, ch chan<- StreamEvent) {
|
||||
func streamOllamaNDJSON(ctx context.Context, r io.Reader, ch chan<- StreamEvent) {
|
||||
scanner := newStreamScanner(r)
|
||||
|
||||
var accumulated []ToolCall
|
||||
|
|
@ -236,7 +236,9 @@ func streamOllamaNDJSON(r io.Reader, ch chan<- StreamEvent) {
|
|||
|
||||
var wire ollamaChatResponse
|
||||
if err := json.Unmarshal(line, &wire); err != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream decode: %v", err)}
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{Done: true, StopReason: fmt.Sprintf("stream decode: %v", err)}) {
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
|
|
@ -252,20 +254,24 @@ func streamOllamaNDJSON(r io.Reader, ch chan<- StreamEvent) {
|
|||
if wire.DoneReason != "" {
|
||||
stopReason = wire.DoneReason
|
||||
}
|
||||
ch <- StreamEvent{
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{
|
||||
Done: true,
|
||||
StopReason: stopReason,
|
||||
Usage: Usage{InputTokens: wire.PromptEvalCount, OutputTokens: wire.EvalCount},
|
||||
ToolCalls: accumulated,
|
||||
}) {
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
if wire.Message.Content != "" {
|
||||
ch <- StreamEvent{Content: wire.Message.Content}
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{Content: wire.Message.Content}) {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream read: %v", err)}
|
||||
sendStreamEvent(ctx, ch, StreamEvent{Done: true, StopReason: fmt.Sprintf("stream read: %v", err)})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
package backend
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
|
|
@ -10,7 +11,7 @@ func collectOllama(ndjson string) []StreamEvent {
|
|||
ch := make(chan StreamEvent, 32)
|
||||
go func() {
|
||||
defer close(ch)
|
||||
streamOllamaNDJSON(strings.NewReader(ndjson), ch)
|
||||
streamOllamaNDJSON(context.Background(), strings.NewReader(ndjson), ch)
|
||||
}()
|
||||
var evs []StreamEvent
|
||||
for ev := range ch {
|
||||
|
|
@ -129,7 +130,7 @@ func TestOllamaNDJSON_ScannerError(t *testing.T) {
|
|||
ch := make(chan StreamEvent, 32)
|
||||
go func() {
|
||||
defer close(ch)
|
||||
streamOllamaNDJSON(r, ch)
|
||||
streamOllamaNDJSON(context.Background(), r, ch)
|
||||
}()
|
||||
var evs []StreamEvent
|
||||
for ev := range ch {
|
||||
|
|
|
|||
|
|
@ -379,7 +379,7 @@ func parseDSMLToolCalls(s string) []ToolCall {
|
|||
// streamOpenAISSE reads an OpenAI-format SSE stream from r and sends
|
||||
// StreamEvents to ch. It is the core parsing loop, separated from HTTP
|
||||
// transport so it can be tested with an io.Reader directly.
|
||||
func streamOpenAISSE(r io.Reader, ch chan<- StreamEvent) {
|
||||
func streamOpenAISSE(ctx context.Context, r io.Reader, ch chan<- StreamEvent) {
|
||||
scanner := newStreamScanner(r)
|
||||
|
||||
// Accumulate tool call fragments by index (OpenAI sends args incrementally).
|
||||
|
|
@ -420,7 +420,9 @@ func streamOpenAISSE(r io.Reader, ch chan<- StreamEvent) {
|
|||
if seenDSML && len(ev.ToolCalls) == 0 {
|
||||
ev.ToolCalls = append(ev.ToolCalls, parseDSMLToolCalls(dsmlBuf.String())...)
|
||||
}
|
||||
ch <- ev
|
||||
if !sendStreamEvent(ctx, ch, ev) {
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
if !strings.HasPrefix(line, "data: ") {
|
||||
|
|
@ -429,7 +431,9 @@ func streamOpenAISSE(r io.Reader, ch chan<- StreamEvent) {
|
|||
|
||||
var wire openAIStreamResponse
|
||||
if err := json.Unmarshal([]byte(line[6:]), &wire); err != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream decode: %v", err)}
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{Done: true, StopReason: fmt.Sprintf("stream decode: %v", err)}) {
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
|
|
@ -454,16 +458,20 @@ func streamOpenAISSE(r io.Reader, ch chan<- StreamEvent) {
|
|||
finishReason = choice.FinishReason
|
||||
}
|
||||
if r := choice.Delta.Reasoning; r != "" {
|
||||
ch <- StreamEvent{Reasoning: r}
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{Reasoning: r}) {
|
||||
return
|
||||
}
|
||||
} else if r := choice.Delta.ReasoningContent; r != "" {
|
||||
ch <- StreamEvent{Reasoning: r}
|
||||
if !sendStreamEvent(ctx, ch, StreamEvent{Reasoning: r}) {
|
||||
return
|
||||
}
|
||||
}
|
||||
if content := choice.Delta.Content; content != "" {
|
||||
if seenDSML || strings.Contains(content, "<|DSML|") {
|
||||
seenDSML = true
|
||||
dsmlBuf.WriteString(content)
|
||||
} else {
|
||||
ch <- StreamEvent{Content: content}
|
||||
} else if !sendStreamEvent(ctx, ch, StreamEvent{Content: content}) {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -483,7 +491,7 @@ func streamOpenAISSE(r io.Reader, ch chan<- StreamEvent) {
|
|||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream read: %v", err)}
|
||||
sendStreamEvent(ctx, ch, StreamEvent{Done: true, StopReason: fmt.Sprintf("stream read: %v", err)})
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
package backend
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
|
|
@ -11,7 +12,7 @@ func collect(sse string) []StreamEvent {
|
|||
ch := make(chan StreamEvent, 32)
|
||||
go func() {
|
||||
defer close(ch)
|
||||
streamOpenAISSE(strings.NewReader(sse), ch)
|
||||
streamOpenAISSE(context.Background(), strings.NewReader(sse), ch)
|
||||
}()
|
||||
var evs []StreamEvent
|
||||
for ev := range ch {
|
||||
|
|
@ -255,7 +256,7 @@ func TestSSE_ScannerError(t *testing.T) {
|
|||
ch := make(chan StreamEvent, 32)
|
||||
go func() {
|
||||
defer close(ch)
|
||||
streamOpenAISSE(r, ch)
|
||||
streamOpenAISSE(context.Background(), r, ch)
|
||||
}()
|
||||
var evs []StreamEvent
|
||||
for ev := range ch {
|
||||
|
|
|
|||
Loading…
Reference in New Issue