Files
openteam/server/internal/proxy/gateway.go
T
SakurasanandClaude 4d4d09ba58 渠道: 模型名称映射(网关生效 + 渠道侧管理)
- 候选渠道携带 upstream_model, prepareUpstream 改写请求体 model 字段为上游名
- 新增渠道视角绑定 CRUD: GET/POST/PATCH/DELETE /admin/channels/:id/models
- 前端渠道页新增"模型映射": 列出绑定、内联改上游名、解除、添加
- rewriteModel 三种协议通用(model 均在顶层)

Co-Authored-By: Claude <noreply@anthropic.com>
2026-08-16 00:09:01 +08:00

257 lines
8.0 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Package proxy API 网关核心:代理 /v1/chat/completions、/v1/responses、/v1/messages、/v1/models。
// M1 直通 OpenAI 渠道;M4 起按客户端协议 × 渠道协议自动转换(见 convert)。
package proxy
import (
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/openteam/server/internal/channel"
"github.com/openteam/server/internal/pkg/apikey"
"github.com/openteam/server/internal/pkg/crypto"
"github.com/openteam/server/internal/pkg/ratelimit"
"github.com/openteam/server/internal/proxy/convert"
"github.com/openteam/server/internal/store"
"github.com/openteam/server/internal/usage"
"gorm.io/gorm"
)
const (
CtxUserID = "proxy_user_id"
CtxKeyID = "proxy_key_id"
CtxTrace = "proxy_trace_id"
)
type Gateway struct {
db *gorm.DB
ch *channel.Service
rec *usage.Recorder
enc *crypto.Encryptor
lim *ratelimit.Limiter
userRPS int
hc *http.Client
}
func NewGateway(db *gorm.DB, enc *crypto.Encryptor, rec *usage.Recorder, lim *ratelimit.Limiter, userRPS int) *Gateway {
return &Gateway{
db: db,
ch: channel.NewService(db, enc),
rec: rec,
enc: enc,
lim: lim,
userRPS: userRPS,
hc: &http.Client{Timeout: 120 * time.Second},
}
}
// Auth 代理鉴权中间件:Bearer sk-xxx → 哈希查表 → 校验状态/过期/模型白名单。
func (g *Gateway) Auth(c *gin.Context) {
// 先按路径确定客户端协议,保证 Auth 阶段错误也按协议格式返回
switch c.Request.URL.Path {
case "/v1/messages":
c.Set("protocol", convert.ProtoMessages)
case "/v1/responses":
c.Set("protocol", convert.ProtoResponses)
default:
c.Set("protocol", convert.ProtoChat)
}
auth := c.GetHeader("Authorization")
key := strings.TrimPrefix(auth, "Bearer ")
key = strings.TrimSpace(key)
if !apikey.Valid(key) {
apiError(c, http.StatusUnauthorized, "invalid_api_key", "Invalid API key format. Expected: Bearer sk-...")
c.Abort()
return
}
hash := apikey.Hash(key)
var k store.APIKey
if err := g.db.Where("key_hash = ? AND status = ?", hash, store.KeyStatusActive).First(&k).Error; err != nil {
apiError(c, http.StatusUnauthorized, "invalid_api_key", "Invalid API key")
c.Abort()
return
}
var u store.User
if err := g.db.First(&u, k.UserID).Error; err != nil || u.Status != store.UserStatusActive {
apiError(c, http.StatusForbidden, "user_disabled", "User account is disabled")
c.Abort()
return
}
if k.ExpiresAt != nil && time.Now().After(*k.ExpiresAt) {
apiError(c, http.StatusUnauthorized, "key_expired", "API key has expired")
c.Abort()
return
}
// 限流与配额(内存计数)
if g.lim != nil {
if k.QuotaRequestsPerDay != nil && !g.lim.AllowRequestDaily(k.ID, *k.QuotaRequestsPerDay) {
apiError(c, http.StatusTooManyRequests, "rate_limit_exceeded", "Daily request quota exceeded for this API key")
c.Abort()
return
}
if k.QuotaTokensPerDay != nil && g.lim.TokensUsed(k.ID) >= *k.QuotaTokensPerDay {
apiError(c, http.StatusTooManyRequests, "rate_limit_exceeded", "Daily token quota exceeded for this API key")
c.Abort()
return
}
if !g.lim.AllowUserRate(u.ID, g.userRPS) {
apiError(c, http.StatusTooManyRequests, "rate_limit_exceeded", "Too many requests. Please slow down.")
c.Abort()
return
}
}
c.Set(CtxUserID, u.ID)
c.Set(CtxKeyID, k.ID)
c.Set(CtxTrace, newTraceID())
g.db.Model(&store.APIKey{}).Where("id = ?", k.ID).Update("last_used_at", time.Now())
c.Next()
}
// Handle 路由到对应协议处理器。
func (g *Gateway) Handle(c *gin.Context) {
switch {
case c.Request.URL.Path == "/v1/chat/completions":
g.chatCompletions(c)
case c.Request.URL.Path == "/v1/responses":
g.responses(c)
case c.Request.URL.Path == "/v1/messages":
g.messages(c)
case c.Request.URL.Path == "/v1/models" && c.Request.Method == http.MethodGet:
g.models(c)
default:
apiError(c, http.StatusNotFound, "not_found", "Unknown endpoint: "+c.Request.URL.Path)
}
}
// models GET /v1/models:返回启用的全局模型(OpenAI 风格)。
func (g *Gateway) models(c *gin.Context) {
var ms []store.Model
if err := g.db.Where("enabled = ?", true).Order("sort ASC, id ASC").Find(&ms).Error; err != nil {
apiError(c, http.StatusInternalServerError, "internal_error", "failed to load models")
return
}
data := make([]gin.H, 0, len(ms))
for _, m := range ms {
data = append(data, gin.H{
"id": m.Name,
"object": "model",
"created": m.CreatedAt.Unix(),
"owned_by": "openteam",
})
}
c.JSON(http.StatusOK, gin.H{"object": "list", "data": data})
}
// candidateChannels 返回可用渠道候选(按模型绑定优先,退化全局;携带模型映射)。
func (g *Gateway) candidateChannels(model string) []channel.Candidate {
return g.ch.Candidates(model)
}
// resolveUser 取当前用户(含余额)。
func (g *Gateway) resolveUser(c *gin.Context) (*store.User, bool) {
uid, _ := c.Get(CtxUserID)
var u store.User
if err := g.db.First(&u, uid).Error; err != nil {
apiError(c, http.StatusUnauthorized, "invalid_api_key", "Invalid API key")
return nil, false
}
return &u, true
}
// checkBalance 余额不足返回 402(PLANNING §4.4.3)。
func (g *Gateway) checkBalance(c *gin.Context, u *store.User) bool {
if u.Balance <= 0 {
apiError(c, http.StatusPaymentRequired, "insufficient_balance", "Insufficient balance. Please recharge or contact admin.")
return false
}
return true
}
// ---------------------------------------------------------------------------
// 协议分派
func upstreamPath(proto string) string {
switch proto {
case convert.ProtoMessages:
return "/v1/messages"
case convert.ProtoResponses:
return "/v1/responses"
default:
return "/v1/chat/completions"
}
}
// upstreamPlan 描述一次代理请求的上游访问方式。
type upstreamPlan struct {
path string // 上游路径
body []byte // 已转换的请求体
lineConv func([]byte) []byte // 流式逐行转换(nil=直通)
bodyConv func([]byte) ([]byte, error) // 非流式响应体转换(nil=直通)
}
// conversionTarget 决定客户端协议在渠道上的处理方式:
// 渠道声明支持该协议则直通(返回原协议);否则转为其首选支持协议(chat > messages > responses)。
func conversionTarget(formats []string, clientProto string) string {
for _, f := range formats {
if f == clientProto {
return clientProto
}
}
for _, p := range []string{convert.ProtoChat, convert.ProtoMessages, convert.ProtoResponses} {
for _, f := range formats {
if f == p {
return p
}
}
}
return ""
}
// prepareUpstream 计算上游访问计划:渠道声明支持客户端协议则直通,否则转换;
// 应用模型名称映射(upstream_model)。
func prepareUpstream(ch *store.Channel, clientProto string, body []byte, upstreamModel string) (*upstreamPlan, error) {
target := conversionTarget(ch.FormatsEffective(), clientProto)
if target == "" {
return nil, fmt.Errorf("channel %q declares no supported protocol format", ch.Name)
}
plan := &upstreamPlan{path: upstreamPath(target), body: body}
if target != clientProto {
converted, err := convert.ConvertRequest(body, clientProto, target)
if err != nil {
return nil, err
}
plan.body = converted
plan.lineConv = convert.NewStreamTransformer(target, clientProto)
plan.bodyConv = func(b []byte) ([]byte, error) { return convert.ConvertResponse(b, target, clientProto) }
}
// 模型名称映射:把请求体 model 字段改写为渠道侧的 upstream_model
if upstreamModel != "" {
if out, err := rewriteModel(plan.body, upstreamModel); err == nil {
plan.body = out
}
}
return plan, nil
}
// rewriteModel 改写请求体中的 model 字段(三种协议 model 都在顶层)。
func rewriteModel(body []byte, model string) ([]byte, error) {
var m map[string]any
if err := json.Unmarshal(body, &m); err != nil {
return body, nil
}
if cur, _ := m["model"].(string); cur == model {
return body, nil
}
m["model"] = model
return json.Marshal(m)
}
var errNoChannel = errors.New("no available channel")