package proxy import ( "bufio" "bytes" "context" "crypto/rand" "encoding/hex" "encoding/json" "fmt" "io" "net/http" "strconv" "strings" "time" "github.com/gin-gonic/gin" "github.com/openteam/server/internal/store" ) // newTraceID 生成请求 trace(用于日志与记账幂等 ref)。 func newTraceID() string { b := make([]byte, 8) _, _ = rand.Read(b) return hex.EncodeToString(b) } // bodyReq 统一取出请求体并解析 model / stream 字段。 type bodyReq struct { Model string `json:"model"` Stream bool `json:"stream"` } // parseBody 读取并回填请求体,解析 model/stream。 func parseBody(c *gin.Context) (*bodyReq, []byte, error) { body, err := io.ReadAll(c.Request.Body) if err != nil { return nil, nil, err } c.Request.Body = io.NopCloser(bytes.NewReader(body)) br := &bodyReq{} _ = json.Unmarshal(body, br) // 解析失败按空处理,直通仍可转发 return br, body, nil } // upstreamURL 组装上游地址:base_url + 客户端路径(/v1/chat/completions 等)。 func upstreamURL(ch *store.Channel, path string) string { base := strings.TrimRight(ch.BaseURL, "/") return base + path } // doPassthrough 通用直通:替换 Authorization 为渠道密钥,转发请求。 // convert 回调用于改写请求体(M1 直通为原样;M3 转换时改写)。 func (g *Gateway) doPassthrough(c *gin.Context, ch *store.Channel, path string, body []byte, stream bool, outUsage func(usageRaw json.RawMessage)) { upKey, err := g.ch.UpstreamKey(ch) if err != nil { openAIError(c, http.StatusInternalServerError, "channel_error", "failed to decrypt channel key") return } upBody := body // 流式 chat:注入 stream_options.include_usage,保证末块带 usage(OpenAI 行为) if stream && path == "/v1/chat/completions" && !bytes.Contains(upBody, []byte(`"include_usage"`)) { var m map[string]any if json.Unmarshal(upBody, &m) == nil { m["stream_options"] = map[string]any{"include_usage": true} if b, err := json.Marshal(m); err == nil { upBody = b } } } ctx, cancel := context.WithTimeout(c.Request.Context(), time.Duration(ch.TimeoutMS)*time.Millisecond) defer cancel() req, err := http.NewRequestWithContext(ctx, http.MethodPost, upstreamURL(ch, path), bytes.NewReader(upBody)) if err != nil { openAIError(c, http.StatusInternalServerError, "internal_error", "failed to build upstream request") return } req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+upKey) req.Header.Set("Accept", c.GetHeader("Accept")) if ua := c.GetHeader("User-Agent"); ua != "" { req.Header.Set("User-Agent", ua) } // 透传 OpenAI 生态请求头(组织/项目等) for _, h := range []string{"OpenAI-Organization", "OpenAI-Project", "OpenAI-Beta"} { if v := c.GetHeader(h); v != "" { req.Header.Set(h, v) } } start := time.Now() resp, err := g.hc.Do(req) if err != nil { status := http.StatusBadGateway msg := "Upstream request failed: " + err.Error() if ctx.Err() == context.DeadlineExceeded { status = http.StatusGatewayTimeout msg = "Upstream request timed out" } openAIError(c, status, "upstream_error", msg) g.recordError(c, ch, nil, start, "upstream_error") return } defer resp.Body.Close() // 非 2xx:透传上游错误体(OpenAI 格式),并记录 error 用量 if resp.StatusCode < 200 || resp.StatusCode >= 300 { errBody, _ := io.ReadAll(resp.Body) status := resp.StatusCode // 上游 5xx → 网关 502/504(重试逻辑 M4) if status >= 500 { status = http.StatusBadGateway } c.DataFromReader(status, int64(len(errBody)), "application/json", bytes.NewReader(errBody), nil) c.Header("Content-Type", "application/json") g.recordError(c, ch, resp, start, "upstream_http_"+strconv.Itoa(resp.StatusCode)) return } // 成功响应 c.Header("Content-Type", resp.Header.Get("Content-Type")) c.Status(http.StatusOK) if stream { g.streamCopy(c, ch, resp.Body, start, outUsage) } else { g.copyAndCapture(c, ch, resp.Body, start, outUsage) } } // copyAndCapture 非流式:整体转发 + 解析 usage + 记账。 func (g *Gateway) copyAndCapture(c *gin.Context, ch *store.Channel, r io.Reader, start time.Time, outUsage func(json.RawMessage)) { data, err := io.ReadAll(r) if err != nil { openAIError(c, http.StatusBadGateway, "upstream_error", "failed reading upstream response") g.recordError(c, ch, nil, start, "read_error") return } // 尝试解析 usage(chat / responses 字段不同) if usageRaw := extractUsage(data); usageRaw != nil { outUsage(usageRaw) } _, _ = c.Writer.Write(data) g.finishUsage(c, ch, start, store.UsageStatusSuccess, "") } // streamCopy 流式:边读上游 SSE 边写客户端,零缓冲转发;扫描 usage 行记账。 // 客户端断连(ctx cancel)即中止上游读取。 func (g *Gateway) streamCopy(c *gin.Context, ch *store.Channel, r io.Reader, start time.Time, outUsage func(json.RawMessage)) { w := c.Writer flusher, ok := w.(http.Flusher) if !ok { flusher = nopFlusher{} } scanner := newSSEScanner(r) for { line, err := scanner.Next() if line != nil { if _, werr := w.Write(line); werr != nil { // 客户端断开:取消上游(ctx cancel 由 request ctx 处理) g.recordError(c, ch, nil, start, "client_disconnect") return } flusher.Flush() if usageRaw := scanUsage(line); usageRaw != nil { outUsage(usageRaw) } } if err != nil { if err == io.EOF { g.finishUsage(c, ch, start, store.UsageStatusSuccess, "") } else { g.recordError(c, ch, nil, start, "stream_read_error") } return } } } type nopFlusher struct{} func (nopFlusher) Flush() {} // --------------------------------------------------------------------------- // usage 提取 // usageShape 兼容 chat (prompt/completion) 与 responses (input/output) 两种命名。 type usageShape struct { PromptTokens int64 `json:"prompt_tokens"` CompletionTokens int64 `json:"completion_tokens"` InputTokens int64 `json:"input_tokens"` OutputTokens int64 `json:"output_tokens"` TotalTokens int64 `json:"total_tokens"` // Claude 缓存口径(M3 接入) CacheReadInputTokens int64 `json:"cache_read_input_tokens"` CacheCreationInputTokens int64 `json:"cache_creation_input_tokens"` } // extractUsage 从完整响应体提取 usage 子对象。 func extractUsage(data []byte) json.RawMessage { var m map[string]json.RawMessage if json.Unmarshal(data, &m) != nil { return nil } if u, ok := m["usage"]; ok && string(u) != "null" { return u } // responses 事件/响应:usage 嵌套在 response 对象内 if respRaw, ok := m["response"]; ok { var resp map[string]json.RawMessage if json.Unmarshal(respRaw, &resp) == nil { if u, ok := resp["usage"]; ok && string(u) != "null" { return u } } } // chat 兜底:choices[].message.usage if choices, ok := m["choices"]; ok { var cs []map[string]json.RawMessage if json.Unmarshal(choices, &cs) == nil { for _, ch := range cs { if msgRaw, ok := ch["message"]; ok { var msg map[string]json.RawMessage if json.Unmarshal(msgRaw, &msg) == nil { if u, ok := msg["usage"]; ok && string(u) != "null" { return u } } } } } } return nil } // scanUsage 从 SSE 一行中提取 usage(OpenAI 末块 / responses completed 事件)。 func scanUsage(line []byte) json.RawMessage { s := string(line) if !strings.Contains(s, `"usage"`) { return nil } if strings.HasPrefix(s, "data: ") { s = strings.TrimPrefix(s, "data: ") } s = strings.TrimSpace(s) if s == "[DONE]" || s == "" { return nil } var m map[string]json.RawMessage if json.Unmarshal([]byte(s), &m) != nil { return nil } if u, ok := m["usage"]; ok && string(u) != "null" { return u } // responses 流式:usage 在 response 对象内(response.completed 事件) if respRaw, ok := m["response"]; ok { var resp map[string]json.RawMessage if json.Unmarshal(respRaw, &resp) == nil { if u, ok := resp["usage"]; ok && string(u) != "null" { return u } } } return nil } // sseScanner 按 SSE 行边界读取(兼容 \n 与 \r\n),保留原始行内容。 // 基于 bufio.Reader:行内可含任意内容,跨 chunk 自动拼接。 type sseScanner struct { r *bufio.Reader } func newSSEScanner(r io.Reader) *sseScanner { return &sseScanner{r: bufio.NewReaderSize(r, 32*1024)} } func (s *sseScanner) Next() ([]byte, error) { line, err := s.r.ReadBytes('\n') if len(line) > 0 { return line, nil } if err != nil { return nil, err } return nil, io.EOF } // --------------------------------------------------------------------------- // 记账 // usageSink 累积流式多次 usage(取最后一次,即最终值)。 type usageSink struct { last json.RawMessage } func (u *usageSink) push(raw json.RawMessage) { if len(raw) > 0 { u.last = raw } } // finishUsage 落账:计算成本并异步写入。 func (g *Gateway) finishUsage(c *gin.Context, ch *store.Channel, start time.Time, status, errCode string) { uid, _ := c.Get(CtxUserID) kid, _ := c.Get(CtxKeyID) trace, _ := c.Get(CtxTrace) var us usageShape if h, ok := c.Get("usage_raw"); ok { if holder, ok := h.(*sinkHolder); ok && holder.sink != nil && len(holder.sink.last) > 0 { _ = json.Unmarshal(holder.sink.last, &us) } } in := us.PromptTokens + us.InputTokens out := us.CompletionTokens + us.OutputTokens cacheRead := us.CacheReadInputTokens cacheCreate := us.CacheCreationInputTokens modelName, _ := c.Get("model_name") mn, _ := modelName.(string) var model store.Model var cost float64 var modelID uint64 _ = g.db.Where("name = ?", mn).First(&model).Error if model.ID > 0 { modelID = model.ID cost = float64(in)/1e6*model.InputPrice + float64(out)/1e6*model.OutputPrice + float64(cacheRead)/1e6*model.CacheReadPrice } else { cost = float64(in)/1e6*0.15 + float64(out)/1e6*0.60 // 无定价模型时按示例价 } proto, _ := c.Get("protocol") p, _ := proto.(string) if p == "" { p = "chat" } traceStr, _ := trace.(string) errMsg := errCode latency := int(time.Since(start).Milliseconds()) // 已写响应头但流中途出错:记 error if status == store.UsageStatusSuccess && c.Writer.Status() >= 400 { status = store.UsageStatusError } g.rec.Record(&store.UsageLog{ RequestID: fmt.Sprintf("trace-%s", traceStr), TraceID: traceStr, UserID: uid.(uint64), KeyID: kid.(uint64), ChannelID: ch.ID, ModelID: modelID, ModelName: mn, Protocol: p, InputTokens: in, OutputTokens: out, CacheReadTokens: cacheRead, CacheCreationTokens: cacheCreate, InputPrice: model.InputPrice, OutputPrice: model.OutputPrice, CacheReadPrice: model.CacheReadPrice, Cost: cost, LatencyMS: latency, Status: status, ErrorCode: &errMsg, CreatedAt: time.Now().UTC(), }) } // recordError 失败请求的记账(不产生扣费,status=error)。 func (g *Gateway) recordError(c *gin.Context, ch *store.Channel, resp *http.Response, start time.Time, code string) { status := store.UsageStatusError _ = resp g.finishUsage(c, ch, start, status, code) } func now() time.Time { return time.Now() }