package workbench import ( "context" "fmt" "net/http" "strings" "aigateway.local/core/internal/apikey" "aigateway.local/core/internal/gateway" tracepkg "aigateway.local/core/internal/trace" ) func (h *RuntimeHTTPHandler) beginTrace(ctx context.Context, principal apikey.Principal, traceType, targetID, targetCode, conversationID string) string { if h.traces == nil { return "" } conversationID = strings.TrimSpace(conversationID) input := tracepkg.StartInput{RequestID: gateway.RequestID(ctx), APIKeyID: principal.APIKeyID, TenantID: principal.TenantID, TraceType: traceType, TargetID: targetID, TargetCode: targetCode, ConversationID: conversationID} item, err := h.traces.Start(context.WithoutCancel(ctx), input) if err != nil { h.logger.Warn("llm trace start failed", "request_id", input.RequestID, "target", targetCode, "error", err) return "" } return item.ID } func (h *RuntimeHTTPHandler) finishTrace(ctx context.Context, traceID, status, errorText string, retrievalCount, modelCallCount, toolCallCount int) { if h.traces == nil || traceID == "" { return } if err := h.traces.Finish(context.WithoutCancel(ctx), traceID, tracepkg.FinishInput{Status: status, Error: errorText, RetrievalCount: retrievalCount, ModelCallCount: modelCallCount, ToolCallCount: toolCallCount}); err != nil { h.logger.Warn("llm trace finish failed", "trace_id", traceID, "error", err) } } func (h *RuntimeHTTPHandler) callGatewayWithTrace(original *http.Request, payload map[string]any, traceID string, round int) (int, http.Header, map[string]any, error) { spanID := "" model, _ := payload["model"].(string) if h.traces != nil && traceID != "" { span, err := h.traces.StartSpan(context.WithoutCancel(original.Context()), tracepkg.SpanInput{TraceID: traceID, SpanType: "model", Name: "chat.completions", Round: round, Model: model, Metadata: map[string]any{"endpoint": "/v1/chat/completions"}}) if err != nil { h.logger.Warn("llm model span start failed", "trace_id", traceID, "error", err) } else { spanID = span.ID } } statusCode, headers, response, callErr := h.callGateway(original, payload) if spanID != "" { inputTokens, outputTokens := responseUsage(response) spanStatus := "success" if callErr != nil || statusCode < 200 || statusCode >= 300 { spanStatus = "error" } metadata := map[string]any{"http_status": statusCode, "round": round} providerCode := "" spanModel := model if provider := headers.Get("X-Gateway-Provider"); provider != "" { providerCode = provider metadata["provider"] = provider } if resolvedModel := strings.TrimSpace(headers.Get("X-Gateway-Model")); resolvedModel != "" { spanModel = resolvedModel } if err := h.traces.FinishSpan(context.WithoutCancel(original.Context()), spanID, tracepkg.SpanFinishInput{Status: spanStatus, Error: errorString(callErr), InputTokens: inputTokens, OutputTokens: outputTokens, ProviderCode: providerCode, Model: spanModel, Metadata: metadata}); err != nil { h.logger.Warn("llm model span finish failed", "span_id", spanID, "error", err) } } return statusCode, headers, response, callErr } func (h *RuntimeHTTPHandler) executeToolWithTrace(ctx context.Context, traceID, name, callID string, round int, execute func() (map[string]any, error)) (map[string]any, error) { spanID := "" if h.traces != nil && traceID != "" { span, err := h.traces.StartSpan(context.WithoutCancel(ctx), tracepkg.SpanInput{TraceID: traceID, SpanType: "tool", Name: name, Round: round, Metadata: map[string]any{"tool_call_id": callID}}) if err != nil { h.logger.Warn("llm tool span start failed", "trace_id", traceID, "tool", name, "error", err) } else { spanID = span.ID } } result, executeErr := execute() if spanID != "" { status := "success" if executeErr != nil { status = "error" } if err := h.traces.FinishSpan(context.WithoutCancel(ctx), spanID, tracepkg.SpanFinishInput{Status: status, Error: errorString(executeErr), Metadata: map[string]any{"tool_call_id": callID}}); err != nil { h.logger.Warn("llm tool span finish failed", "span_id", spanID, "error", err) } } return result, executeErr } func (h *RuntimeHTTPHandler) searchWithTrace(ctx context.Context, traceID, knowledgeBaseID, query string, topK int) ([]SearchHit, error) { spanID := "" if h.traces != nil && traceID != "" { span, err := h.traces.StartSpan(context.WithoutCancel(ctx), tracepkg.SpanInput{TraceID: traceID, SpanType: "retrieval", Name: "knowledge.search", Metadata: map[string]any{"knowledge_base_id": knowledgeBaseID, "top_k": topK}}) if err != nil { h.logger.Warn("llm retrieval span start failed", "trace_id", traceID, "error", err) } else { spanID = span.ID } } hits, searchErr := h.retriever.Search(ctx, knowledgeBaseID, query, topK) if spanID != "" { status := "success" if searchErr != nil { status = "error" } metadata := map[string]any{"knowledge_base_id": knowledgeBaseID, "hit_count": len(hits)} if err := h.traces.FinishSpan(context.WithoutCancel(ctx), spanID, tracepkg.SpanFinishInput{Status: status, Error: errorString(searchErr), Metadata: metadata}); err != nil { h.logger.Warn("llm retrieval span finish failed", "span_id", spanID, "error", err) } } return hits, searchErr } func responseUsage(response map[string]any) (int64, int64) { if response == nil { return 0, 0 } usage, _ := response["usage"].(map[string]any) return numberValue(usage["prompt_tokens"], usage["input_tokens"]), numberValue(usage["completion_tokens"], usage["output_tokens"]) } func numberValue(values ...any) int64 { for _, value := range values { switch number := value.(type) { case float64: return int64(number) case float32: return int64(number) case int: return int64(number) case int64: return number } } return 0 } func errorString(err error) string { if err == nil { return "" } return fmt.Sprint(err) }