Files
oci-portal/internal/api/aigateway_test.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

227 lines
8.7 KiB
Go

package api
import (
"context"
"encoding/json"
"net/http"
"strconv"
"strings"
"testing"
"time"
"gorm.io/gorm"
"oci-portal/internal/model"
)
// seedGatewayModel 直插启用渠道与模型缓存,绕过云同步。
func seedGatewayModel(t *testing.T, db *gorm.DB, name string) {
t.Helper()
ch := &model.AiChannel{Name: "t-" + name, OciConfigID: 1, Region: "r-" + name, Enabled: true, Priority: 1, Weight: 1}
if err := db.Create(ch).Error; err != nil {
t.Fatalf("seed channel: %v", err)
}
cache := &model.AiModelCache{ChannelID: ch.ID, ModelOcid: "ocid1.." + name, Name: name, Vendor: "v", SyncedAt: time.Now()}
if err := db.Create(cache).Error; err != nil {
t.Fatalf("seed model cache: %v", err)
}
}
func TestAiGatewayKeyModelRestrict(t *testing.T) {
r, auth, _, db := newTestRouterDB(t)
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
if err != nil {
t.Fatalf("login: %v", err)
}
// 受限密钥:白名单含脏输入,入库应规范化为 2 项(cohere + ghost)
w := doRequest(t, r, http.MethodPost, "/api/v1/ai-keys", token,
`{"name":"limited","value":"limited-key-1234","models":[" cohere.command-a-03-2025 ","cohere.command-a-03-2025","","ghost-model"]}`)
if w.Code != http.StatusCreated {
t.Fatalf("创建受限密钥 status = %d, body %s", w.Code, w.Body.String())
}
var created struct {
Item model.AiKey `json:"item"`
}
if err := json.Unmarshal(w.Body.Bytes(), &created); err != nil {
t.Fatalf("decode: %v", err)
}
if len(created.Item.Models) != 2 || created.Item.Models[0] != "cohere.command-a-03-2025" {
t.Fatalf("创建后 models 未规范化: %v", created.Item.Models)
}
// 不限密钥对照
w = doRequest(t, r, http.MethodPost, "/api/v1/ai-keys", token, `{"name":"open","value":"open-key-12345"}`)
if w.Code != http.StatusCreated {
t.Fatalf("创建不限密钥 status = %d, body %s", w.Code, w.Body.String())
}
seedGatewayModel(t, db, "cohere.command-a-03-2025")
seedGatewayModel(t, db, "meta.llama-3.3-70b-instruct")
deny := []string{"model_not_found", "无权访问"}
pass := []string{"未知模型"} // 穿过白名单闸门,由编排层以「未知模型」拒绝(池内无 ghost-model)
tests := []struct {
name string
key string
path string
body string
wantCode int
wantSub []string
}{
{"responses 流式同样拦截", "limited-key-1234", "/ai/v1/responses",
`{"model":"meta.llama-3.3-70b-instruct","stream":true,"input":"hi"}`, 404, deny},
{"responses 白名单内穿透", "limited-key-1234", "/ai/v1/responses",
`{"model":"ghost-model","input":"hi"}`, 404, pass},
{"embeddings 白名单外拦截", "limited-key-1234", "/ai/v1/embeddings",
`{"model":"meta.llama-3.3-70b-instruct","input":["hi"]}`, 404, deny},
{"messages 白名单外拦截为 Anthropic 体", "limited-key-1234", "/ai/v1/messages",
`{"model":"meta.llama-3.3-70b-instruct","max_tokens":16,"messages":[{"role":"user","content":"hi"}]}`, 404,
[]string{`"type":"error"`, "model_not_found", "无权访问"}},
{"responses 白名单外拦截", "limited-key-1234", "/ai/v1/responses",
`{"model":"meta.llama-3.3-70b-instruct","input":"hi"}`, 404, deny},
{"不限密钥穿透闸门", "open-key-12345", "/ai/v1/responses",
`{"model":"ghost-model","input":"hi"}`, 404, pass},
{"chat completions 白名单外拦截", "limited-key-1234", "/ai/v1/chat/completions",
`{"model":"meta.llama-3.3-70b-instruct","messages":[{"role":"user","content":"hi"}]}`, 404, deny},
{"chat completions 白名单内穿透", "limited-key-1234", "/ai/v1/chat/completions",
`{"model":"ghost-model","messages":[{"role":"user","content":"hi"}]}`, 404, pass},
{"messages 缺 max_tokens 按默认值放行", "open-key-12345", "/ai/v1/messages",
`{"model":"ghost-model","messages":[{"role":"user","content":"hi"}]}`, 404, pass},
{"chat completions 缺 messages 拒绝", "open-key-12345", "/ai/v1/chat/completions",
`{"model":"ghost-model"}`, 400, []string{"invalid_request_error"}},
{"chat completions 非 function 工具拒绝", "open-key-12345", "/ai/v1/chat/completions",
`{"model":"ghost-model","messages":[{"role":"user","content":"hi"}],"tools":[{"type":"web_search","function":{}}]}`,
400, []string{"仅支持 function"}},
{"chat completions 不支持内容块拒绝", "open-key-12345", "/ai/v1/chat/completions",
`{"model":"ghost-model","messages":[{"role":"user","content":[{"type":"input_audio"}]}]}`,
400, []string{"invalid_request_error"}},
{"audio speech 白名单外拦截", "limited-key-1234", "/ai/v1/audio/speech",
`{"model":"meta.llama-3.3-70b-instruct","input":"你好"}`, 404, deny},
{"rerank 白名单外拦截", "limited-key-1234", "/ai/v1/rerank",
`{"model":"meta.llama-3.3-70b-instruct","query":"q","documents":["d"]}`, 404, deny},
{"rerank 缺 documents 拒绝", "open-key-12345", "/ai/v1/rerank",
`{"model":"ghost-model","query":"q"}`, 400, []string{"invalid_request_error"}},
{"moderations 空 input 拒绝", "open-key-12345", "/ai/v1/moderations",
`{"input":[]}`, 400, []string{"invalid_request_error"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := doRequest(t, r, http.MethodPost, tt.path, tt.key, tt.body)
if w.Code != tt.wantCode {
t.Fatalf("status = %d, want %d, body %s", w.Code, tt.wantCode, w.Body.String())
}
for _, sub := range tt.wantSub {
if !strings.Contains(w.Body.String(), sub) {
t.Errorf("body 缺少 %q: %s", sub, w.Body.String())
}
}
})
}
}
func TestAiGatewayModelsFilteredByKey(t *testing.T) {
r, auth, _, db := newTestRouterDB(t)
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
if err != nil {
t.Fatalf("login: %v", err)
}
seedGatewayModel(t, db, "cohere.command-a-03-2025")
seedGatewayModel(t, db, "meta.llama-3.3-70b-instruct")
w := doRequest(t, r, http.MethodPost, "/api/v1/ai-keys", token,
`{"name":"limited","value":"limited-key-1234","models":["cohere.command-a-03-2025"]}`)
if w.Code != http.StatusCreated {
t.Fatalf("创建密钥 status = %d", w.Code)
}
w = doRequest(t, r, http.MethodPost, "/api/v1/ai-keys", token, `{"name":"open","value":"open-key-12345"}`)
if w.Code != http.StatusCreated {
t.Fatalf("创建密钥 status = %d", w.Code)
}
tests := []struct {
name string
key string
want []string
notWant []string
}{
{"受限密钥仅见交集", "limited-key-1234", []string{"cohere.command-a-03-2025"}, []string{"meta.llama-3.3-70b-instruct"}},
{"不限密钥全量可见", "open-key-12345", []string{"cohere.command-a-03-2025", "meta.llama-3.3-70b-instruct"}, nil},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := doRequest(t, r, http.MethodGet, "/ai/v1/models", tt.key, "")
if w.Code != http.StatusOK {
t.Fatalf("status = %d, body %s", w.Code, w.Body.String())
}
for _, sub := range tt.want {
if !strings.Contains(w.Body.String(), sub) {
t.Errorf("应包含 %q: %s", sub, w.Body.String())
}
}
for _, sub := range tt.notWant {
if strings.Contains(w.Body.String(), sub) {
t.Errorf("不应包含 %q: %s", sub, w.Body.String())
}
}
})
}
}
func TestAiKeyModelsUpdateRoundTrip(t *testing.T) {
r, auth, _, _ := newTestRouterDB(t)
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
if err != nil {
t.Fatalf("login: %v", err)
}
w := doRequest(t, r, http.MethodPost, "/api/v1/ai-keys", token, `{"name":"k","value":"round-key-1234"}`)
if w.Code != http.StatusCreated {
t.Fatalf("创建 status = %d", w.Code)
}
var created struct {
Item model.AiKey `json:"item"`
}
_ = json.Unmarshal(w.Body.Bytes(), &created)
id := created.Item.ID
listModels := func() []string {
w := doRequest(t, r, http.MethodGet, "/api/v1/ai-keys", token, "")
var resp struct {
Items []model.AiKey `json:"items"`
}
_ = json.Unmarshal(w.Body.Bytes(), &resp)
for _, k := range resp.Items {
if k.ID == id {
return k.Models
}
}
t.Fatalf("密钥 %d 不在列表中", id)
return nil
}
tests := []struct {
name string
body string
want []string
}{
{"设置白名单", `{"models":[" a-model ","a-model","b-model"]}`, []string{"a-model", "b-model"}},
{"不传 models 保持不变", `{"name":"k2"}`, []string{"a-model", "b-model"}},
{"空数组清空恢复不限", `{"models":[]}`, nil},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := doRequest(t, r, http.MethodPut, "/api/v1/ai-keys/"+strconv.Itoa(int(id)), token, tt.body)
if w.Code != http.StatusNoContent {
t.Fatalf("update status = %d, body %s", w.Code, w.Body.String())
}
got := listModels()
if len(got) != len(tt.want) {
t.Fatalf("models = %v, want %v", got, tt.want)
}
for i := range tt.want {
if got[i] != tt.want[i] {
t.Fatalf("models = %v, want %v", got, tt.want)
}
}
})
}
}