Files
ai-gateway-go/internal/workbench/runtime_http.go
T
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

717 lines
26 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package workbench
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"strings"
"time"
"aigateway.local/core/internal/apikey"
"aigateway.local/core/internal/factcheck"
"aigateway.local/core/internal/gateway"
"aigateway.local/core/internal/platform/apiresponse"
tracepkg "aigateway.local/core/internal/trace"
)
type RuntimeHTTPHandler struct {
service *Service
tools *ToolService
retriever Retriever
auth apikey.PrincipalAuthenticator
gateway http.Handler
factCheck *factcheck.Engine
traces *tracepkg.Store
logger *slog.Logger
mux *http.ServeMux
market MarketplaceDeps
}
// MarketplaceDeps carries the resource-marketplace services into the runtime
// handler (MCP servers, skills, digital employees, installations).
type MarketplaceDeps struct {
MCPServers *MCPServerService
Skills *SkillService
Employees *DigitalEmployeeService
Market *MarketplaceService
MCPClient *MCPClient
}
func NewRuntimeHTTPHandler(service *Service, tools *ToolService, retriever Retriever, auth apikey.PrincipalAuthenticator, gatewayHandler http.Handler, market MarketplaceDeps) *RuntimeHTTPHandler {
h := &RuntimeHTTPHandler{service: service, tools: tools, retriever: retriever, auth: auth, gateway: gatewayHandler, mux: http.NewServeMux(), market: market}
h.logger = slog.Default()
h.mux.HandleFunc("GET /v1/prompts", h.listPrompts)
h.mux.HandleFunc("POST /v1/prompts/{name}/render", h.renderPrompt)
h.mux.HandleFunc("POST /v1/knowledge/search", h.searchKnowledge)
h.mux.HandleFunc("POST /v1/knowledge/{id}/search", h.searchKnowledge)
h.mux.HandleFunc("GET /v1/tools", h.listTools)
h.mux.HandleFunc("POST /v1/tools/{code}/invoke", h.invokeTool)
h.mux.HandleFunc("POST /v1/applications/{code}/chat/completions", h.runApplication)
h.mux.HandleFunc("POST /v1/skills/{code}/render", h.renderSkill)
h.mux.HandleFunc("GET /v1/mcp-servers", h.listMCPServers)
h.mux.HandleFunc("GET /v1/mcp-servers/{code}/tools", h.mcpServerTools)
h.mux.HandleFunc("POST /v1/mcp-servers/{code}/tools/{tool}/invoke", h.invokeMCPTool)
h.mux.HandleFunc("POST /v1/digital-employees/{code}/chat/completions", h.runDigitalEmployee)
return h
}
func (h *RuntimeHTTPHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { h.mux.ServeHTTP(w, r) }
// SetLogger wires a logger for best-effort diagnostics (fact-check skips etc.).
func (h *RuntimeHTTPHandler) SetLogger(logger *slog.Logger) {
if logger != nil {
h.logger = logger
}
}
// SetFactCheckEngine enables post-answer fact-checking on application
// conversations. When nil (the default) fact-checking is skipped entirely.
func (h *RuntimeHTTPHandler) SetFactCheckEngine(engine *factcheck.Engine) { h.factCheck = engine }
// SetTraceStore enables metadata-only LLM Trace recording for application and
// digital-employee runs. Trace persistence is best effort and never changes
// the runtime response when the database is unavailable.
func (h *RuntimeHTTPHandler) SetTraceStore(store *tracepkg.Store) { h.traces = store }
// factCheckRetriever adapts the workbench Retriever to the fact-check engine's
// EvidenceRetriever interface, reusing the same knowledge-base search path that
// application prompts already use.
type factCheckRetriever struct{ inner Retriever }
func NewFactCheckRetriever(inner Retriever) *factCheckRetriever {
if inner == nil {
return nil
}
return &factCheckRetriever{inner: inner}
}
func (a *factCheckRetriever) Search(ctx context.Context, knowledgeBaseID, query string, topK int) ([]factcheck.EvidenceHit, error) {
hits, err := a.inner.Search(ctx, knowledgeBaseID, query, topK)
if err != nil {
return nil, err
}
out := make([]factcheck.EvidenceHit, 0, len(hits))
for _, hit := range hits {
out = append(out, factcheck.EvidenceHit{DocumentTitle: hit.DocumentTitle, Content: hit.Content})
}
return out, nil
}
// VerifyFactCheck satisfies factcheck.Verifier. It routes a non-streaming chat
// completion through the same governed gateway, reusing the original request's
// credential headers so the fact-check call is authenticated and rate-limited
// exactly like the application call that produced the answer.
func (h *RuntimeHTTPHandler) VerifyFactCheck(ctx context.Context, original *http.Request, model, system, user string, timeout time.Duration) (string, error) {
ctx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()
payload := map[string]any{
"model": model,
"temperature": 0,
"stream": false,
"messages": []map[string]any{
{"role": "system", "content": system},
{"role": "user", "content": user},
},
}
raw, _ := json.Marshal(payload)
request := original.Clone(ctx)
request.Method = http.MethodPost
request.URL.Path = "/v1/chat/completions"
request.URL.RawPath = ""
request.Body = ioNopCloser{bytes.NewReader(raw)}
request.ContentLength = int64(len(raw))
request.Header = request.Header.Clone()
request.Header.Set("Content-Type", "application/json")
recorder := newBoundedRecorder()
h.gateway.ServeHTTP(recorder, request)
if recorder.overrun > 0 {
return "", errors.New("事实核查响应超过 2MB 上限")
}
result := recorder.Result()
defer result.Body.Close()
var decoded map[string]any
if json.NewDecoder(result.Body).Decode(&decoded) != nil {
return "", errors.New("事实核查响应无法解析")
}
if result.StatusCode < 200 || result.StatusCode >= 300 {
return "", fmt.Errorf("事实核查模型调用失败(HTTP %d", result.StatusCode)
}
content, _ := firstChoiceMessage(decoded)["content"].(string)
if strings.TrimSpace(content) == "" {
return "", errors.New("事实核查模型未返回文本")
}
return content, nil
}
func (h *RuntimeHTTPHandler) principal(w http.ResponseWriter, r *http.Request) (apikey.Principal, bool) {
secret := strings.TrimSpace(r.Header.Get("X-Gateway-API-Key"))
if secret == "" {
value := strings.TrimSpace(r.Header.Get("Authorization"))
if strings.HasPrefix(strings.ToLower(value), "bearer ") {
secret = strings.TrimSpace(value[7:])
}
}
principal, err := h.auth.AuthenticatePrincipal(r.Context(), secret)
if err != nil {
apiresponse.Error(w, 401, "API Key 无效或已过期")
return principal, false
}
return principal, true
}
func visible(departments []string, principal apikey.Principal, secure bool) bool {
// fail-closed:无 APIKeyID 的匿名主体不视为"可见一切"。
// 今天认证器总是返回 bootstrap 或真实 key ID,但任何未来认证路径的
// 变化都不应静默放开所有部门作用域资产。
if principal.APIKeyID == "" {
return false
}
// 无部门限定的资源是全局资源:所有已认证主体可见。secure 只标记
// "执行敏感能力"类资源,不改变可见性规则——否则全局工具/MCP 对
// 所有人不可见,绑定它们的应用会在运行时失败。
if len(departments) == 0 {
return true
}
if principal.TenantID == nil {
return false
}
for _, id := range departments {
if id == *principal.TenantID {
return true
}
}
return false
}
func (h *RuntimeHTTPHandler) listPrompts(w http.ResponseWriter, r *http.Request) {
principal, ok := h.principal(w, r)
if !ok {
return
}
items, err := h.service.ListPrompts(r.Context())
if err != nil {
apiresponse.Error(w, 503, "Prompt 服务暂不可用")
return
}
result := []map[string]any{}
for _, item := range items {
if item.Enabled && item.Current != nil && visible(item.DepartmentIDs, principal, false) {
result = append(result, map[string]any{"name": item.Name, "description": item.Description, "tags": item.Tags, "version": item.Current.Version, "variables": item.Current.Variables})
}
}
writeRuntime(w, 200, map[string]any{"object": "list", "data": result})
}
func (h *RuntimeHTTPHandler) renderPrompt(w http.ResponseWriter, r *http.Request) {
principal, ok := h.principal(w, r)
if !ok {
return
}
var input struct {
Variables map[string]any `json:"variables"`
}
if !decodeRuntime(w, r, &input) {
return
}
items, err := h.service.ListPrompts(r.Context())
if err != nil {
runtimeError(w, 503, "Prompt 服务暂不可用")
return
}
var selected *PromptTemplate
for i := range items {
if items[i].Name == r.PathValue("name") && items[i].Enabled && visible(items[i].DepartmentIDs, principal, false) {
selected = &items[i]
break
}
}
if selected == nil || selected.Current == nil {
runtimeError(w, 404, "Prompt 不存在或不可见")
return
}
rendered, err := RenderPrompt(selected.Current.Content, selected.Current.Variables, input.Variables)
if err != nil {
runtimeError(w, 400, err.Error())
return
}
writeRuntime(w, 200, map[string]any{"name": selected.Name, "version": selected.Current.Version, "rendered": rendered})
}
func (h *RuntimeHTTPHandler) searchKnowledge(w http.ResponseWriter, r *http.Request) {
principal, ok := h.principal(w, r)
if !ok {
return
}
var input struct {
KnowledgeBaseID string `json:"knowledge_base_id"`
Query string `json:"query"`
TopK int `json:"top_k"`
}
if !decodeRuntime(w, r, &input) {
return
}
id := r.PathValue("id")
if id == "" {
id = input.KnowledgeBaseID
}
kb, err := h.service.GetKnowledgeBase(r.Context(), id)
if err != nil || !kb.Enabled || !visible(kb.DepartmentIDs, principal, false) {
runtimeError(w, 404, "知识库不存在或不可见")
return
}
hits, err := h.retriever.Search(r.Context(), id, input.Query, input.TopK)
if err != nil {
runtimeError(w, 400, err.Error())
return
}
writeRuntime(w, 200, map[string]any{"knowledge_base_id": id, "results": hits})
}
func (h *RuntimeHTTPHandler) listTools(w http.ResponseWriter, r *http.Request) {
principal, ok := h.principal(w, r)
if !ok {
return
}
items, err := h.tools.List(r.Context())
if err != nil {
runtimeError(w, 503, "工具服务暂不可用")
return
}
result := []map[string]any{}
for _, tool := range items {
if tool.Enabled && visible(tool.DepartmentIDs, principal, true) {
result = append(result, map[string]any{"code": tool.Code, "name": tool.Name, "description": tool.Description, "input_schema": tool.InputSchema})
}
}
writeRuntime(w, 200, map[string]any{"object": "list", "data": result})
}
func (h *RuntimeHTTPHandler) invokeTool(w http.ResponseWriter, r *http.Request) {
principal, ok := h.principal(w, r)
if !ok {
return
}
var input struct {
Input map[string]any `json:"input"`
}
if !decodeRuntime(w, r, &input) {
return
}
tool, err := h.tools.GetByCode(r.Context(), r.PathValue("code"))
if err != nil || !visible(tool.DepartmentIDs, principal, true) {
runtimeError(w, 404, "工具不存在或不可调用")
return
}
result, err := h.tools.Execute(r.Context(), tool, input.Input, principal.APIKeyID, gateway.RequestID(r.Context()))
if err != nil {
runtimeError(w, 502, err.Error())
return
}
writeRuntime(w, 200, result)
}
type applicationRequest struct {
Messages []map[string]any `json:"messages"`
Variables map[string]any `json:"variables"`
}
func (h *RuntimeHTTPHandler) runApplication(w http.ResponseWriter, r *http.Request) {
principal, ok := h.principal(w, r)
if !ok {
return
}
var input applicationRequest
if !decodeRuntime(w, r, &input) {
return
}
app, err := h.service.GetPublishedApplicationByCode(r.Context(), r.PathValue("code"))
if err != nil || app.PublishedConfig == nil || !visible(app.DepartmentIDs, principal, false) {
runtimeError(w, 404, "应用不存在、未发布或不可见")
return
}
started := time.Now()
status := "error"
runError := ""
retrievalCount := 0
toolCount := 0
modelCallCount := 0
conversationID := strings.TrimSpace(r.Header.Get("X-Gateway-Conversation-ID"))
traceID := h.beginTrace(r.Context(), principal, "application", app.ID, app.Code, conversationID)
defer func() {
traceCtx := context.WithoutCancel(r.Context())
h.finishTrace(traceCtx, traceID, status, runError, retrievalCount, modelCallCount, toolCount)
runID, idErr := newUUID()
if idErr == nil {
_, _ = h.service.pool.Exec(traceCtx, `INSERT INTO gateway.application_runs(id,application_id,version,api_key_id,trace_id,request_id,status,latency_ms,retrieval_count,tool_count,error) VALUES($1,$2,nullif($3,0),nullif($4,'')::uuid,nullif($5,'')::uuid,$6,$7,$8,$9,$10,$11)`, runID, app.ID, valueOrZero(app.PublishedVersion), principal.APIKeyID, traceID, gateway.RequestID(r.Context()), status, time.Since(started).Milliseconds(), retrievalCount, toolCount, runError)
}
}()
payload, toolsByCode, prepareErr := h.prepareApplication(r.Context(), app, input, principal, &retrievalCount, traceID)
if prepareErr != nil {
runError = prepareErr.Error()
runtimeError(w, 400, runError)
return
}
config := *app.PublishedConfig
var response map[string]any
var responseHeaders http.Header
var statusCode int
for round := 0; ; round++ {
modelCallCount++
statusCode, responseHeaders, response, err = h.callGatewayWithTrace(r, payload, traceID, round)
if err != nil {
runError = err.Error()
copyHeaders(w.Header(), responseHeaders)
runtimeError(w, statusCode, runError)
return
}
calls := extractToolCalls(response)
if len(calls) == 0 {
break
}
if round >= config.MaxToolRounds {
runError = "工具调用轮次已达上限"
runtimeError(w, 502, runError)
return
}
choice := firstChoiceMessage(response)
payload["messages"] = append(payload["messages"].([]map[string]any), choice)
for _, call := range calls {
tool, exists := toolsByCode[call.Name]
if !exists {
runError = "模型请求了未授权工具 " + call.Name
runtimeError(w, 400, runError)
return
}
var args map[string]any
if json.Unmarshal([]byte(call.Arguments), &args) != nil {
args = map[string]any{}
}
result, executeErr := h.executeToolWithTrace(r.Context(), traceID, call.Name, call.ID, round, func() (map[string]any, error) {
return h.tools.Execute(r.Context(), tool, args, principal.APIKeyID, gateway.RequestID(r.Context()))
})
if executeErr != nil {
runError = executeErr.Error()
runtimeError(w, 502, runError)
return
}
encoded, _ := json.Marshal(result["body"])
payload["messages"] = append(payload["messages"].([]map[string]any), map[string]any{"role": "tool", "tool_call_id": call.ID, "name": call.Name, "content": string(encoded)})
toolCount++
}
}
if h.factCheck != nil {
h.applyFactCheck(r, input, response, app.DepartmentIDs)
}
response["application"] = map[string]any{"code": app.Code, "name": app.Name, "version": valueOrZero(app.PublishedVersion), "retrieval_count": retrievalCount, "tool_calls": toolCount}
status = "success"
copyHeaders(w.Header(), responseHeaders)
writeRuntime(w, statusCode, response)
}
// applyFactCheck verifies the assistant answer against configured knowledge
// bases and applies the policy action. It must never fail the chat: any error
// is logged and the answer is returned unchanged. departments 用于选择
// department:<uuid> 作用域的策略,空列表只应用 global 策略。
func (h *RuntimeHTTPHandler) applyFactCheck(r *http.Request, input applicationRequest, response map[string]any, departments []string) {
answer, _ := assistantAnswer(response)
lastQuestion := lastUserMessage(input.Messages)
if strings.TrimSpace(answer) == "" || strings.TrimSpace(lastQuestion) == "" {
return
}
scope := ""
if len(departments) > 0 {
scope = "department:" + departments[0]
}
verifier := func(ctx context.Context, model, system, user string, timeout time.Duration) (string, error) {
return h.VerifyFactCheck(ctx, r, model, system, user, timeout)
}
event, err := h.factCheck.Check(r.Context(), gateway.RequestID(r.Context()), scope, lastQuestion, answer, factcheck.VerifierFunc(verifier))
if err != nil {
h.logger.Warn("fact-check skipped", "request_id", gateway.RequestID(r.Context()), "error", err)
return
}
if event.ID == "" {
return
}
switch {
case event.Action == "block" && event.Verdict == "unsupported":
overrideAnswer(response, "无法回答:该回复与知识库事实不符,已被事实核查拦截。")
response["fact_check"] = map[string]any{"event_id": event.ID, "verdict": event.Verdict, "support_score": event.SupportScore, "blocked": true}
case event.Action == "annotate":
response["fact_check"] = map[string]any{"event_id": event.ID, "verdict": event.Verdict, "support_score": event.SupportScore, "blocked": false}
}
}
// lastUserMessage returns the content of the last user message in the request.
func lastUserMessage(messages []map[string]any) string {
last := ""
for _, message := range messages {
role, _ := message["role"].(string)
if role != "user" {
continue
}
if content, ok := message["content"].(string); ok {
last = content
}
}
return last
}
// assistantAnswer extracts the final assistant text from a gateway response.
func assistantAnswer(response map[string]any) (string, bool) {
content, _ := firstChoiceMessage(response)["content"].(string)
return content, strings.TrimSpace(content) != ""
}
// overrideAnswer rewrites the assistant message content in place so the portal
// and runtime consumers of response["choices"][0]["message"]["content"] all see
// the fact-checked text.
func overrideAnswer(response map[string]any, content string) {
if message := firstChoiceMessage(response); message != nil {
message["content"] = content
}
}
func (h *RuntimeHTTPHandler) prepareApplication(ctx context.Context, app Application, input applicationRequest, principal apikey.Principal, retrievalCount *int, traceID string) (map[string]any, map[string]Tool, error) {
config := *app.PublishedConfig
messages := make([]map[string]any, 0, len(input.Messages)+2)
total := 0
lastQuestion := ""
for _, message := range input.Messages {
role, _ := message["role"].(string)
content, contentOK := message["content"].(string)
if (role != "user" && role != "assistant") || !contentOK {
return nil, nil, errors.New("应用对话只接受 user/assistant 文本消息")
}
total += len(content)
if total > 100000 {
return nil, nil, errors.New("对话历史超过 100000 字符")
}
messages = append(messages, map[string]any{"role": role, "content": content})
if role == "user" {
lastQuestion = content
}
}
if lastQuestion == "" {
return nil, nil, errors.New("至少需要一条用户消息")
}
system := []string{}
if config.PromptTemplateID != "" {
prompt, err := h.service.GetPrompt(ctx, config.PromptTemplateID)
// The prompt must be visible to this principal, mirroring the direct
// render/list entry points, or a shared application could leak a
// department-scoped prompt across departments.
if err != nil || prompt.Current == nil || !prompt.Enabled || !visible(prompt.DepartmentIDs, principal, false) {
return nil, nil, errors.New("应用绑定的 Prompt 当前不可用")
}
rendered, err := RenderPrompt(prompt.Current.Content, prompt.Current.Variables, input.Variables)
if err != nil {
return nil, nil, err
}
system = append(system, rendered)
}
evidence := []string{}
for _, kbID := range config.KnowledgeBaseIDs {
kb, err := h.service.GetKnowledgeBase(ctx, kbID)
if err != nil || !kb.Enabled || !visible(kb.DepartmentIDs, principal, false) {
return nil, nil, fmt.Errorf("应用绑定的知识库 %s 当前不可用", kbID)
}
hits, searchErr := h.searchWithTrace(ctx, traceID, kbID, lastQuestion, config.RetrievalTopK)
if searchErr != nil {
continue
}
for _, hit := range hits {
*retrievalCount++
evidence = append(evidence, fmt.Sprintf("[资料%d%s]\n%s", len(evidence)+1, hit.DocumentTitle, hit.Content))
}
}
if len(evidence) > 0 {
system = append(system, "请优先依据以下企业资料回答;资料不足时明确说明不确定,不得编造。引用时使用[资料N]。\n\n"+strings.Join(evidence, "\n\n"))
}
if len(system) > 0 {
messages = append([]map[string]any{{"role": "system", "content": strings.Join(system, "\n\n")}}, messages...)
}
toolsByCode := map[string]Tool{}
schemas := []map[string]any{}
for _, toolID := range config.ToolIDs {
tool, err := h.tools.Get(ctx, toolID)
// Enforce the same department visibility as the direct tool invoke
// entry point (secure=true because tools carry embedded request
// headers); otherwise an app shared across departments could trigger a
// department-only tool and borrow its stored credentials.
if err != nil || !tool.Enabled || !visible(tool.DepartmentIDs, principal, true) {
return nil, nil, fmt.Errorf("应用绑定的工具 %s 当前不可用", toolID)
}
toolsByCode[tool.Code] = tool
var schema any
_ = json.Unmarshal(tool.InputSchema, &schema)
schemas = append(schemas, map[string]any{"type": "function", "function": map[string]any{"name": tool.Code, "description": tool.Description, "parameters": schema}})
}
payload := map[string]any{"model": config.Model, "messages": messages, "stream": false, "temperature": config.Temperature}
if len(schemas) > 0 {
payload["tools"] = schemas
payload["tool_choice"] = "auto"
}
return payload, toolsByCode, nil
}
func (h *RuntimeHTTPHandler) callGateway(original *http.Request, payload map[string]any) (int, http.Header, map[string]any, error) {
raw, _ := json.Marshal(payload)
request := original.Clone(original.Context())
request.Method = http.MethodPost
request.URL.Path = "/v1/chat/completions"
request.URL.RawPath = ""
request.Body = http.NoBody
if len(raw) > 0 {
request.Body = ioNopCloser{bytes.NewReader(raw)}
}
request.ContentLength = int64(len(raw))
request.Header = request.Header.Clone()
request.Header.Set("Content-Type", "application/json")
recorder := newBoundedRecorder()
h.gateway.ServeHTTP(recorder, request)
if recorder.overrun > 0 {
return http.StatusBadGateway, recorder.Header(), nil,
errors.New("模型响应超过 2MB 上限,已截断")
}
result := recorder.Result()
defer result.Body.Close()
var decoded map[string]any
if json.NewDecoder(result.Body).Decode(&decoded) != nil {
return result.StatusCode, result.Header, nil, errors.New("模型返回无法解析")
}
if result.StatusCode < 200 || result.StatusCode >= 300 {
message := "应用模型调用失败"
if value, ok := decoded["error"].(map[string]any); ok {
if text, ok := value["message"].(string); ok {
message = text
}
}
return result.StatusCode, result.Header, decoded, errors.New(message)
}
return result.StatusCode, result.Header, decoded, nil
}
type ioNopCloser struct{ *bytes.Reader }
func (ioNopCloser) Close() error { return nil }
// maxGatewayResponseBytes caps how much of a model response a buffered app
// call may hold in memory. Non-streaming conversations go through a recorder
// that buffers the full upstream response; without a cap a long completion
// could exhaust process memory under concurrent app conversations.
const maxGatewayResponseBytes = 2 << 20 // 2 MiB
// boundedRecorder is a minimal http.ResponseWriter that buffers the response
// up to maxGatewayResponseBytes. Anything beyond the cap is discarded (but
// counted) so a runaway upstream completion can never exhaust memory; callGateway
// turns an overrun into an explicit error instead of decoding truncated JSON.
type boundedRecorder struct {
code int
header http.Header
body bytes.Buffer
overrun int64
}
func newBoundedRecorder() *boundedRecorder {
// code 初始为 0:WriteHeader 只在首次调用时生效,若网关从未调用
// WriteHeader,则 Write 时默认回退 200。
return &boundedRecorder{header: make(http.Header)}
}
func (r *boundedRecorder) Header() http.Header { return r.header }
func (r *boundedRecorder) WriteHeader(code int) {
if r.code != 0 {
return
}
r.code = code
}
func (r *boundedRecorder) Write(data []byte) (int, error) {
if r.code == 0 {
r.code = http.StatusOK
}
remaining := int64(maxGatewayResponseBytes) - int64(r.body.Len())
if remaining > 0 {
written := data
if int64(len(written)) > remaining {
written = written[:remaining]
}
_, _ = r.body.Write(written)
}
r.overrun += int64(len(data)) - min(remaining, int64(len(data)))
return len(data), nil
}
func (r *boundedRecorder) Flush() {}
func (r *boundedRecorder) Result() *http.Response {
return &http.Response{
StatusCode: r.code,
Header: r.header,
Body: io.NopCloser(bytes.NewReader(r.body.Bytes())),
}
}
type toolCall struct{ ID, Name, Arguments string }
func extractToolCalls(response map[string]any) []toolCall {
message := firstChoiceMessage(response)
rawCalls, ok := message["tool_calls"].([]any)
if !ok {
return nil
}
calls := []toolCall{}
for _, raw := range rawCalls {
item, ok := raw.(map[string]any)
if !ok {
continue
}
fn, _ := item["function"].(map[string]any)
calls = append(calls, toolCall{ID: toString(item["id"]), Name: toString(fn["name"]), Arguments: toString(fn["arguments"])})
}
return calls
}
func firstChoiceMessage(response map[string]any) map[string]any {
choices, ok := response["choices"].([]any)
if !ok || len(choices) == 0 {
return map[string]any{}
}
choice, _ := choices[0].(map[string]any)
message, _ := choice["message"].(map[string]any)
return message
}
func valueOrZero(value *int) int {
if value == nil {
return 0
}
return *value
}
func decodeRuntime(w http.ResponseWriter, r *http.Request, target any) bool {
decoder := json.NewDecoder(http.MaxBytesReader(w, r.Body, 2<<20))
decoder.UseNumber()
if err := decoder.Decode(target); err != nil {
runtimeError(w, 400, "请求格式无效")
return false
}
return true
}
func runtimeError(w http.ResponseWriter, status int, message string) {
writeRuntime(w, status, map[string]any{"error": map[string]any{"message": message, "type": "application_error"}})
}
func writeRuntime(w http.ResponseWriter, status int, value any) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(value)
}
func copyHeaders(target, source http.Header) {
for key, values := range source {
lower := strings.ToLower(key)
if lower == "connection" || lower == "content-length" || lower == "transfer-encoding" {
continue
}
target[key] = append([]string(nil), values...)
}
}