Files
superidou 9501751792 0.10.1: 安全与业务逻辑加固、新品牌与部署加固
三轮审查修复(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 币种维度)
2026-08-13 10:50:51 +08:00

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)
}