package api import ( "crypto/rand" "encoding/hex" "encoding/json" "errors" "fmt" "io" "net/http" "strings" "time" "github.com/gin-gonic/gin" "oci-portal/internal/aiwire" "oci-portal/internal/model" "oci-portal/internal/oci" "oci-portal/internal/service" ) // aiGatewayHandler 处理 /ai/v1/* 对外网关端点(OpenAI / Anthropic 协议)。 type aiGatewayHandler struct { gw *service.AiGatewayService } const aiKeyCtx = "aiKey" // auth 是网关鉴权中间件:Bearer 或 x-api-key 双头识别,失败按端点协议返回错误体。 func (h *aiGatewayHandler) auth(c *gin.Context) { key, err := h.gw.VerifyKey(c.Request.Context(), extractAiKey(c)) if err != nil { aiError(c, http.StatusUnauthorized, "authentication_error", "无效或已禁用的 API 密钥") c.Abort() return } c.Set(aiKeyCtx, key) c.Next() } func extractAiKey(c *gin.Context) string { if header := c.GetHeader("Authorization"); strings.HasPrefix(header, "Bearer ") { return strings.TrimSpace(header[7:]) } return strings.TrimSpace(c.GetHeader("x-api-key")) } // aiError 按端点协议输出错误体:/ai/v1/messages 用 Anthropic 格式,其余 OpenAI 格式。 func aiError(c *gin.Context, status int, code, msg string) { if strings.HasSuffix(c.FullPath(), "/messages") { c.JSON(status, aiwire.AnthErrorBody{Type: "error", Error: aiwire.AnthErrorDetail{Type: code, Message: msg}}) return } c.JSON(status, aiwire.ErrorBody{Error: aiwire.ErrorDetail{Message: msg, Type: code}}) } // upstreamError 把编排层错误映射为网关响应(未知模型 404 / 无渠道 503 / 上游状态透传)。 func upstreamError(c *gin.Context, err error) { switch { case errors.Is(err, service.ErrAiUnknownModel): aiError(c, http.StatusNotFound, "model_not_found", err.Error()) case errors.Is(err, service.ErrAiNoChannel): aiError(c, http.StatusServiceUnavailable, "overloaded_error", err.Error()) case errors.Is(err, service.ErrAiUnsupportedBlock): aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) default: if status, ok := oci.ServiceStatus(err); ok { aiError(c, status, "upstream_error", oci.CompactError(err)) return } aiError(c, http.StatusBadGateway, "upstream_error", oci.CompactError(err)) } } func aiRandID(prefix string) string { buf := make([]byte, 12) _, _ = rand.Read(buf) return prefix + hex.EncodeToString(buf) } // keyGroup 取当前请求密钥的分组(空 = 不限分组)。 func keyGroup(c *gin.Context) string { if key, ok := c.Get(aiKeyCtx); ok { return key.(*model.AiKey).Group } return "" } // logEntry 组装调用日志骨架;usage / status 由调用方补齐。 func (h *aiGatewayHandler) logEntry(c *gin.Context, endpoint, modelName string, stream bool, meta service.ChatMeta, start time.Time) model.AiCallLog { entry := model.AiCallLog{ Endpoint: endpoint, Model: modelName, Stream: stream, ChannelID: meta.ChannelID, ChannelName: meta.ChannelName, Retries: meta.Retries, LatencyMs: time.Since(start).Milliseconds(), ClientIP: requestIP(c), } if key, ok := c.Get(aiKeyCtx); ok { k := key.(*model.AiKey) entry.KeyID, entry.KeyName = k.ID, k.Name } return entry } // validateIR 做协议无关的入参检查(模型名、消息、不支持的内容块)。 func validateIR(ir aiwire.ChatRequest) error { if strings.TrimSpace(ir.Model) == "" { return fmt.Errorf("model 不能为空") } if len(ir.Messages) == 0 { return fmt.Errorf("messages 不能为空") } for _, m := range ir.Messages { if m.Content.HasUnsupported() { return service.ErrAiUnsupportedBlock } } return nil } // maybeLogContent 在调用密钥显式开启内容日志且未过期时写正文(红线例外); // respBody 为 nil 时只记请求(流式响应与向量结果不记录);callLogID 关联同次调用日志。 func (h *aiGatewayHandler) maybeLogContent(c *gin.Context, callLogID uint, endpoint, modelName string, stream bool, reqBody, respBody any) { keyVal, ok := c.Get(aiKeyCtx) if !ok { return } key := keyVal.(*model.AiKey) if key.ContentLogUntil == nil || time.Now().After(*key.ContentLogUntil) { return } entry := model.AiContentLog{CallLogID: callLogID, KeyID: key.ID, KeyName: key.Name, Endpoint: endpoint, Model: modelName, Stream: stream} if b, err := json.Marshal(reqBody); err == nil { entry.RequestBody = string(b) } if respBody != nil { if b, err := json.Marshal(respBody); err == nil { entry.ResponseBody = string(b) } } h.gw.LogContent(entry) } // logFailure 记失败调用日志,并在密钥内容日志开启时留请求正文用于排障(错误信息已在 ErrMsg,不记响应)。 func (h *aiGatewayHandler) logFailure(c *gin.Context, entry model.AiCallLog, reqBody any) { entry.Status = c.Writer.Status() callID := h.gw.LogCall(entry) h.maybeLogContent(c, callID, entry.Endpoint, entry.Model, entry.Stream, reqBody, nil) } // chatCompletions 是 OpenAI /ai/v1/chat/completions 端点。 // // @Summary OpenAI 兼容对话补全 // @Tags AI 网关 // @Param body body object true "OpenAI chat/completions 请求体(支持 stream)" // @Success 200 {object} map[string]any "OpenAI 兼容响应(流式为 SSE)" // @Router /ai/v1/chat/completions [post] func (h *aiGatewayHandler) chatCompletions(c *gin.Context) { var ir aiwire.ChatRequest if err := c.ShouldBindJSON(&ir); err != nil { aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) return } if err := validateIR(ir); err != nil { aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) return } if ir.Stream { h.streamOpenAI(c, ir) return } start := time.Now() resp, meta, err := h.gw.Chat(c.Request.Context(), ir, keyGroup(c)) entry := h.logEntry(c, "openai", ir.Model, false, meta, start) if err != nil { upstreamError(c, err) entry.ErrMsg = err.Error() h.logFailure(c, entry, ir) return } resp.ID = aiRandID("chatcmpl-") if resp.Created == 0 { resp.Created = time.Now().Unix() } entry.Status = http.StatusOK fillUsage(&entry, resp.Usage) callID := h.gw.LogCall(entry) h.maybeLogContent(c, callID, "openai", ir.Model, false, ir, resp) c.JSON(http.StatusOK, resp) } func fillUsage(entry *model.AiCallLog, u *aiwire.Usage) { if u == nil { return } entry.PromptTokens, entry.CompletionTokens, entry.TotalTokens = u.PromptTokens, u.CompletionTokens, u.TotalTokens entry.CachedTokens = u.CachedTokens() } // embeddings 是 OpenAI /ai/v1/embeddings 端点(非流式)。 // // @Summary OpenAI 兼容向量嵌入 // @Tags AI 网关 // @Param body body object true "OpenAI embeddings 请求体" // @Success 200 {object} map[string]any "OpenAI 兼容响应" // @Router /ai/v1/embeddings [post] func (h *aiGatewayHandler) embeddings(c *gin.Context) { var req aiwire.EmbeddingsRequest if err := c.ShouldBindJSON(&req); err != nil { aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) return } if strings.TrimSpace(req.Model) == "" || len(req.Input) == 0 { aiError(c, http.StatusBadRequest, "invalid_request_error", "model 与 input 不能为空") return } if req.EncodingFormat != "" && req.EncodingFormat != "float" { aiError(c, http.StatusBadRequest, "invalid_request_error", "仅支持 encoding_format=float") return } start := time.Now() resp, meta, err := h.gw.Embeddings(c.Request.Context(), req, keyGroup(c)) entry := h.logEntry(c, "embeddings", req.Model, false, meta, start) if err != nil { upstreamError(c, err) entry.ErrMsg = err.Error() h.logFailure(c, entry, req) return } entry.Status = http.StatusOK if resp.Usage != nil { entry.PromptTokens, entry.TotalTokens = resp.Usage.PromptTokens, resp.Usage.TotalTokens } callID := h.gw.LogCall(entry) h.maybeLogContent(c, callID, "embeddings", req.Model, false, req, nil) c.JSON(http.StatusOK, resp) } // streamOpenAI 以 OpenAI SSE 直通流式响应,末尾发 [DONE]。 func (h *aiGatewayHandler) streamOpenAI(c *gin.Context, ir aiwire.ChatRequest) { start := time.Now() stream, meta, err := h.gw.OpenStream(c.Request.Context(), ir, keyGroup(c)) entry := h.logEntry(c, "openai", ir.Model, true, meta, start) if err != nil { upstreamError(c, err) entry.ErrMsg = err.Error() h.logFailure(c, entry, ir) return } defer stream.Close() sseHeaders(c) id, created := aiRandID("chatcmpl-"), time.Now().Unix() var usage *aiwire.Usage for { chunk, err := stream.Next() if err != nil { if !errors.Is(err, io.EOF) { entry.ErrMsg = err.Error() } break } chunk.ID, chunk.Created = id, created if chunk.Usage != nil { usage = chunk.Usage } writeSSEData(c, chunk) } c.Writer.WriteString("data: [DONE]\n\n") c.Writer.Flush() entry.Status = http.StatusOK entry.LatencyMs = time.Since(start).Milliseconds() fillUsage(&entry, usage) callID := h.gw.LogCall(entry) h.maybeLogContent(c, callID, "openai", ir.Model, true, ir, nil) } func sseHeaders(c *gin.Context) { c.Header("Content-Type", "text/event-stream") c.Header("Cache-Control", "no-cache") c.Header("X-Accel-Buffering", "no") c.Writer.Flush() } func writeSSEData(c *gin.Context, v any) { b, err := json.Marshal(v) if err != nil { return } c.Writer.WriteString("data: ") c.Writer.Write(b) c.Writer.WriteString("\n\n") c.Writer.Flush() } // messages 是 Anthropic /ai/v1/messages 端点。 // // @Summary Anthropic Messages 兼容端点 // @Tags AI 网关 // @Param body body object true "Anthropic messages 请求体(支持 stream)" // @Success 200 {object} map[string]any "Anthropic 兼容响应(流式为 SSE)" // @Router /ai/v1/messages [post] func (h *aiGatewayHandler) messages(c *gin.Context) { var req aiwire.MessagesRequest if err := c.ShouldBindJSON(&req); err != nil { aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) return } if req.MaxTokens <= 0 { aiError(c, http.StatusBadRequest, "invalid_request_error", "max_tokens 必填且需大于 0") return } ir, err := service.AnthropicToIR(req) if err != nil { aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) return } if err := validateIR(ir); err != nil { aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) return } if req.Stream { h.streamAnthropic(c, ir) return } start := time.Now() resp, meta, err := h.gw.Chat(c.Request.Context(), ir, keyGroup(c)) entry := h.logEntry(c, "anthropic", ir.Model, false, meta, start) if err != nil { upstreamError(c, err) entry.ErrMsg = err.Error() h.logFailure(c, entry, req) return } entry.Status = http.StatusOK fillUsage(&entry, resp.Usage) callID := h.gw.LogCall(entry) out := service.IRRespToAnthropic(resp, aiRandID("msg_")) h.maybeLogContent(c, callID, "anthropic", ir.Model, false, req, out) c.JSON(http.StatusOK, out) } // streamAnthropic 经状态机把 IR chunk 流聚合为 Anthropic SSE 事件流。 func (h *aiGatewayHandler) streamAnthropic(c *gin.Context, ir aiwire.ChatRequest) { start := time.Now() stream, meta, err := h.gw.OpenStream(c.Request.Context(), ir, keyGroup(c)) entry := h.logEntry(c, "anthropic", ir.Model, true, meta, start) if err != nil { upstreamError(c, err) entry.ErrMsg = err.Error() h.logFailure(c, entry, ir) return } defer stream.Close() sseHeaders(c) st := service.NewAnthStream(aiRandID("msg_"), ir.Model) for { chunk, err := stream.Next() if err != nil { if !errors.Is(err, io.EOF) { entry.ErrMsg = err.Error() } break } writeAnthEvents(c, st.Feed(chunk)) } writeAnthEvents(c, st.Finish()) entry.Status = http.StatusOK entry.LatencyMs = time.Since(start).Milliseconds() usage := st.Usage() entry.PromptTokens, entry.CompletionTokens = usage.InputTokens, usage.OutputTokens entry.TotalTokens = usage.InputTokens + usage.OutputTokens entry.CachedTokens = usage.CacheReadInputTokens callID := h.gw.LogCall(entry) h.maybeLogContent(c, callID, "anthropic", ir.Model, true, ir, nil) } func writeAnthEvents(c *gin.Context, events []service.AnthEvent) { for _, ev := range events { b, err := json.Marshal(ev.Data) if err != nil { continue } c.Writer.WriteString("event: " + ev.Event + "\ndata: ") c.Writer.Write(b) c.Writer.WriteString("\n\n") } c.Writer.Flush() } // listModels 是 /ai/v1/models 端点(OpenAI 格式,从启用渠道的模型缓存聚合)。 // // @Summary 可用模型列表 // @Tags AI 网关 // @Success 200 {object} map[string]any "OpenAI 兼容 models 列表" // @Router /ai/v1/models [get] func (h *aiGatewayHandler) listModels(c *gin.Context) { list, err := h.gw.GatewayModels(c.Request.Context(), keyGroup(c)) if err != nil { aiError(c, http.StatusInternalServerError, "api_error", "查询模型列表失败") return } c.JSON(http.StatusOK, list) } // responses 是 OpenAI /ai/v1/responses 端点(无状态子集)。 // // @Summary OpenAI Responses 兼容端点 // @Tags AI 网关 // @Param body body object true "OpenAI responses 请求体(支持 stream)" // @Success 200 {object} map[string]any "OpenAI 兼容响应(流式为 SSE)" // @Router /ai/v1/responses [post] func (h *aiGatewayHandler) responses(c *gin.Context) { var req aiwire.RespRequest if err := c.ShouldBindJSON(&req); err != nil { aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) return } ir, err := service.ResponsesToIR(req) if err != nil { aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) return } if err := validateIR(ir); err != nil { aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) return } if req.Stream { h.streamResponses(c, ir) return } start := time.Now() resp, meta, err := h.gw.Chat(c.Request.Context(), ir, keyGroup(c)) entry := h.logEntry(c, "responses", ir.Model, false, meta, start) if err != nil { upstreamError(c, err) entry.ErrMsg = err.Error() h.logFailure(c, entry, req) return } entry.Status = http.StatusOK fillUsage(&entry, resp.Usage) callID := h.gw.LogCall(entry) out := service.IRRespToResponses(resp, aiRandID("resp_"), time.Now().Unix()) h.maybeLogContent(c, callID, "responses", ir.Model, false, req, out) c.JSON(http.StatusOK, out) } // streamResponses 经状态机把 IR chunk 流聚合为 Responses 语义事件流。 func (h *aiGatewayHandler) streamResponses(c *gin.Context, ir aiwire.ChatRequest) { start := time.Now() stream, meta, err := h.gw.OpenStream(c.Request.Context(), ir, keyGroup(c)) entry := h.logEntry(c, "responses", ir.Model, true, meta, start) if err != nil { upstreamError(c, err) entry.ErrMsg = err.Error() h.logFailure(c, entry, ir) return } defer stream.Close() sseHeaders(c) st := service.NewRespStream(aiRandID("resp_"), ir.Model, time.Now().Unix()) for { chunk, err := stream.Next() if err != nil { if !errors.Is(err, io.EOF) { entry.ErrMsg = err.Error() } break } writeRespEvents(c, st.Feed(chunk)) } writeRespEvents(c, st.Finish()) entry.Status = http.StatusOK entry.LatencyMs = time.Since(start).Milliseconds() fillUsage(&entry, st.Usage()) callID := h.gw.LogCall(entry) h.maybeLogContent(c, callID, "responses", ir.Model, true, ir, nil) } func writeRespEvents(c *gin.Context, events []service.RespEvent) { for _, ev := range events { b, err := json.Marshal(ev.Data) if err != nil { continue } c.Writer.WriteString("event: " + ev.Event + "\ndata: ") c.Writer.Write(b) c.Writer.WriteString("\n\n") } c.Writer.Flush() }