- 新增 tokenizer 包(tiktoken-go)按模型估算 token, 未知模型回退 cl100k_base - passthrough 提取请求输入文本 + SSE 已生成内容, 客户端断开时记 canceled - usage 计费范围扩展: canceled(流式中断)按已生成部分收费 Co-Authored-By: Claude <noreply@anthropic.com>
154 lines
4.1 KiB
Go
154 lines
4.1 KiB
Go
// Package usage 异步记账:请求完成后写入 usage_logs,批量落库(PLANNING §3.2)。
|
||
// 每个请求在 flush 时同步完成:写明细 + 扣余额 + 写流水 + 日聚合。
|
||
package usage
|
||
|
||
import (
|
||
"log"
|
||
"sync"
|
||
"time"
|
||
|
||
"github.com/openteam/server/internal/store"
|
||
"gorm.io/gorm"
|
||
"gorm.io/gorm/clause"
|
||
)
|
||
|
||
// Recorder 异步记账器:缓冲队列 + 批量事务落库。
|
||
type Recorder struct {
|
||
db *gorm.DB
|
||
ch chan *store.UsageLog
|
||
wg sync.WaitGroup
|
||
closed chan struct{}
|
||
}
|
||
|
||
const batchSize = 32
|
||
|
||
func NewRecorder(db *gorm.DB) *Recorder {
|
||
r := &Recorder{
|
||
db: db,
|
||
ch: make(chan *store.UsageLog, 512),
|
||
closed: make(chan struct{}),
|
||
}
|
||
r.wg.Add(1)
|
||
go r.run()
|
||
return r
|
||
}
|
||
|
||
// Record 提交一条用量(非阻塞;队列满时同步写入,保证不丢账)。
|
||
func (r *Recorder) Record(l *store.UsageLog) {
|
||
select {
|
||
case r.ch <- l:
|
||
default:
|
||
if err := r.flush([]*store.UsageLog{l}); err != nil {
|
||
log.Printf("usage: sync write failed: %v", err)
|
||
}
|
||
}
|
||
}
|
||
|
||
func (r *Recorder) Close() {
|
||
close(r.closed)
|
||
r.wg.Wait()
|
||
close(r.ch)
|
||
}
|
||
|
||
func (r *Recorder) run() {
|
||
defer r.wg.Done()
|
||
buf := make([]*store.UsageLog, 0, batchSize)
|
||
tick := time.NewTicker(2 * time.Second)
|
||
defer tick.Stop()
|
||
for {
|
||
select {
|
||
case l, ok := <-r.ch:
|
||
if !ok {
|
||
return
|
||
}
|
||
buf = append(buf, l)
|
||
if len(buf) >= batchSize {
|
||
if err := r.flush(buf); err != nil {
|
||
log.Printf("usage: batch write failed: %v", err)
|
||
}
|
||
buf = buf[:0]
|
||
}
|
||
case <-r.closed:
|
||
if len(buf) > 0 {
|
||
if err := r.flush(buf); err != nil {
|
||
log.Printf("usage: final batch write failed: %v", err)
|
||
}
|
||
}
|
||
return
|
||
case <-tick.C:
|
||
if len(buf) > 0 {
|
||
if err := r.flush(buf); err != nil {
|
||
log.Printf("usage: batch write failed: %v", err)
|
||
}
|
||
buf = buf[:0]
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
// flush 批量插入用量明细,并同步更新余额、余额流水与日聚合。
|
||
// 记账口径:单次成本 = in×in_price + out×out_price + cache_read×cache_read_price(每百万 token)。
|
||
func (r *Recorder) flush(logs []*store.UsageLog) error {
|
||
if len(logs) == 0 {
|
||
return nil
|
||
}
|
||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||
if err := tx.Create(logs).Error; err != nil {
|
||
return err
|
||
}
|
||
for _, l := range logs {
|
||
// 计费范围:success(正常完成)与 canceled(流式中断,按已生成部分收费)
|
||
if (l.Status != store.UsageStatusSuccess && l.Status != store.UsageStatusCanceled) || l.Cost <= 0 {
|
||
continue
|
||
}
|
||
// 扣余额(余额可为负:流式请求不中断;后续请求被拒)
|
||
var user store.User
|
||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&user, l.UserID).Error; err != nil {
|
||
continue
|
||
}
|
||
newBalance := user.Balance - l.Cost
|
||
if err := tx.Model(&store.User{}).Where("id = ?", l.UserID).Update("balance", newBalance).Error; err != nil {
|
||
continue
|
||
}
|
||
tx.Create(&store.BalanceLog{
|
||
UserID: l.UserID,
|
||
Change: -l.Cost,
|
||
BalanceAfter: newBalance,
|
||
Type: store.BalanceTypeUsage,
|
||
RefID: usageRefID(l.TraceID),
|
||
Remark: "usage: " + l.ModelName,
|
||
})
|
||
|
||
// 日聚合 upsert
|
||
date := l.CreatedAt.UTC().Format("2006-01-02")
|
||
tx.Clauses(clause.OnConflict{
|
||
Columns: []clause.Column{{Name: "user_id"}, {Name: "model_id"}, {Name: "date"}},
|
||
DoUpdates: clause.Assignments(map[string]any{
|
||
"requests": gorm.Expr("requests + 1"),
|
||
"input_tokens": gorm.Expr("input_tokens + ?", l.InputTokens),
|
||
"output_tokens": gorm.Expr("output_tokens + ?", l.OutputTokens),
|
||
"cache_read_tokens": gorm.Expr("cache_read_tokens + ?", l.CacheReadTokens),
|
||
"cost": gorm.Expr("cost + ?", l.Cost),
|
||
}),
|
||
}).Create(&store.UsageDaily{
|
||
UserID: l.UserID,
|
||
ModelID: l.ModelID,
|
||
Date: date,
|
||
Requests: 1,
|
||
InputTokens: l.InputTokens,
|
||
OutputTokens: l.OutputTokens,
|
||
CacheReadTokens: l.CacheReadTokens,
|
||
Cost: l.Cost,
|
||
})
|
||
}
|
||
return nil
|
||
})
|
||
}
|
||
|
||
func usageRefID(traceID string) string {
|
||
if traceID == "" {
|
||
traceID = "unknown"
|
||
}
|
||
return "usage:" + traceID
|
||
}
|