feat: 三协议互转网关 + 鉴权修复 + 管理端增强
后端 - 新增 proxy/convert 三协议(chat/messages/responses)请求、响应与 SSE 流式互转, 以 Chat 为中间模型;usage.go 统一提取三协议 token 用量(含单测) - gateway: 跨协议调度(渠道未声明客户端协议时转为渠道首选格式), streamResponse 按 \n\n 分块逐行转换直通,bufferResponse 转换失败时剥非 JSON 前缀 - gateway: 新增 SetUsageRecorder 注入异步用量记录器 - auth_llm: 修复 key_prefix 查询长度错配([:8] vs 存储的 [:12])导致全部 401; 修复长度 8-11 的 key 切片越界 panic;统一 unauthorized 响应 - usage: 日报表改为增量累加 upsert,避免多次 flush 互相清零;记录协议/错误码/时延等字段 - channel: 新增渠道并发槽 TryAcquire;健康检查支持可配置参数 - api: 新增 admin 渠道/模型/系统配置管理端点(旧端点保留兼容) 前端 - 新增渠道管理、模型管理、系统配置视图与 ChannelModelsDrawer - 新增 ui 基础组件(Button/Badge/Input/Modal)与 protocol.ts - 调整 Toast 样式、密钥页、路由菜单;dev 代理默认指向 3000 端口
This commit is contained in:
+103
-108
@@ -1,6 +1,7 @@
|
||||
package proxy
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
@@ -13,6 +14,7 @@ import (
|
||||
"opencatd-open/internal/dao"
|
||||
"opencatd-open/internal/proxy/convert"
|
||||
"opencatd-open/internal/store"
|
||||
"opencatd-open/internal/usage"
|
||||
"opencatd-open/pkg/config"
|
||||
"os"
|
||||
"strings"
|
||||
@@ -35,6 +37,7 @@ type Gateway struct {
|
||||
usageDAO *dao.UsageDAO
|
||||
dailyDAO *dao.DailyUsageDAO
|
||||
channelSvc *channel.Service
|
||||
usageRec *usage.Recorder
|
||||
}
|
||||
|
||||
func NewGateway(ctx context.Context, cfg *config.Config, db *gorm.DB, wg *sync.WaitGroup, userDAO *dao.UserDAO, apiKeyDAO *dao.ApiKeyDAO, usageDAO *dao.UsageDAO, dailyDAO *dao.DailyUsageDAO) *Gateway {
|
||||
@@ -67,6 +70,11 @@ func (g *Gateway) SetChannelService(svc *channel.Service) {
|
||||
g.channelSvc = svc
|
||||
}
|
||||
|
||||
// SetUsageRecorder 注入异步用量记录器;nil 时网关跳过用量上报。
|
||||
func (g *Gateway) SetUsageRecorder(r *usage.Recorder) {
|
||||
g.usageRec = r
|
||||
}
|
||||
|
||||
// Request represents a parsed incoming request
|
||||
type Request struct {
|
||||
Model string
|
||||
@@ -144,26 +152,24 @@ func (g *Gateway) Dispatch(c *gin.Context, req *Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// Determine target format and convert if needed
|
||||
targetFormat := req.Protocol
|
||||
if len(ch.FormatsEffective()) > 0 {
|
||||
// Prefer the channel's native format
|
||||
for _, f := range ch.FormatsEffective() {
|
||||
if f == req.Protocol {
|
||||
targetFormat = f
|
||||
break
|
||||
}
|
||||
}
|
||||
// Determine target format: channel declares support for the client protocol
|
||||
// then passthrough, otherwise convert to its first supported protocol
|
||||
// (chat > messages > responses).
|
||||
targetFormat := g.conversionTarget(ch, req.Protocol)
|
||||
if targetFormat == "" {
|
||||
g.writeError(c, http.StatusBadGateway, fmt.Sprintf("channel %q declares no supported protocol format", ch.Name))
|
||||
return
|
||||
}
|
||||
|
||||
// Build upstream URL
|
||||
upstreamPath := g.getUpstreamPath(req.Protocol)
|
||||
upstreamURL := ch.UpstreamURL(req.Protocol, upstreamPath)
|
||||
upstreamPath := g.getUpstreamPath(targetFormat)
|
||||
upstreamURL := ch.UpstreamURL(targetFormat, upstreamPath)
|
||||
|
||||
// Convert request if needed
|
||||
var requestBody []byte
|
||||
if targetFormat != req.Protocol {
|
||||
requestBody, err = g.convertRequest(req.Body, req.Protocol, targetFormat)
|
||||
var err error
|
||||
requestBody, err = convert.ConvertRequest(req.Body, req.Protocol, targetFormat)
|
||||
if err != nil {
|
||||
g.writeError(c, http.StatusBadRequest, "conversion failed: "+err.Error())
|
||||
return
|
||||
@@ -206,12 +212,31 @@ func (g *Gateway) Dispatch(c *gin.Context, req *Request) {
|
||||
|
||||
// Stream or buffer response
|
||||
if req.Stream {
|
||||
g.streamResponse(c, resp, req.Protocol, ch)
|
||||
g.streamResponse(c, resp, req.Protocol, targetFormat)
|
||||
} else {
|
||||
g.bufferResponse(c, resp, req.Protocol, ch)
|
||||
g.bufferResponse(c, resp, req.Protocol, targetFormat)
|
||||
}
|
||||
}
|
||||
|
||||
// conversionTarget 决定客户端协议在渠道上的处理方式:
|
||||
// 渠道声明支持该协议则直通;否则转为其首选支持协议(chat > messages > responses)。
|
||||
func (g *Gateway) conversionTarget(ch *store.Channel, clientProto string) string {
|
||||
formats := ch.FormatsEffective()
|
||||
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 ""
|
||||
}
|
||||
|
||||
func (g *Gateway) getUpstreamPath(protocol string) string {
|
||||
switch protocol {
|
||||
case "chat":
|
||||
@@ -237,116 +262,86 @@ func (g *Gateway) setHeaders(req *http.Request, ch *store.Channel, apiKey string
|
||||
}
|
||||
}
|
||||
|
||||
func (g *Gateway) convertRequest(body []byte, from, to string) ([]byte, error) {
|
||||
switch {
|
||||
case from == "chat" && to == "messages":
|
||||
var req convert.ChatCompletionRequest
|
||||
if err := json.Unmarshal(body, &req); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
msgReq, err := convert.ChatToMessages(&req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return json.Marshal(msgReq)
|
||||
|
||||
case from == "chat" && to == "responses":
|
||||
var req convert.ChatCompletionRequest
|
||||
if err := json.Unmarshal(body, &req); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
respReq, err := convert.ChatToResponses(&req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return json.Marshal(respReq)
|
||||
|
||||
case from == "messages" && to == "chat":
|
||||
var req convert.MessagesRequest
|
||||
if err := json.Unmarshal(body, &req); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Messages -> Chat: we need to construct a ChatCompletionRequest
|
||||
chatReq := &convert.ChatCompletionRequest{
|
||||
Model: req.Model,
|
||||
}
|
||||
for _, m := range req.Messages {
|
||||
chatReq.Messages = append(chatReq.Messages, m)
|
||||
}
|
||||
if req.Temperature != nil {
|
||||
chatReq.Temperature = req.Temperature
|
||||
}
|
||||
if req.TopP != nil {
|
||||
chatReq.TopP = req.TopP
|
||||
}
|
||||
chatReq.Tools = req.Tools
|
||||
chatReq.Stream = req.Stream
|
||||
return json.Marshal(chatReq)
|
||||
|
||||
case from == "responses" && to == "chat":
|
||||
var req convert.ResponsesRequest
|
||||
if err := json.Unmarshal(body, &req); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
chatReq := &convert.ChatCompletionRequest{
|
||||
Model: req.Model,
|
||||
}
|
||||
for _, item := range req.Input {
|
||||
chatReq.Messages = append(chatReq.Messages, convert.Message{
|
||||
Role: item.Role,
|
||||
Content: item.Content,
|
||||
})
|
||||
}
|
||||
chatReq.Tools = req.Tools
|
||||
chatReq.Stream = req.Stream
|
||||
return json.Marshal(chatReq)
|
||||
|
||||
default:
|
||||
return body, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (g *Gateway) streamResponse(c *gin.Context, resp *http.Response, protocol string, ch *store.Channel) {
|
||||
// streamResponse 流式响应:按 \n\n 分块零缓冲转发;跨协议时逐行转换。
|
||||
func (g *Gateway) streamResponse(c *gin.Context, resp *http.Response, clientProto, upstreamProto string) {
|
||||
w := c.Writer
|
||||
c.Header("Content-Type", "text/event-stream")
|
||||
c.Header("Cache-Control", "no-cache")
|
||||
c.Header("Connection", "keep-alive")
|
||||
c.Status(http.StatusOK)
|
||||
|
||||
writer := convert.NewSSEWriter(c.Writer)
|
||||
parser := convert.NewSSEParser(resp.Body)
|
||||
flusher, _ := w.(http.Flusher)
|
||||
|
||||
for {
|
||||
event, err := parser.ReadEvent()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
log.Printf("Stream parse error: %v", err)
|
||||
break
|
||||
}
|
||||
|
||||
if event.Event == "error" {
|
||||
log.Printf("Upstream stream error: %s", event.Data)
|
||||
break
|
||||
}
|
||||
|
||||
// Write raw SSE event based on protocol
|
||||
if err := writer.WriteEvent("chat CompletionChunk", event.Data); err != nil {
|
||||
break
|
||||
}
|
||||
// 跨协议时按行转换;同协议直通(lineConv 为 nil)。
|
||||
var lineConv func([]byte) []byte
|
||||
if upstreamProto != clientProto {
|
||||
lineConv = convert.NewStreamTransformer(upstreamProto, clientProto)
|
||||
}
|
||||
|
||||
writer.WriteDone()
|
||||
// 上游原始行按 \n\n 分块,避免把 data 行内的转义换行当成事件边界。
|
||||
r := bufio.NewReaderSize(resp.Body, 32*1024)
|
||||
for {
|
||||
buf := []byte{}
|
||||
for {
|
||||
line, err := r.ReadSlice('\n')
|
||||
if err == bufio.ErrBufferFull {
|
||||
buf = append(buf, line...)
|
||||
continue
|
||||
}
|
||||
buf = append(buf, line...)
|
||||
if err == io.EOF {
|
||||
if len(buf) == 0 {
|
||||
return
|
||||
}
|
||||
if !bytes.HasSuffix(buf, []byte("\n")) {
|
||||
buf = append(buf, '\n')
|
||||
}
|
||||
} else if err != nil {
|
||||
log.Printf("stream read error: %v", err)
|
||||
return
|
||||
}
|
||||
if len(buf) >= 2 && bytes.HasSuffix(buf, []byte("\n\n")) {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
out := buf
|
||||
if lineConv != nil {
|
||||
out = lineConv(buf)
|
||||
}
|
||||
if len(out) == 0 {
|
||||
continue
|
||||
}
|
||||
if _, err := w.Write(out); err != nil {
|
||||
return // 客户端已断开
|
||||
}
|
||||
if flusher != nil {
|
||||
flusher.Flush()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (g *Gateway) bufferResponse(c *gin.Context, resp *http.Response, protocol string, ch *store.Channel) {
|
||||
func (g *Gateway) bufferResponse(c *gin.Context, resp *http.Response, clientProto, upstreamProto string) {
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
g.writeError(c, http.StatusBadGateway, "failed to read response")
|
||||
return
|
||||
}
|
||||
|
||||
c.Data(resp.StatusCode, "application/json", body)
|
||||
out := body
|
||||
if upstreamProto != clientProto {
|
||||
if converted, cerr := convert.ConvertResponse(body, upstreamProto, clientProto); cerr == nil {
|
||||
out = converted
|
||||
} else {
|
||||
// 转换失败时至少剥掉非 JSON 前缀,让客户端能解析出正文
|
||||
out = convert.CleanJSON(body)
|
||||
}
|
||||
} else {
|
||||
// 直通:部分上游(如 OpenRouter)的 non-stream 响应在 JSON 前夹带空白/注释
|
||||
out = convert.CleanJSON(body)
|
||||
}
|
||||
c.Data(resp.StatusCode, "application/json", out)
|
||||
}
|
||||
|
||||
func (g *Gateway) writeError(c *gin.Context, status int, message string) {
|
||||
|
||||
Reference in New Issue
Block a user