AI网关:Responses直通codex兼容与流式升级回退
This commit is contained in:
+106
-12
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user