Files
openteam/server/internal/proxy/passthrough.go
T
Sakurasan 360c6b33a6 M0+M1: 基建 + 用户/密钥/核心代理
后端 (Go/Gin/GORM):
- 配置(viper+env)、SQLite/Postgres 迁移、argon2id、AES-GCM 渠道密钥、JWT+refresh cookie
- 用户注册/登录/刷新/登出、API Key CRUD(仅存哈希、明文一次展示)
- 代理网关: /v1/chat/completions、/v1/responses、/v1/models 直通 OpenAI 渠道
  非流式+流式(SSE 零缓冲转发), 用量捕获(chat 末块/responses completed 嵌套),
  OpenAI 错误格式(401/402/404/502), 余额检查
- 异步批量记账 + 余额流水 + 日聚合, admin 用户/余额/配置 API
- 单测: crypto/jwt/apikey/流式 usage 提取

前端 (Vue3+TS+Vite+Tailwind v4):
- taste-skill 设计 tokens: 深色仪表盘, 石墨+信号铜色, Outfit+JetBrains Mono
- Landing/登录/注册, 控制台(仪表盘图表/密钥管理/用量明细)
- 基础组件 Button/Input/Badge/Modal, ECharts 用量图

部署: docker-compose(nginx+api+postgres), 双 Dockerfile, nginx SSE 反代
联调: scripts/mockupstream 本地 mock 上游, 端到端验证通过
2026-08-15 13:10:47 +08:00

383 lines
11 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
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() }