Files
oci-portal/internal/api/aigateway_extras.go
T
wangdefa 0a86b5a291
CI / test (push) Successful in 32s
Release / release (push) Successful in 1m4s
AI网关新增TTS/重排/审核端点,xAI工具扩展,swagger修缺
- 新端点 /ai/v1/audio/speech(xai.grok-tts)、/rerank(cohere.rerank-v4)、/moderations(OCI Guardrails)
- Responses 放行 code_interpreter 与远程 mcp 工具,web_search/x_search 解除仅非流式限制
- 模型能力映射扩展:TEXT_RERANK→RERANK、TEXT_TO_AUDIO→TTS
- AI 网关文档独立 docs/ai-gateway.md,字段兼容矩阵只列支持项;README 精简引用
- swagger 修缺:135 处响应注解具体化,RawMessage/联合类型统一渲染 AnyJSON,overrides 迁至 docs/.swaggo
- CHANGELOG 0.4.0,版本段不再记日期;DASH_VERSION v0.4.0
2026-07-13 20:17:06 +08:00

120 lines
4.1 KiB
Go

package api
import (
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"oci-portal/internal/aiwire"
"oci-portal/internal/service"
)
// audioSpeech 是 OpenAI /ai/v1/audio/speech 端点(TTS,直通兼容面)。
//
// @Summary OpenAI Audio Speech 兼容端点(文本转语音)
// @Tags AI 网关
// @Param body body aiwire.SpeechRequest true "OpenAI audio speech 请求体(model/input 必填,voice 见 xAI Grok Voice 列表,language 缺省 auto;未列字段原样透传)"
// @Success 200 {file} binary "音频字节(Content-Type 透传上游,默认 audio/mpeg)"
// @Router /ai/v1/audio/speech [post]
func (h *aiGatewayHandler) audioSpeech(c *gin.Context) {
raw, err := c.GetRawData()
if err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
modelName, body, err := service.SpeechBodyNormalize(raw)
if err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
if !checkKeyModel(c, modelName) {
return
}
start := time.Now()
audio, contentType, meta, err := h.gw.Speech(c.Request.Context(), modelName, body, keyGroup(c))
entry := h.logEntry(c, "speech", modelName, false, meta, start)
if err != nil {
upstreamError(c, err)
entry.ErrMsg = err.Error()
h.logFailure(c, entry, string(raw))
return
}
entry.Status = http.StatusOK
callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "speech", modelName, false, string(raw), nil)
if contentType == "" {
contentType = "audio/mpeg"
}
c.Data(http.StatusOK, contentType, audio)
}
// rerank 是 /ai/v1/rerank 端点(Jina / Cohere 风格文档重排)。
//
// @Summary 文档重排端点(Cohere Rerank)
// @Tags AI 网关
// @Param body body aiwire.RerankRequest true "重排请求体(model/query/documents 必填,可选 top_n/return_documents)"
// @Success 200 {object} aiwire.RerankResponse "重排结果(results[].index 指向入参下标)"
// @Router /ai/v1/rerank [post]
func (h *aiGatewayHandler) rerank(c *gin.Context) {
var req aiwire.RerankRequest
if err := c.ShouldBindJSON(&req); err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
if strings.TrimSpace(req.Model) == "" || strings.TrimSpace(req.Query) == "" || len(req.Documents) == 0 {
aiError(c, http.StatusBadRequest, "invalid_request_error", "model、query 与 documents 不能为空")
return
}
if !checkKeyModel(c, req.Model) {
return
}
start := time.Now()
resp, meta, err := h.gw.Rerank(c.Request.Context(), req, keyGroup(c))
entry := h.logEntry(c, "rerank", 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
callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "rerank", req.Model, false, req, resp)
c.JSON(http.StatusOK, resp)
}
// moderations 是 /ai/v1/moderations 端点(OpenAI 外壳映射 OCI Guardrails)。
//
// @Summary 内容审核端点(OCI Guardrails)
// @Tags AI 网关
// @Param body body aiwire.ModerationsRequest true "审核请求体(input 为字符串或字符串数组,单次至多 8 条;model 接受但忽略)"
// @Success 200 {object} aiwire.ModerationsResponse "审核结果(categories/category_scores 为 overall/blocklist/prompt_injection,pii 为扩展字段)"
// @Router /ai/v1/moderations [post]
func (h *aiGatewayHandler) moderations(c *gin.Context) {
var req aiwire.ModerationsRequest
if err := c.ShouldBindJSON(&req); err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
inputs, err := service.ModerationInputs(req.Input)
if err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
start := time.Now()
resp, meta, err := h.gw.Moderations(c.Request.Context(), aiRandID("modr_"), inputs, keyGroup(c))
entry := h.logEntry(c, "moderations", "oci-guardrails", false, meta, start)
if err != nil {
upstreamError(c, err)
entry.ErrMsg = err.Error()
h.logFailure(c, entry, req)
return
}
entry.Status = http.StatusOK
callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "moderations", "oci-guardrails", false, req, resp)
c.JSON(http.StatusOK, resp)
}