From cfc69f6f703b3549aaa60f8ce998ac6b13c29d20 Mon Sep 17 00:00:00 2001 From: Viet Tran Date: Fri, 3 Apr 2026 19:38:51 +0700 Subject: [PATCH] fix(tracing): use per-request model/provider overrides in trace spans (#667) (#668) Trace spans and cost calculations were recording the agent's default model/provider instead of the effective overridden values from heartbeat requests. Introduced functional options pattern (spanOption) so all span emitters resolve the correct model/provider. --- internal/agent/loop.go | 6 ++--- internal/agent/loop_run.go | 9 ++++++- internal/agent/loop_tracing.go | 46 +++++++++++++++++++++++++++------- 3 files changed, 48 insertions(+), 13 deletions(-) diff --git a/internal/agent/loop.go b/internal/agent/loop.go index 4e5ee449..345550dd 100644 --- a/internal/agent/loop.go +++ b/internal/agent/loop.go @@ -290,7 +290,7 @@ func (l *Loop) runLoop(ctx context.Context, req RunRequest) (result *RunResult, callCtx = providers.WithReasoningDecision(callCtx, reasoningDecision) } llmSpanStart := time.Now().UTC() - llmSpanID := l.emitLLMSpanStart(callCtx, llmSpanStart, rs.iteration, messages) + llmSpanID := l.emitLLMSpanStart(callCtx, llmSpanStart, rs.iteration, messages, withModel(model), withProvider(provider.Name())) if req.Stream { resp, err = provider.ChatStream(callCtx, chatReq, func(chunk providers.StreamChunk) { @@ -316,11 +316,11 @@ func (l *Loop) runLoop(ctx context.Context, req RunRequest) (result *RunResult, } if err != nil { - l.emitLLMSpanEnd(callCtx, llmSpanID, llmSpanStart, nil, err) + l.emitLLMSpanEnd(callCtx, llmSpanID, llmSpanStart, nil, err, withModel(model), withProvider(provider.Name())) return nil, fmt.Errorf("LLM call failed (iteration %d): %w", rs.iteration, err) } - l.emitLLMSpanEnd(callCtx, llmSpanID, llmSpanStart, resp, nil) + l.emitLLMSpanEnd(callCtx, llmSpanID, llmSpanStart, resp, nil, withModel(model), withProvider(provider.Name())) // For non-streaming responses, emit thinking and content as single events if !req.Stream { diff --git a/internal/agent/loop_run.go b/internal/agent/loop_run.go index 4ff8603f..37f80372 100644 --- a/internal/agent/loop_run.go +++ b/internal/agent/loop_run.go @@ -120,7 +120,14 @@ func (l *Loop) Run(ctx context.Context, req RunRequest) (*RunResult, error) { // Emit running agent span immediately so it's visible in the trace UI. if agentSpanID != uuid.Nil { - l.emitAgentSpanStart(ctx, agentSpanID, runStart, req.Message) + var agentSpanOpts []spanOption + if req.ModelOverride != "" { + agentSpanOpts = append(agentSpanOpts, withModel(req.ModelOverride)) + } + if req.ProviderOverride != nil { + agentSpanOpts = append(agentSpanOpts, withProvider(req.ProviderOverride.Name())) + } + l.emitAgentSpanStart(ctx, agentSpanID, runStart, req.Message, agentSpanOpts...) } // Child trace (announce run): set parent trace back to "running" while diff --git a/internal/agent/loop_tracing.go b/internal/agent/loop_tracing.go index 54be372a..4aec88a6 100644 --- a/internal/agent/loop_tracing.go +++ b/internal/agent/loop_tracing.go @@ -31,6 +31,31 @@ func (l *Loop) Model() string { return l.model } // IsRunning returns whether the agent is currently processing. func (l *Loop) IsRunning() bool { return l.activeRuns.Load() > 0 } +// --------------------------------------------------------------------------- +// Span options — functional options for overriding model/provider in spans. +// --------------------------------------------------------------------------- + +// spanOption overrides span metadata (model, provider) when per-request +// overrides are active (e.g. heartbeat with a cheaper model). +type spanOption func(*spanOverrides) + +type spanOverrides struct { + model string + provider string +} + +func withModel(m string) spanOption { return func(o *spanOverrides) { o.model = m } } +func withProvider(p string) spanOption { return func(o *spanOverrides) { o.provider = p } } + +// resolveSpan returns (model, provider) applying any overrides on top of agent defaults. +func (l *Loop) resolveSpan(opts []spanOption) (string, string) { + o := spanOverrides{model: l.model, provider: l.provider.Name()} + for _, fn := range opts { + fn(&o) + } + return o.model, o.provider +} + // --------------------------------------------------------------------------- // Two-phase LLM span: start (running) + end (completed/error) // --------------------------------------------------------------------------- @@ -38,24 +63,25 @@ func (l *Loop) IsRunning() bool { return l.activeRuns.Load() > 0 } // emitLLMSpanStart emits a "running" LLM span before the LLM call begins. // Returns the span ID so the caller can later call emitLLMSpanEnd to finalize it. // Goroutine-safe: only reads immutable Loop fields and does a channel send. -func (l *Loop) emitLLMSpanStart(ctx context.Context, start time.Time, iteration int, messages []providers.Message) uuid.UUID { +func (l *Loop) emitLLMSpanStart(ctx context.Context, start time.Time, iteration int, messages []providers.Message, opts ...spanOption) uuid.UUID { collector := tracing.CollectorFromContext(ctx) traceID := tracing.TraceIDFromContext(ctx) if collector == nil || traceID == uuid.Nil { return uuid.Nil } + model, providerName := l.resolveSpan(opts) spanID := store.GenNewID() span := store.SpanData{ ID: spanID, TraceID: traceID, SpanType: store.SpanTypeLLMCall, - Name: fmt.Sprintf("%s/%s #%d", l.provider.Name(), l.model, iteration), + Name: fmt.Sprintf("%s/%s #%d", providerName, model, iteration), StartTime: start, Status: store.SpanStatusRunning, Level: store.SpanLevelDefault, - Model: l.model, - Provider: l.provider.Name(), + Model: model, + Provider: providerName, CreatedAt: start, } if parentID := tracing.ParentSpanIDFromContext(ctx); parentID != uuid.Nil { @@ -96,7 +122,7 @@ func (l *Loop) emitLLMSpanStart(ctx context.Context, start time.Time, iteration // emitLLMSpanEnd finalizes a running LLM span with results. // Uses EmitSpanUpdate (channel send) — does NOT depend on ctx being alive, // so it works correctly even after ctx cancellation or deadline exceeded. -func (l *Loop) emitLLMSpanEnd(ctx context.Context, spanID uuid.UUID, start time.Time, resp *providers.ChatResponse, callErr error) { +func (l *Loop) emitLLMSpanEnd(ctx context.Context, spanID uuid.UUID, start time.Time, resp *providers.ChatResponse, callErr error, opts ...spanOption) { if spanID == uuid.Nil { return // tracing disabled — no running span was emitted } @@ -139,7 +165,8 @@ func (l *Loop) emitLLMSpanEnd(ctx context.Context, spanID uuid.UUID, start time. } } // Calculate cost if pricing config is available. - if pricing := tracing.LookupPricing(l.modelPricing, l.provider.Name(), l.model); pricing != nil { + model, providerName := l.resolveSpan(opts) + if pricing := tracing.LookupPricing(l.modelPricing, providerName, model); pricing != nil { cost := tracing.CalculateCost(pricing, resp.Usage) if cost > 0 { updates["total_cost"] = cost @@ -282,7 +309,7 @@ func (l *Loop) emitToolSpanEnd(ctx context.Context, spanID uuid.UUID, start time // emitAgentSpanStart emits a "running" root agent span at the beginning of a run. // The span is identified by agentSpanID (pre-generated, same ID used as ParentSpanID // for child LLM/tool spans). -func (l *Loop) emitAgentSpanStart(ctx context.Context, agentSpanID uuid.UUID, start time.Time, inputPreview string) { +func (l *Loop) emitAgentSpanStart(ctx context.Context, agentSpanID uuid.UUID, start time.Time, inputPreview string, opts ...spanOption) { collector := tracing.CollectorFromContext(ctx) traceID := tracing.TraceIDFromContext(ctx) if collector == nil || traceID == uuid.Nil { @@ -291,6 +318,7 @@ func (l *Loop) emitAgentSpanStart(ctx context.Context, agentSpanID uuid.UUID, st previewLimit := previewLimitForVerbose(collector.Verbose()) + model, providerName := l.resolveSpan(opts) spanName := l.id span := store.SpanData{ ID: agentSpanID, @@ -300,8 +328,8 @@ func (l *Loop) emitAgentSpanStart(ctx context.Context, agentSpanID uuid.UUID, st StartTime: start, Status: store.SpanStatusRunning, Level: store.SpanLevelDefault, - Model: l.model, - Provider: l.provider.Name(), + Model: model, + Provider: providerName, InputPreview: tracing.TruncateMid(inputPreview, previewLimit), CreatedAt: start, }