9501751792
三轮审查修复(60+ 项),相对远端 main(b536672)的关键变更:
- 安全: 数据面 SSRF 拨号防护(防 DNS rebinding)/上游凭据剥离/登录防枚举
与锁定态统一/可信代理(X-Forwarded-For)限流加固/会话版本失效机制/
撤销即时传播/弱密钥拒绝启动/脱敏字节级重写(保签名契约)
- 业务逻辑: 裸 body 上传 panic/bootstrap 审计管线卡死/定价通配符优先级/
全局工具可见性/调度器停机补跑/TOTP 挑战令牌消费顺序/熔断探针语义/
>4MB 响应 token 计量/管理员重置密码作废会话 等
- 前端: 新 logo(语枢 AI 网关主题)/Provider 凭据异常警示/删除入口/
后端错误消息透传/localStorage 敏感数据收敛
- 部署: CREDENTIAL_MASTER_KEY 持久化与弱值拒绝/Provider DELETE 接口/
nginx 安全头/worker 内存限制
- 新增迁移 000029(key_hash 索引)/000030(usage_daily 币种维度)
149 lines
5.8 KiB
Go
149 lines
5.8 KiB
Go
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)
|
|
}
|