462 lines
14 KiB
Go
462 lines
14 KiB
Go
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 端点。
|
|
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 端点(非流式)。
|
|
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 端点。
|
|
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 格式,从启用渠道的模型缓存聚合)。
|
|
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 端点(无状态子集)。
|
|
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()
|
|
}
|