762 lines
27 KiB
Go
762 lines
27 KiB
Go
package api
|
|
|
|
import (
|
|
"bufio"
|
|
"bytes"
|
|
"crypto/rand"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"errors"
|
|
"io"
|
|
"log"
|
|
"net/http"
|
|
"slices"
|
|
"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 ""
|
|
}
|
|
|
|
// keyModels 取当前请求密钥的模型白名单(空 = 不限模型)。
|
|
func keyModels(c *gin.Context) []string {
|
|
if key, ok := c.Get(aiKeyCtx); ok {
|
|
return key.(*model.AiKey).Models
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// checkKeyModel 校验请求模型是否在密钥白名单内;拒绝时按端点协议返回
|
|
// 404 model_not_found(与未知模型同口径,不泄露密钥配置细节)并返回 false。
|
|
func checkKeyModel(c *gin.Context, modelName string) bool {
|
|
allowed := keyModels(c)
|
|
if len(allowed) == 0 || slices.Contains(allowed, modelName) {
|
|
return true
|
|
}
|
|
aiError(c, http.StatusNotFound, "model_not_found", "模型不存在或无权访问: "+modelName)
|
|
return false
|
|
}
|
|
|
|
// 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 做协议无关的入参检查(模型名、消息、不支持的内容块)。
|
|
// 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)
|
|
}
|
|
|
|
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 aiwire.EmbeddingsRequest true "OpenAI embeddings 请求体"
|
|
// @Success 200 {object} aiwire.EmbeddingsResponse "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
|
|
}
|
|
if !checkKeyModel(c, req.Model) {
|
|
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)
|
|
}
|
|
|
|
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()
|
|
}
|
|
|
|
// anthDefaultMaxTokens 是 max_tokens 缺省时的默认输出上限:Anthropic 协议
|
|
// 该字段必填,但部分客户端当选填不传,按默认值放行而非 400 拒绝。
|
|
const anthDefaultMaxTokens = 8192
|
|
|
|
// messages 是 Anthropic /ai/v1/messages 端点。
|
|
//
|
|
// @Summary Anthropic Messages 兼容端点
|
|
// @Tags AI 网关
|
|
// @Param body body aiwire.MessagesRequest true "Anthropic messages 请求体(支持 stream;经 OCI OpenAI 兼容面直通;max_tokens 可缺省,默认 8192)"
|
|
// @Success 200 {object} aiwire.MessagesResponse "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 {
|
|
req.MaxTokens = anthDefaultMaxTokens
|
|
}
|
|
if len(req.Messages) == 0 {
|
|
aiError(c, http.StatusBadRequest, "invalid_request_error", "messages 不能为空")
|
|
return
|
|
}
|
|
if !checkKeyModel(c, req.Model) {
|
|
return
|
|
}
|
|
body, err := service.AnthropicToResponsesBody(req)
|
|
if err != nil {
|
|
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
|
return
|
|
}
|
|
if req.Stream {
|
|
h.streamAnthropic(c, body, req)
|
|
return
|
|
}
|
|
start := time.Now()
|
|
payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), body, req.Model, keyGroup(c))
|
|
entry := h.logEntry(c, "anthropic", req.Model, false, meta, start)
|
|
if err != nil {
|
|
upstreamError(c, err)
|
|
entry.ErrMsg = err.Error()
|
|
h.logFailure(c, entry, req)
|
|
return
|
|
}
|
|
out, err := service.ResponsesToAnthropic(payload, aiRandID("msg_"))
|
|
if err != nil {
|
|
aiError(c, http.StatusBadGateway, "api_error", err.Error())
|
|
entry.ErrMsg = err.Error()
|
|
h.logFailure(c, entry, req)
|
|
return
|
|
}
|
|
entry.Status = http.StatusOK
|
|
fillUsage(&entry, service.RespPassthroughUsage(payload))
|
|
callID := h.gw.LogCall(entry)
|
|
h.maybeLogContent(c, callID, "anthropic", req.Model, false, req, out)
|
|
c.JSON(http.StatusOK, out)
|
|
}
|
|
|
|
// streamAnthropic 流式直通并桥接:上游 Responses SSE 事件经 AnthRespBridge
|
|
// 转为 Anthropic 事件序列逐块写出;usage 从 completed 事件记账。
|
|
func (h *aiGatewayHandler) streamAnthropic(c *gin.Context, body []byte, req aiwire.MessagesRequest) {
|
|
start := time.Now()
|
|
upstream, meta, err := h.gw.RespPassthroughStream(c.Request.Context(), body, req.Model, keyGroup(c))
|
|
entry := h.logEntry(c, "anthropic", req.Model, true, meta, start)
|
|
if err != nil {
|
|
upstreamError(c, err)
|
|
entry.ErrMsg = err.Error()
|
|
h.logFailure(c, entry, req)
|
|
return
|
|
}
|
|
defer upstream.Close()
|
|
sseHeaders(c)
|
|
bridge := service.NewAnthRespBridge(aiRandID("msg_"), req.Model)
|
|
emitted := 0
|
|
if err := forwardSSEData(upstream, func(data []byte) {
|
|
evs := bridge.Feed(data)
|
|
emitted += len(evs)
|
|
writeAnthEvents(c, evs)
|
|
}); err != nil {
|
|
entry.ErrMsg = err.Error()
|
|
}
|
|
// 上游断流且客户端尚未收到任何事件(OCI 兼容面对大请求 + 推理模型的
|
|
// 流式通道会在 reasoning 阶段掐断):降级非流式重做,结果按事件序列推送
|
|
if emitted == 0 && !bridge.SawTerminal() && c.Request.Context().Err() == nil {
|
|
if h.anthFallback(c, &entry, req) {
|
|
entry.Status = http.StatusOK
|
|
entry.LatencyMs = time.Since(start).Milliseconds()
|
|
callID := h.gw.LogCall(entry)
|
|
h.maybeLogContent(c, callID, "anthropic", req.Model, true, req, nil)
|
|
return
|
|
}
|
|
}
|
|
writeAnthEvents(c, bridge.Finish())
|
|
// 上游错误事件 / 提前终止:SSE 头已发出维持 200,错误落日志可查
|
|
if msg := bridge.Err(); msg != "" && entry.ErrMsg == "" {
|
|
entry.ErrMsg = msg
|
|
}
|
|
entry.Status = http.StatusOK
|
|
entry.LatencyMs = time.Since(start).Milliseconds()
|
|
usage := bridge.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", req.Model, true, req, nil)
|
|
}
|
|
|
|
// anthFallback 用非流式重做同一请求并把完整结果按事件序列推送;
|
|
// 成功返回 true 并把渠道 / 用量 / 降级标注写入日志条目。
|
|
func (h *aiGatewayHandler) anthFallback(c *gin.Context, entry *model.AiCallLog, req aiwire.MessagesRequest) bool {
|
|
req.Stream = false
|
|
body, err := service.AnthropicToResponsesBody(req)
|
|
if err != nil {
|
|
return false
|
|
}
|
|
payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), body, req.Model, keyGroup(c))
|
|
if err != nil {
|
|
return false
|
|
}
|
|
out, err := service.ResponsesToAnthropic(payload, aiRandID("msg_"))
|
|
if err != nil {
|
|
return false
|
|
}
|
|
writeAnthEvents(c, service.AnthMessageEvents(out))
|
|
entry.ChannelID, entry.ChannelName = meta.ChannelID, meta.ChannelName
|
|
entry.Retries++
|
|
entry.ErrMsg = "流式上游断流,已降级非流式完成"
|
|
fillUsage(entry, service.RespPassthroughUsage(payload))
|
|
return true
|
|
}
|
|
|
|
// forwardSSEData 逐行读上游 SSE,把 data 行交给 emit 即时转换写出;
|
|
// Anthropic 与 Chat Completions 两条流式桥共用。
|
|
func forwardSSEData(upstream io.Reader, emit func([]byte)) error {
|
|
reader := bufio.NewReader(upstream)
|
|
for {
|
|
line, err := reader.ReadBytes('\n')
|
|
if len(line) > 0 {
|
|
trimmed := bytes.TrimSpace(line)
|
|
if data, ok := bytes.CutPrefix(trimmed, []byte("data: ")); ok {
|
|
emit(data)
|
|
}
|
|
}
|
|
if err != nil {
|
|
if errors.Is(err, io.EOF) {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
}
|
|
}
|
|
|
|
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()
|
|
}
|
|
|
|
// chatCompletions 是 OpenAI /ai/v1/chat/completions 端点(Tier 2 兼容层,
|
|
// 承接只会说 Chat Completions 的存量客户端)。
|
|
//
|
|
// @Summary OpenAI Chat Completions 兼容端点
|
|
// @Tags AI 网关
|
|
// @Param body body aiwire.ChatRequest true "OpenAI chat/completions 请求体(支持 stream;经 Responses 转换直通)"
|
|
// @Success 200 {object} aiwire.ChatResponse "OpenAI 兼容响应(非流式;流式为 SSE,末尾 data: [DONE])"
|
|
// @Router /ai/v1/chat/completions [post]
|
|
func (h *aiGatewayHandler) chatCompletions(c *gin.Context) {
|
|
var req aiwire.ChatRequest
|
|
if err := c.ShouldBindJSON(&req); err != nil {
|
|
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
|
return
|
|
}
|
|
if strings.TrimSpace(req.Model) == "" || len(req.Messages) == 0 {
|
|
aiError(c, http.StatusBadRequest, "invalid_request_error", "model 与 messages 不能为空")
|
|
return
|
|
}
|
|
if !checkKeyModel(c, req.Model) {
|
|
return
|
|
}
|
|
body, err := service.ChatToResponsesBody(req)
|
|
if err != nil {
|
|
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
|
return
|
|
}
|
|
if req.Stream {
|
|
h.streamChat(c, body, req)
|
|
return
|
|
}
|
|
h.chatOnce(c, body, req)
|
|
}
|
|
|
|
// chatOnce 处理非流式:直通上游 → 转回 chat.completion → 记账。
|
|
func (h *aiGatewayHandler) chatOnce(c *gin.Context, body []byte, req aiwire.ChatRequest) {
|
|
start := time.Now()
|
|
payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), body, req.Model, keyGroup(c))
|
|
entry := h.logEntry(c, "openai", req.Model, false, meta, start)
|
|
if err != nil {
|
|
upstreamError(c, err)
|
|
entry.ErrMsg = err.Error()
|
|
h.logFailure(c, entry, req)
|
|
return
|
|
}
|
|
out, err := service.ResponsesToChat(payload, aiRandID("chatcmpl-"), time.Now().Unix())
|
|
if err != nil {
|
|
aiError(c, http.StatusBadGateway, "api_error", err.Error())
|
|
entry.ErrMsg = err.Error()
|
|
h.logFailure(c, entry, req)
|
|
return
|
|
}
|
|
entry.Status = http.StatusOK
|
|
fillUsage(&entry, out.Usage)
|
|
callID := h.gw.LogCall(entry)
|
|
h.maybeLogContent(c, callID, "openai", req.Model, false, req, out)
|
|
c.JSON(http.StatusOK, out)
|
|
}
|
|
|
|
// streamChat 流式直通并桥接:上游 Responses SSE 经 ChatRespBridge 转为
|
|
// chat.completion.chunk 序列逐块写出,末尾发 [DONE];usage 从 completed 事件记账。
|
|
func (h *aiGatewayHandler) streamChat(c *gin.Context, body []byte, req aiwire.ChatRequest) {
|
|
start := time.Now()
|
|
upstream, meta, err := h.gw.RespPassthroughStream(c.Request.Context(), body, req.Model, keyGroup(c))
|
|
entry := h.logEntry(c, "openai", req.Model, true, meta, start)
|
|
if err != nil {
|
|
upstreamError(c, err)
|
|
entry.ErrMsg = err.Error()
|
|
h.logFailure(c, entry, req)
|
|
return
|
|
}
|
|
defer upstream.Close()
|
|
sseHeaders(c)
|
|
includeUsage := req.StreamOptions != nil && req.StreamOptions.IncludeUsage
|
|
bridge := service.NewChatRespBridge(aiRandID("chatcmpl-"), req.Model, time.Now().Unix(), includeUsage)
|
|
emitted := 0
|
|
if err := forwardSSEData(upstream, func(data []byte) {
|
|
chunks := bridge.Feed(data)
|
|
emitted += len(chunks)
|
|
writeChatChunks(c, chunks)
|
|
}); err != nil {
|
|
entry.ErrMsg = err.Error()
|
|
}
|
|
// 上游断流且客户端尚未收到任何块:降级非流式重做(成因同 messages 端点)
|
|
if emitted == 0 && bridge.Usage() == nil && c.Request.Context().Err() == nil {
|
|
if h.chatFallback(c, &entry, req, includeUsage) {
|
|
entry.Status = http.StatusOK
|
|
entry.LatencyMs = time.Since(start).Milliseconds()
|
|
callID := h.gw.LogCall(entry)
|
|
h.maybeLogContent(c, callID, "openai", req.Model, true, req, nil)
|
|
return
|
|
}
|
|
}
|
|
writeChatChunks(c, bridge.Finish())
|
|
c.Writer.WriteString("data: [DONE]\n\n")
|
|
c.Writer.Flush()
|
|
// 未见 completed 即结束:usage 缺失说明上游流异常提前终止,落日志可查
|
|
if bridge.Usage() == nil && entry.ErrMsg == "" {
|
|
entry.ErrMsg = "上游流提前终止,未返回终态事件"
|
|
}
|
|
entry.Status = http.StatusOK
|
|
entry.LatencyMs = time.Since(start).Milliseconds()
|
|
fillUsage(&entry, bridge.Usage())
|
|
callID := h.gw.LogCall(entry)
|
|
h.maybeLogContent(c, callID, "openai", req.Model, true, req, nil)
|
|
}
|
|
|
|
// chatFallback 用非流式重做同一请求并把完整结果按 chunk 序列推送;
|
|
// 成功返回 true 并把渠道 / 用量 / 降级标注写入日志条目。
|
|
func (h *aiGatewayHandler) chatFallback(c *gin.Context, entry *model.AiCallLog, req aiwire.ChatRequest, includeUsage bool) bool {
|
|
req.Stream = false
|
|
body, err := service.ChatToResponsesBody(req)
|
|
if err != nil {
|
|
return false
|
|
}
|
|
payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), body, req.Model, keyGroup(c))
|
|
if err != nil {
|
|
return false
|
|
}
|
|
out, err := service.ResponsesToChat(payload, aiRandID("chatcmpl-"), time.Now().Unix())
|
|
if err != nil {
|
|
return false
|
|
}
|
|
writeChatChunks(c, service.ChatResponseChunks(out, includeUsage))
|
|
c.Writer.WriteString("data: [DONE]\n\n")
|
|
c.Writer.Flush()
|
|
entry.ChannelID, entry.ChannelName = meta.ChannelID, meta.ChannelName
|
|
entry.Retries++
|
|
entry.ErrMsg = "流式上游断流,已降级非流式完成"
|
|
fillUsage(entry, service.RespPassthroughUsage(payload))
|
|
return true
|
|
}
|
|
|
|
// writeChatChunks 逐块写出 chunk 的 SSE data 行并 flush。
|
|
func writeChatChunks(c *gin.Context, chunks []aiwire.ChatChunk) {
|
|
for _, ch := range chunks {
|
|
b, err := json.Marshal(ch)
|
|
if err != nil {
|
|
continue
|
|
}
|
|
c.Writer.WriteString("data: ")
|
|
c.Writer.Write(b)
|
|
c.Writer.WriteString("\n\n")
|
|
}
|
|
if len(chunks) > 0 {
|
|
c.Writer.Flush()
|
|
}
|
|
}
|
|
|
|
// listModels 是 /ai/v1/models 端点(OpenAI 格式,从启用渠道的模型缓存聚合)。
|
|
//
|
|
// @Summary 可用模型列表
|
|
// @Tags AI 网关
|
|
// @Success 200 {object} aiwire.ModelList "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
|
|
}
|
|
if allowed := keyModels(c); len(allowed) > 0 {
|
|
kept := make([]aiwire.Model, 0, len(list.Data))
|
|
for _, m := range list.Data {
|
|
if slices.Contains(allowed, m.ID) {
|
|
kept = append(kept, m)
|
|
}
|
|
}
|
|
list.Data = kept
|
|
}
|
|
c.JSON(http.StatusOK, list)
|
|
}
|
|
|
|
// responses 是 OpenAI /ai/v1/responses 端点(无状态子集)。
|
|
//
|
|
// @Summary OpenAI Responses 兼容端点
|
|
// @Tags AI 网关
|
|
// @Param body body aiwire.RespRequest true "OpenAI responses 请求体(支持 stream;服务端工具 web_search/x_search/code_interpreter/mcp 含流式;codex 兼容:namespace 工具组拍平为限定名 function 并在响应还原,custom 工具转 function 包装并回转 custom_tool_call(apply_patch 丢弃),tool_search 剥离,web_search.external_web_access 上游不支持自动处理,超 76KB 流式请求自动改非流式合成 SSE;未列字段原样透传上游)"
|
|
// @Success 200 {object} aiwire.RespResponse "OpenAI 兼容响应(非流式;流式为 SSE);直通仅建模常用字段,未列字段原样返回"
|
|
// @Router /ai/v1/responses [post]
|
|
func (h *aiGatewayHandler) responses(c *gin.Context) {
|
|
raw, err := c.GetRawData()
|
|
if err != nil {
|
|
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
|
return
|
|
}
|
|
var req aiwire.RespRequest
|
|
if err := json.Unmarshal(raw, &req); err != nil {
|
|
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
|
return
|
|
}
|
|
if !checkKeyModel(c, req.Model) {
|
|
return
|
|
}
|
|
h.responsesPassthrough(c, raw, req)
|
|
}
|
|
|
|
// responsesPassthrough 直通链路(唯一上游):原始 body 改写后直达
|
|
// OCI /actions/v1/responses,响应原样透传。
|
|
func (h *aiGatewayHandler) responsesPassthrough(c *gin.Context, raw []byte, req aiwire.RespRequest) {
|
|
if err := service.RespPassthroughValidate(req); err != nil {
|
|
log.Printf("responses 直通(model=%s): 校验拒绝: %v", req.Model, err)
|
|
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
|
return
|
|
}
|
|
body, compat, err := service.RespPassthroughBody(raw)
|
|
if err != nil {
|
|
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
|
return
|
|
}
|
|
logRespCompat(req.Model, compat)
|
|
if req.Stream {
|
|
if len(body) > service.RespStreamUpgradeLimit {
|
|
h.responsesStreamUpgrade(c, body, req, compat)
|
|
return
|
|
}
|
|
h.responsesPassthroughStream(c, body, req, compat)
|
|
return
|
|
}
|
|
start := time.Now()
|
|
payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), body, req.Model, keyGroup(c))
|
|
entry := h.logEntry(c, "responses", req.Model, false, meta, start)
|
|
if err != nil {
|
|
upstreamError(c, err)
|
|
entry.ErrMsg = err.Error()
|
|
h.logFailure(c, entry, req)
|
|
return
|
|
}
|
|
payload = service.RespRestoreToolCalls(payload, compat)
|
|
entry.Status = http.StatusOK
|
|
fillUsage(&entry, service.RespPassthroughUsage(payload))
|
|
callID := h.gw.LogCall(entry)
|
|
h.maybeLogContent(c, callID, "responses", req.Model, false, req, json.RawMessage(payload))
|
|
c.Data(http.StatusOK, "application/json; charset=utf-8", payload)
|
|
}
|
|
|
|
// logRespCompat 记录直通请求的 codex 兼容改写动作(观测)。
|
|
func logRespCompat(model string, compat service.RespCompat) {
|
|
if len(compat.Flattened) == 0 && len(compat.Dropped) == 0 && len(compat.Converted) == 0 {
|
|
return
|
|
}
|
|
var parts []string
|
|
if len(compat.Flattened) > 0 {
|
|
parts = append(parts, "拍平: "+strings.Join(compat.Flattened, ", "))
|
|
}
|
|
if len(compat.Converted) > 0 {
|
|
parts = append(parts, "转换: "+strings.Join(compat.Converted, ", "))
|
|
}
|
|
if len(compat.Dropped) > 0 {
|
|
parts = append(parts, "剥离: "+strings.Join(compat.Dropped, ", "))
|
|
}
|
|
log.Printf("responses 直通(model=%s): %s", model, strings.Join(parts, "; "))
|
|
}
|
|
|
|
// responsesStreamUpgrade 流式升级回退:上游对超大流式请求会在推理途中掐断
|
|
// (实测 >~82KB,见 RespStreamUpgradeLimit),超限时改调非流式上游拿完整响应,
|
|
// 本地合成最小 SSE 事件序列回给客户端;丢失增量输出,换会话不中断。
|
|
func (h *aiGatewayHandler) responsesStreamUpgrade(c *gin.Context, body []byte, req aiwire.RespRequest, compat service.RespCompat) {
|
|
start := time.Now()
|
|
nsBody, err := service.RespDisableStream(body)
|
|
if err != nil {
|
|
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
|
return
|
|
}
|
|
log.Printf("responses 直通(model=%s): 请求体 %dKB 超流式安全上限,改走非流式合成 SSE", req.Model, len(body)/1024)
|
|
payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), nsBody, req.Model, keyGroup(c))
|
|
entry := h.logEntry(c, "responses", req.Model, true, meta, start)
|
|
if err != nil {
|
|
upstreamError(c, err)
|
|
entry.ErrMsg = err.Error()
|
|
h.logFailure(c, entry, req)
|
|
return
|
|
}
|
|
payload = service.RespRestoreToolCalls(payload, compat)
|
|
writeSynthSSE(c, payload)
|
|
entry.Status = http.StatusOK
|
|
entry.LatencyMs = time.Since(start).Milliseconds()
|
|
fillUsage(&entry, service.RespPassthroughUsage(payload))
|
|
callID := h.gw.LogCall(entry)
|
|
h.maybeLogContent(c, callID, "responses", req.Model, true, req, json.RawMessage(payload))
|
|
}
|
|
|
|
// writeSynthSSE 把完整响应按合成事件序列写出;合成失败时降级为一次性 JSON,
|
|
// 客户端至少拿到完整结果。
|
|
func writeSynthSSE(c *gin.Context, payload []byte) {
|
|
events, err := service.RespSynthSSEEvents(payload)
|
|
if err != nil {
|
|
log.Printf("responses 直通: 合成 SSE 失败,降级 JSON 返回: %v", err)
|
|
c.Data(http.StatusOK, "application/json; charset=utf-8", payload)
|
|
return
|
|
}
|
|
sseHeaders(c)
|
|
for _, ev := range events {
|
|
c.Writer.Write([]byte("data: "))
|
|
c.Writer.Write(ev)
|
|
c.Writer.Write([]byte("\n\n"))
|
|
}
|
|
c.Writer.Flush()
|
|
}
|
|
|
|
// responsesPassthroughStream 流式直通:SSE 事件原样转发(推理增量等直达客户端),
|
|
// 逐行扫描 completed 事件提取 usage 记账;refs 非空时对 function_call 事件做
|
|
// namespace 还原后再转发。
|
|
func (h *aiGatewayHandler) responsesPassthroughStream(c *gin.Context, body []byte, req aiwire.RespRequest, compat service.RespCompat) {
|
|
start := time.Now()
|
|
upstream, meta, err := h.gw.RespPassthroughStream(c.Request.Context(), body, req.Model, keyGroup(c))
|
|
entry := h.logEntry(c, "responses", req.Model, true, meta, start)
|
|
if err != nil {
|
|
upstreamError(c, err)
|
|
entry.ErrMsg = err.Error()
|
|
h.logFailure(c, entry, req)
|
|
return
|
|
}
|
|
defer upstream.Close()
|
|
sseHeaders(c)
|
|
usage, upErr, err := forwardSSE(c, upstream, compat)
|
|
if err != nil {
|
|
entry.ErrMsg = err.Error()
|
|
} else if upErr != "" {
|
|
// 上游错误事件已原样转发给客户端,这里落日志可查
|
|
entry.ErrMsg = upErr
|
|
} else if usage == nil {
|
|
entry.ErrMsg = "上游流提前终止,未返回终态事件"
|
|
}
|
|
entry.Status = http.StatusOK
|
|
entry.LatencyMs = time.Since(start).Milliseconds()
|
|
fillUsage(&entry, usage)
|
|
callID := h.gw.LogCall(entry)
|
|
h.maybeLogContent(c, callID, "responses", req.Model, true, req, nil)
|
|
}
|
|
|
|
// forwardSSE 把上游 SSE 逐行转发给客户端,空行(事件边界)即 flush;
|
|
// 顺带从 data 行提取 response.completed 的 usage 与错误事件消息;
|
|
// refs 非空时 data 行先做 namespace 还原(未改动的行原样转发)。
|
|
func forwardSSE(c *gin.Context, upstream io.Reader, compat service.RespCompat) (*aiwire.Usage, string, error) {
|
|
reader := bufio.NewReader(upstream)
|
|
var usage *aiwire.Usage
|
|
var upErr string
|
|
for {
|
|
line, err := reader.ReadBytes('\n')
|
|
if len(line) > 0 {
|
|
trimmed := bytes.TrimSpace(line)
|
|
if data, ok := bytes.CutPrefix(trimmed, []byte("data: ")); ok {
|
|
c.Writer.Write(restoreSSELine(line, data, compat))
|
|
if u := service.RespStreamCompletedUsage(data); u != nil {
|
|
usage = u
|
|
}
|
|
if m := service.RespStreamErrorMsg(data); m != "" && upErr == "" {
|
|
upErr = m
|
|
}
|
|
} else {
|
|
c.Writer.Write(line)
|
|
if len(trimmed) == 0 {
|
|
c.Writer.Flush()
|
|
}
|
|
}
|
|
}
|
|
if err != nil {
|
|
c.Writer.Flush()
|
|
if errors.Is(err, io.EOF) {
|
|
return usage, upErr, nil
|
|
}
|
|
return usage, upErr, err
|
|
}
|
|
}
|
|
}
|
|
|
|
// restoreSSELine 对一行 data 事件做工具调用项还原(namespace/custom),
|
|
// 未改动时原行透传(字节级直通)。
|
|
func restoreSSELine(line, data []byte, compat service.RespCompat) []byte {
|
|
if !compat.NeedRestore() {
|
|
return line
|
|
}
|
|
restored, changed := service.RespRestoreToolCallsEvent(data, compat)
|
|
if !changed {
|
|
return line
|
|
}
|
|
out := make([]byte, 0, len(restored)+8)
|
|
out = append(out, "data: "...)
|
|
out = append(out, restored...)
|
|
out = append(out, '\n')
|
|
return out
|
|
}
|