Prevent backend stream goroutine leaks

This commit is contained in:
Ollie Agent 2026-08-17 18:27:35 +02:00
parent 49e43248b0
commit 701e9bcaad
8 changed files with 70 additions and 40 deletions

View File

@ -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)})
}
}

View File

@ -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 {

View File

@ -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{

View File

@ -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
}
}

View File

@ -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)})
}
}

View File

@ -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 {

View File

@ -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)})
}
}

View File

@ -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 {