AI网关:Responses直通codex兼容与流式升级回退
CI / test (push) Successful in 31s
Release / release (push) Successful in 39s

This commit is contained in:
2026-07-15 10:39:56 +08:00
parent a8bde89b56
commit e1f8a0539c
8 changed files with 1339 additions and 221 deletions
+106 -12
View File
@@ -8,6 +8,7 @@ import (
"encoding/json"
"errors"
"io"
"log"
"net/http"
"slices"
"strings"
@@ -552,7 +553,7 @@ func (h *aiGatewayHandler) listModels(c *gin.Context) {
//
// @Summary OpenAI Responses 兼容端点
// @Tags AI 网关
// @Param body body aiwire.RespRequest true "OpenAI responses 请求体(支持 stream;服务端工具 web_search/x_search/code_interpreter/mcp 含流式;未列字段原样透传上游)"
// @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) {
@@ -576,16 +577,22 @@ func (h *aiGatewayHandler) responses(c *gin.Context) {
// 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, err := service.RespPassthroughBody(raw)
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 {
h.responsesPassthroughStream(c, body, req)
if len(body) > service.RespStreamUpgradeLimit {
h.responsesStreamUpgrade(c, body, req, compat)
return
}
h.responsesPassthroughStream(c, body, req, compat)
return
}
start := time.Now()
@@ -597,6 +604,7 @@ func (h *aiGatewayHandler) responsesPassthrough(c *gin.Context, raw []byte, req
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)
@@ -604,9 +612,74 @@ func (h *aiGatewayHandler) responsesPassthrough(c *gin.Context, raw []byte, req
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 记账
func (h *aiGatewayHandler) responsesPassthroughStream(c *gin.Context, body []byte, req aiwire.RespRequest) {
// 逐行扫描 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)
@@ -618,7 +691,7 @@ func (h *aiGatewayHandler) responsesPassthroughStream(c *gin.Context, body []byt
}
defer upstream.Close()
sseHeaders(c)
usage, upErr, err := forwardSSE(c, upstream)
usage, upErr, err := forwardSSE(c, upstream, compat)
if err != nil {
entry.ErrMsg = err.Error()
} else if upErr != "" {
@@ -635,25 +708,29 @@ func (h *aiGatewayHandler) responsesPassthroughStream(c *gin.Context, body []byt
}
// forwardSSE 把上游 SSE 逐行转发给客户端,空行(事件边界)即 flush;
// 顺带从 data 行提取 response.completed 的 usage 与错误事件消息
func forwardSSE(c *gin.Context, upstream io.Reader) (*aiwire.Usage, string, error) {
// 顺带从 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 {
c.Writer.Write(line)
trimmed := bytes.TrimSpace(line)
if len(trimmed) == 0 {
c.Writer.Flush()
} else if data, ok := bytes.CutPrefix(trimmed, []byte("data: ")); ok {
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 {
@@ -665,3 +742,20 @@ func forwardSSE(c *gin.Context, upstream io.Reader) (*aiwire.Usage, string, erro
}
}
}
// 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
}