@@ -4,6 +4,8 @@ import (
"bufio"
"bytes"
"context"
"crypto/rand"
"encoding/hex"
"encoding/json"
"fmt"
"io"
@@ -36,6 +38,7 @@ type Gateway struct {
apiKeyDAO * dao . ApiKeyDAO
usageDAO * dao . UsageDAO
dailyDAO * dao . DailyUsageDAO
modelDAO * dao . ModelDAO
channelSvc * channel . Service
usageRec * usage . Recorder
}
@@ -62,6 +65,7 @@ func NewGateway(ctx context.Context, cfg *config.Config, db *gorm.DB, wg *sync.W
apiKeyDAO : apiKeyDAO ,
usageDAO : usageDAO ,
dailyDAO : dailyDAO ,
modelDAO : dao . NewModelDAO ( db ) ,
channelSvc : nil ,
}
}
@@ -75,6 +79,15 @@ func (g *Gateway) SetUsageRecorder(r *usage.Recorder) {
g . usageRec = r
}
// generateRequestID 生成请求级唯一 ID,用于用量明细关联与排障。
func generateRequestID ( ) string {
b := make ( [ ] byte , 12 )
if _ , err := rand . Read ( b ) ; err != nil {
return fmt . Sprintf ( "req-%d" , time . Now ( ) . UnixNano ( ) )
}
return "req-" + hex . EncodeToString ( b )
}
// Request represents a parsed incoming request
type Request struct {
Model string
@@ -83,6 +96,8 @@ type Request struct {
Body [ ] byte
APIKey * store . APIKey
UserID uint64
KeyID uint64
RequestID string
}
// ParseRequest parses the incoming request and extracts key fields
@@ -96,9 +111,13 @@ func (g *Gateway) ParseRequest(c *gin.Context, protocol string) (*Request, error
userID , _ := c . Get ( "user_id" )
req := & Request {
Protocol : protocol ,
Body : body ,
UserID : userID . ( uint64 ) ,
Protocol : protocol ,
Body : body ,
UserID : userID . ( uint64 ) ,
RequestID : c . GetHeader ( "X-Request-Id" ) ,
}
if req . RequestID == "" {
req . RequestID = generateRequestID ( )
}
if ak , ok := apiKey . ( * store . APIKey ) ; ok {
@@ -133,89 +152,190 @@ func (g *Gateway) ParseRequest(c *gin.Context, protocol string) (*Request, error
return req , nil
}
// Dispatch routes the request to the appropriate upstream
// Dispatch routes the request to the appropriate upstream.
// 遍历候选渠道(绑定优先,全局回退;按优先级/权重排序),可重试性失败自动故障转移。
func ( g * Gateway ) Dispatch ( c * gin . Context , req * Request ) {
if g . channelSvc == nil {
g . writeError ( c , http . StatusBadGateway , "channel service not available" )
return
}
ch , err := g . channelSvc . SelectChannel ( g . ctx , req . Model )
if err != nil {
g . writeError ( c , http . StatusBadGateway , err . Error ( ) )
cands := g . channelSvc . Candidates ( req . Model )
// 内存健康过滤:连续失败进入 cooldown 的渠道不再尝试(渠道级健康自愈靠冷却过期)。
cands = g . channelSvc . FilterHealthy ( cands )
if len ( cands ) == 0 {
g . writeError ( c , http . StatusServiceUnavailable , "no enabled channels for model: " + req . Model )
g . recordUsage ( req , nil , nil , usage . Event {
IsError : true , ErrorCode : "no_channel" ,
} , convert . TokenUsage { } )
return
}
apiKey , err := g . channelSvc . GetAPIKey ( ch )
if err != nil {
g . writeError ( c , http . StatusBadGateway , "failed to decrypt API key" )
return
}
var lastCh * store . Channel
_ = lastCh // 保留变量名便于断点排查;失败渠道已在循环内各自 RecordFailure
lastErrStatus := http . StatusBadGateway
lastErrBody := "all upstream channels failed"
// 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
}
for i := range cands {
cand := & cands [ i ]
ch := cand . Channel
lastCh = ch
// Build upstream URL
upstreamPath := g . getUpstreamPath ( targetFormat )
upstreamURL := ch . UpstreamURL ( targetFormat , upstreamPath )
// Convert request if needed
var requestBody [ ] byte
if targetFormat != req . Protocol {
var err error
requestBody , err = convert . ConvertRequest ( req . Body , req . Protocol , targetFormat )
apiKey , err := g . channelSvc . GetAPIKey ( ch )
if err != nil {
g . writeError ( c , http . StatusBadRequest , "conversion failed: " + err . Error ( ) )
lastErrStatus , lastErrBody = http . StatusBadGateway , "failed to decrypt API key"
continue
}
// 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 == "" {
continue // 渠道不支持该协议,换下一个
}
// Build upstream URL
upstreamURL := ch . UpstreamURL ( targetFormat , g . getUpstreamPath ( targetFormat ) )
// Convert request if needed
var requestBody [ ] byte
if targetFormat != req . Protocol {
var err error
requestBody , err = convert . ConvertRequest ( req . Body , req . Protocol , targetFormat )
if err != nil {
lastErrStatus , lastErrBody = http . StatusBadRequest , "conversion failed: " + err . Error ( )
continue
}
} else {
requestBody = req . Body
}
// 绑定了 upstream_model 时把请求体里的 model 重写为上游模型名(别名映射)。
if cand . Binding != nil && cand . Binding . UpstreamModel != "" &&
cand . Binding . UpstreamModel != req . Model {
requestBody = rewriteModel ( requestBody , cand . Binding . UpstreamModel )
}
// Create upstream request
httpReq , err := http . NewRequestWithContext ( g . ctx , "POST" , upstreamURL , bytes . NewReader ( requestBody ) )
if err != nil {
lastErrStatus , lastErrBody = http . StatusBadGateway , "failed to create request"
continue
}
g . setHeaders ( httpReq , ch , apiKey , targetFormat )
// Execute request
start := time . Now ( )
resp , err := g . httpClient . Do ( httpReq )
if err != nil {
g . channelSvc . RecordFailure ( ch . ID )
lastErrStatus = http . StatusBadGateway
lastErrBody = fmt . Sprintf ( "upstream error: %v" , err )
g . recordUsage ( req , cand , ch , usage . Event {
IsError : true ,
ErrorCode : "upstream_error" ,
LatencyMS : int ( time . Since ( start ) . Milliseconds ( ) ) ,
} , convert . TokenUsage { } )
continue // 可重试:换下一个渠道
}
// Handle upstream error responses
if resp . StatusCode >= 400 {
body , _ := io . ReadAll ( resp . Body )
resp . Body . Close ( )
log . Printf ( "Upstream error: status=%d body=%s" , resp . StatusCode , string ( body ) )
g . recordUsage ( req , cand , ch , usage . Event {
IsError : true ,
ErrorCode : fmt . Sprintf ( "upstream_%d" , resp . StatusCode ) ,
LatencyMS : int ( time . Since ( start ) . Milliseconds ( ) ) ,
} , convert . TokenUsage { } )
// 429/5xx 可换渠道重试;4xx 直接透传
if resp . StatusCode == http . StatusTooManyRequests || resp . StatusCode >= 500 {
lastErrStatus , lastErrBody = resp . StatusCode , string ( body )
continue
}
c . Data ( resp . StatusCode , "application/json" , body )
return
}
} else {
requestBody = req . Body
g . channelSvc . RecordSuccess ( ch . ID )
// Stream or buffer response; tok 从上游响应(SSE usage 块或非流式 JSON)提取。
var tok convert . TokenUsage
if req . Stream {
tok = g . streamResponse ( c , resp , req . Protocol , targetFormat )
} else {
tok = g . bufferResponse ( c , resp , req . Protocol , targetFormat )
}
resp . Body . Close ( )
// 成功记录:用量 + 定价计费。
g . recordUsage ( req , cand , ch , usage . Event {
LatencyMS : int ( time . Since ( start ) . Milliseconds ( ) ) ,
} , tok )
return
}
// Create upstream request
httpReq , err := http . NewRequestWithContext ( g . ctx , "POST" , upstreamURL , bytes . NewReader ( requestBody ) )
// 全部候选失败(每个候选失败时已各自 RecordFailure,不再重复计数)
g . writeError ( c , lastErrStatus , lastErrBody )
}
// rewriteModel 把 JSON 请求体顶层的 model 字段替换为 upstreamModel。
func rewriteModel ( body [ ] byte , upstreamModel string ) [ ] byte {
var m map [ string ] json . RawMessage
if json . Unmarshal ( body , & m ) != nil {
return body
}
if _ , ok := m [ "model" ] ; ! ok {
return body
}
m [ "model" ] , _ = json . Marshal ( upstreamModel )
out , err := json . Marshal ( m )
if err != nil {
g . writeError ( c , http . StatusBadGateway , "failed to create request" )
return body
}
return out
}
// recordUsage 汇总一次请求的用量事件并异步落库。tok 为从上游响应提取的用量。
// cand/ch 可为 nil(无可用渠道的失败场景)。
func ( g * Gateway ) recordUsage ( req * Request , cand * channel . Candidate , ch * store . Channel , ev usage . Event , tok convert . TokenUsage ) {
if g . usageRec == nil {
return
}
// Set headers
g . setHeaders ( httpReq , ch , apiKey , targetFormat )
// Execute request
start := time . Now ( )
resp , err := g . httpClient . Do ( httpReq )
latency := time . Since ( start )
if err != nil {
g . channelSvc . RecordFailure ( ch . ID )
g . writeError ( c , http . StatusBadGateway , fmt . Sprintf ( "upstream error: %v (latency: %v)" , err , latency ) )
return
ev . UserID = req . UserID
ev . ModelName = req . Model
ev . Protocol = req . Protocol
ev . RequestID = req . RequestID
if req . APIKey != nil {
ev . KeyID = req . APIKey . ID
}
defer resp . Body . Close ( )
// Record success
g . channelSvc . RecordSuccess ( ch . ID )
// Handle response
if resp . StatusCode >= 400 {
body , _ := io . ReadAll ( resp . Body )
log . Printf ( "Upstream error: status=%d body=%s" , resp . StatusCode , string ( body ) )
c . Data ( resp . StatusCode , "application/json" , body )
return
if ch != nil {
ev . ChannelID = ch . ID
}
// Stream or buffer response
if req . Stream {
g . streamResponse ( c , resp , req . Protocol , targetFormat )
} else {
g . bufferResponse ( c , resp , req . Protocol , targetFormat )
if cand != nil && cand . Binding != nil {
ev . ModelID = cand . Binding . ModelID
}
ev . PromptTokens = tok . InputTokens
ev . CompletionTokens = tok . OutputTokens
ev . CacheReadTokens = tok . CacheReadTokens
ev . CacheCreationTokens = tok . CacheCreationTokens
// 定价与成本(价格按每百万 token 的 USD 单价)。
// 成本口径:非缓存输入 × 输入价 + 缓存读 × 缓存价 + 缓存写与输出 × 输出价。
if ev . ModelID != 0 {
if m , err := g . modelDAO . GetByID ( ev . ModelID ) ; err == nil {
ev . InputPrice = m . InputPrice
ev . OutputPrice = m . OutputPrice
ev . CacheReadPrice = m . CacheReadPrice
}
}
if ! ev . IsError {
ev . Cost = ( float64 ( ev . PromptTokens - ev . CacheReadTokens ) * ev . InputPrice +
float64 ( ev . CacheReadTokens ) * ev . CacheReadPrice +
float64 ( ev . CacheCreationTokens ) * ev . OutputPrice +
float64 ( ev . CompletionTokens ) * ev . OutputPrice ) / 1e6
}
g . usageRec . Record ( ev )
}
// conversionTarget 决定客户端协议在渠道上的处理方式:
@@ -264,7 +384,8 @@ func (g *Gateway) setHeaders(req *http.Request, ch *store.Channel, apiKey string
// streamResponse 流式响应:按 \n\n 分块零缓冲转发;跨协议时逐行转换。
func ( g * Gateway ) streamResponse ( c * gin . Context , resp * http . Response , clientProto , upstreamProto string ) {
// 返回从上游 SSE usage 块累计的 token 用量(按上游协议解析)。
func ( g * Gateway ) streamResponse ( c * gin . Context , resp * http . Response , clientProto , upstreamProto string ) convert . TokenUsage {
w := c . Writer
c . Header ( "Content-Type" , "text/event-stream" )
c . Header ( "Cache-Control" , "no-cache" )
@@ -280,7 +401,9 @@ func (g *Gateway) streamResponse(c *gin.Context, resp *http.Response, clientProt
}
// 上游原始行按 \n\n 分块,避免把 data 行内的转义换行当成事件边界。
// 同时喂入用量累计器(usage 块可能出现在任一事件)。
r := bufio . NewReaderSize ( resp . Body , 32 * 1024 )
accum := convert . NewStreamUsageAccum ( )
for {
buf := [ ] byte { }
for {
@@ -292,20 +415,25 @@ func (g *Gateway) streamResponse(c *gin.Context, resp *http.Response, clientProt
buf = append ( buf , line ... )
if err == io . EOF {
if len ( buf ) == 0 {
return
return accum . Usage ( )
}
if ! bytes . HasSuffix ( buf , [ ] byte ( "\n" ) ) {
buf = append ( buf , '\n' )
}
} else if err != nil {
log . Printf ( "stream read error: %v" , err )
return
return accum . Usage ( )
}
if len ( buf ) >= 2 && bytes . HasSuffix ( buf , [ ] byte ( "\n\n" ) ) {
break
}
}
// 先解析用量(data: {...} 行),再决定转发内容。
for _ , data := range sseDataPayloads ( buf ) {
accum . Feed ( data , upstreamProto )
}
out := buf
if lineConv != nil {
out = lineConv ( buf )
@@ -314,21 +442,67 @@ func (g *Gateway) streamResponse(c *gin.Context, resp *http.Response, clientProt
continue
}
if _ , err := w . Write ( out ) ; err != nil {
return // 客户端已断开
return accum . Usage ( ) // 客户端已断开
}
if flusher != nil {
flusher . Flush ( )
}
// 流结束标记:chat/messages 上游以 data: [DONE] 收尾。部分上游(keep-alive)
// 发完 [DONE] 后不关连接,继续读会阻塞到超时;据此主动收尾。
// responses 协议没有 [DONE],以 response.completed 事件收尾。
if streamTerminated ( buf , upstreamProto ) {
return accum . Usage ( )
}
}
}
func ( g * Gateway ) bufferResponse ( c * gin . Context , resp * http . Response , clientProto , upstreamProto string ) {
// streamTerminated 判断一块 SSE 是否为上游流的结束事件。
func streamTerminated ( chunk [ ] byte , proto string ) bool {
switch proto {
case convert . ProtoChat :
// chat 上游以 data: [DONE] 收尾;keep-alive 上游发完不关连接。
return bytes . Contains ( chunk , [ ] byte ( "data: [DONE]" ) )
case convert . ProtoMessages :
// messages 上游以 message_stop 事件结束(无 [DONE])。
return bytes . Contains ( chunk , [ ] byte ( ` "type":"message_stop" ` ) ) ||
bytes . Contains ( chunk , [ ] byte ( ` "type": "message_stop" ` ) ) ||
bytes . Contains ( chunk , [ ] byte ( "data: [DONE]" ) )
case convert . ProtoResponses :
return bytes . Contains ( chunk , [ ] byte ( ` "response.completed" ` ) ) ||
bytes . Contains ( chunk , [ ] byte ( ` "type":"response.completed" ` ) )
}
return false
}
// sseDataPayloads 从一块 SSE(一个完整事件,\n\n 结尾)中取出所有 data 行的原始载荷。
func sseDataPayloads ( chunk [ ] byte ) [ ] [ ] byte {
var out [ ] [ ] byte
for _ , line := range bytes . Split ( chunk , [ ] byte ( "\n" ) ) {
line = bytes . TrimSuffix ( line , [ ] byte ( "\r" ) )
if ! bytes . HasPrefix ( line , [ ] byte ( "data:" ) ) {
continue
}
payload := bytes . TrimPrefix ( line , [ ] byte ( "data:" ) )
payload = bytes . TrimPrefix ( payload , [ ] byte ( " " ) )
if len ( payload ) == 0 || bytes . Equal ( payload , [ ] byte ( "[DONE]" ) ) {
continue
}
out = append ( out , payload )
}
return out
}
func ( g * Gateway ) bufferResponse ( c * gin . Context , resp * http . Response , clientProto , upstreamProto string ) convert . TokenUsage {
body , err := io . ReadAll ( resp . Body )
if err != nil {
g . writeError ( c , http . StatusBadGateway , "failed to read response" )
return
return convert . TokenUsage { }
}
// 用量从上游原始响应体提取(先于转换,转换会改字段名)。
tok , _ := convert . ExtractUsageJSON ( body , upstreamProto )
out := body
if upstreamProto != clientProto {
if converted , cerr := convert . ConvertResponse ( body , upstreamProto , clientProto ) ; cerr == nil {
@@ -342,6 +516,7 @@ func (g *Gateway) bufferResponse(c *gin.Context, resp *http.Response, clientProt
out = convert . CleanJSON ( body )
}
c . Data ( resp . StatusCode , "application/json" , out )
return tok
}
func ( g * Gateway ) writeError ( c * gin . Context , status int , message string ) {