Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 32 additions & 9 deletions pkg/tracing/traced_llm.go
Original file line number Diff line number Diff line change
Expand Up @@ -287,12 +287,20 @@ func (m *TracedLLM) GenerateWithToolsStream(ctx context.Context, prompt string,
// Start a goroutine to proxy events and handle span completion
go func() {
defer close(wrappedChan)

var streamErr error

defer func() {
// When streaming is complete, create tool call spans and end main span
endTime := time.Now()
duration := endTime.Sub(startTime)
span.SetAttribute("duration_ms", duration.Milliseconds())

// Record streaming error on span if one was captured (#328)
if streamErr != nil {
span.RecordError(streamErr)
}

// Get tool calls from context and create spans using TraceGeneration if any exist
toolCalls := GetToolCallsFromContext(ctx)

Expand All @@ -306,25 +314,40 @@ func (m *TracedLLM) GenerateWithToolsStream(ctx context.Context, prompt string,
model = streamingLLM.Name()
}

// Build metadata, including error if the stream failed mid-flight (#328)
metadata := map[string]any{
"streaming": true,
"tools": len(tools),
}
if streamErr != nil {
metadata["error"] = streamErr.Error()
}

responseText := "streaming_response"
if streamErr != nil {
if m.shouldIncludeContent() {
responseText = "<error>"
} else {
responseText = "<redacted>"
}
}

// Create spans using TraceGeneration which handles tool calls correctly
if adapter, ok := m.tracer.(*OTELTracerAdapter); ok {
_, _ = adapter.otelTracer.TraceGeneration(ctx, model, prompt, "streaming_response", startTime, endTime, map[string]any{ //nolint:gosec
"streaming": true,
"tools": len(tools),
})
_, _ = adapter.otelTracer.TraceGeneration(ctx, model, prompt, responseText, startTime, endTime, metadata) //nolint:gosec
} else if tracer, ok := m.tracer.(*OTELLangfuseTracer); ok {
_, _ = tracer.TraceGeneration(ctx, model, prompt, "streaming_response", startTime, endTime, map[string]any{ //nolint:gosec
"streaming": true,
"tools": len(tools),
})
_, _ = tracer.TraceGeneration(ctx, model, prompt, responseText, startTime, endTime, metadata) //nolint:gosec
}
}

span.End()
}()

// Proxy all events from the original channel
// Proxy all events from the original channel, capturing any error event (#328)
for event := range originalChan {
if event.Type == interfaces.StreamEventError && event.Error != nil {
streamErr = event.Error
}
wrappedChan <- event
}
}()
Expand Down
142 changes: 142 additions & 0 deletions pkg/tracing/traced_llm_error_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,142 @@
package tracing

import (
"context"
"errors"
"testing"
"time"

"github.com/Ingenimax/agent-sdk-go/pkg/interfaces"
sdktrace "go.opentelemetry.io/otel/sdk/trace"
"go.opentelemetry.io/otel/sdk/trace/tracetest"
)

type errorStreamLLM struct {
streamErr error
}

func (*errorStreamLLM) Generate(context.Context, string, ...interfaces.GenerateOption) (string, error) {
return "", nil
}

func (*errorStreamLLM) GenerateWithTools(context.Context, string, []interfaces.Tool, ...interfaces.GenerateOption) (string, error) {
return "", nil
}

func (*errorStreamLLM) GenerateDetailed(context.Context, string, ...interfaces.GenerateOption) (*interfaces.LLMResponse, error) {
return &interfaces.LLMResponse{}, nil
}

func (*errorStreamLLM) GenerateWithToolsDetailed(context.Context, string, []interfaces.Tool, ...interfaces.GenerateOption) (*interfaces.LLMResponse, error) {
return &interfaces.LLMResponse{}, nil
}

func (*errorStreamLLM) GenerateStream(context.Context, string, ...interfaces.GenerateOption) (<-chan interfaces.StreamEvent, error) {
return nil, nil
}

func (m *errorStreamLLM) GenerateWithToolsStream(ctx context.Context, _ string, _ []interfaces.Tool, _ ...interfaces.GenerateOption) (<-chan interfaces.StreamEvent, error) {
AddToolCallToContext(ctx, ToolCall{Name: "test_tool", Timestamp: time.Now().Format(time.RFC3339)})
ch := make(chan interfaces.StreamEvent, 1)
if m.streamErr != nil {
ch <- interfaces.StreamEvent{Type: interfaces.StreamEventError, Error: m.streamErr}
} else {
ch <- interfaces.StreamEvent{Type: interfaces.StreamEventContentComplete, Content: "complete"}
}
close(ch)
return ch, nil
}

func (*errorStreamLLM) Name() string { return "test" }
func (*errorStreamLLM) SupportsStreaming() bool { return true }

func TestGenerateWithToolsStreamGenerationSpan(t *testing.T) {
for _, tc := range []struct {
name string
includeContent bool
streamErr error
wantError string
wantResponse string
}{
{
name: "error with content included",
includeContent: true,
streamErr: errors.New("stream interrupted"),
wantError: "stream interrupted",
wantResponse: "<error>",
},
{
name: "error with content redacted",
includeContent: false,
streamErr: errors.New("stream interrupted"),
wantError: "stream interrupted",
wantResponse: "<redacted>",
},
{
name: "successful stream preserves response placeholder",
includeContent: false,
wantResponse: "streaming_response",
},
} {
t.Run(tc.name, func(t *testing.T) {
exporter := tracetest.NewInMemoryExporter()
provider := sdktrace.NewTracerProvider(sdktrace.WithSyncer(exporter))
t.Cleanup(func() { _ = provider.Shutdown(context.Background()) })

otelTracer := &OTELLangfuseTracer{
tracerProvider: provider,
tracer: provider.Tracer("test"),
enabled: true,
IncludeContent: tc.includeContent,
}
traced := NewTracedLLM(&errorStreamLLM{streamErr: tc.streamErr}, NewOTELTracerAdapter(otelTracer)).(*TracedLLM)

stream, err := traced.GenerateWithToolsStream(context.Background(), "secret prompt", nil)
if err != nil {
t.Fatalf("GenerateWithToolsStream() error = %v", err)
}
for range stream {
}

var generation sdktrace.ReadOnlySpan
for _, span := range exporter.GetSpans() {
for _, attr := range span.Attributes {
if string(attr.Key) == "langfuse.observation.type" && attr.Value.AsString() == "generation" {
generation = span.Snapshot()
}
}
}
if generation == nil {
t.Fatal("generation span was not exported")
}

attributes := make(map[string]string)
for _, attr := range generation.Attributes() {
attributes[string(attr.Key)] = attr.Value.AsString()
}
if got := attributes["langfuse.observation.metadata.error"]; got != tc.wantError {
t.Errorf("error metadata = %q, want %q", got, tc.wantError)
}
if got := attributes["gen_ai.completion.0.content"]; got != tc.wantResponse {
t.Errorf("generation response = %q, want %q", got, tc.wantResponse)
}

if tc.streamErr != nil {
var foundErrorEvent bool
for _, span := range exporter.GetSpans() {
if span.Name != "llm.generate_with_tools_stream" {
continue
}
for _, event := range span.Events {
if event.Name == "exception" {
foundErrorEvent = true
}
}
}
if !foundErrorEvent {
t.Error("main streaming span did not record the stream error")
}
}
})
}
}