diff --git a/pkg/llm/llm_test.go b/pkg/llm/llm_test.go index d2f64b7ba..ac8e6c0c3 100644 --- a/pkg/llm/llm_test.go +++ b/pkg/llm/llm_test.go @@ -48,11 +48,12 @@ func (m *mockProvider) ChatCompletionStream(_ context.Context, _ *llm.ChatComple } type mockStream struct { - events []llm.ChatCompletionStreamEvent - idx int - current llm.ChatCompletionStreamEvent - err error - closed bool + events []llm.ChatCompletionStreamEvent + idx int + current llm.ChatCompletionStreamEvent + err error + closeErr error + closed bool } func (s *mockStream) Next() bool { @@ -66,7 +67,7 @@ func (s *mockStream) Next() bool { func (s *mockStream) Event() llm.ChatCompletionStreamEvent { return s.current } func (s *mockStream) Err() error { return s.err } -func (s *mockStream) Close() error { s.closed = true; return nil } +func (s *mockStream) Close() error { s.closed = true; return s.closeErr } func newTestClient(provider llm.Provider) (*llm.Client, *tracetest.SpanRecorder) { recorder := tracetest.NewSpanRecorder() @@ -459,6 +460,33 @@ func TestChatCompletionStream(t *testing.T) { require.Len(t, spans, 1, "span should be ended by Close even without exhausting stream") }) + t.Run("close error traced on span", func(t *testing.T) { + t.Parallel() + + ms := &mockStream{ + events: []llm.ChatCompletionStreamEvent{ + {Delta: llm.MessageDelta{Content: "partial"}}, + }, + closeErr: errors.New("broken pipe"), + } + + client, recorder := newTestClient(&mockProvider{streamResp: ms}) + stream, err := client.ChatCompletionStream(context.Background(), &llm.ChatCompletionRequest{ + Model: "test-model", + Messages: []llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "Hi"}}}}, + }) + require.NoError(t, err) + + require.True(t, stream.Next()) + closeErr := stream.Close() + require.ErrorContains(t, closeErr, "broken pipe") + + spans := recorder.Ended() + require.Len(t, spans, 1) + assert.Equal(t, codes.Error, spans[0].Status().Code) + assert.Contains(t, spans[0].Status().Description, "broken pipe") + }) + t.Run("stream span records usage and finish reason", func(t *testing.T) { t.Parallel() diff --git a/pkg/llm/trace.go b/pkg/llm/trace.go index a1b6bd373..7e285171d 100644 --- a/pkg/llm/trace.go +++ b/pkg/llm/trace.go @@ -112,8 +112,13 @@ func (s *tracedStream) Err() error { } func (s *tracedStream) Close() error { + err := s.inner.Close() + if err != nil { + s.span.RecordError(err) + s.span.SetStatus(codes.Error, err.Error()) + } s.finalizeSpan() - return s.inner.Close() + return err } func (s *tracedStream) finalizeSpan() {