180 lines
4.9 KiB
Go
180 lines
4.9 KiB
Go
package agent
|
||
|
||
// ========================================================================
|
||
// TokenBudget —— 单次 AI 任务的 Token 预算累加器
|
||
// ========================================================================
|
||
// 作用:
|
||
// 给一次 AI 生成任务设置"总 token 上限",跨多轮 LLM 调用累加,
|
||
// 超过预算时主动中止,避免:
|
||
// 1. Agent 死循环/工具反复触发导致单请求烧几十万 token(成本失控)
|
||
// 2. 模型无限复读(罕见但发生过,会撑爆 max_tokens 之外的总账)
|
||
// 3. 用户构造超长 prompt 攻击
|
||
//
|
||
// 与厂商 max_tokens 的区别:
|
||
// - 厂商 max_tokens:控制"单次响应"输出 token 数,无法跨轮累加
|
||
// - TokenBudget :控制"整次任务"累计 token 数(prompt+completion 一起算),
|
||
// 是 Agent 客户端层面的二级防护
|
||
//
|
||
// 使用方式:
|
||
// budget := NewTokenBudget(8000) // 后台配置 ai_token_budget_per_request
|
||
// for {
|
||
// maxTok := budget.CalcMaxTokensForCall(cfg.MaxPerCall)
|
||
// result := client.ChatWithOpts(ctx, msgs, tools, ChatOpts{MaxTokens: maxTok})
|
||
// budget.Consume(result.PromptTokens, result.CompletionTokens)
|
||
// if budget.IsExceeded() { break } // 触发软中止
|
||
// }
|
||
// ========================================================================
|
||
|
||
import (
|
||
"fmt"
|
||
"sync"
|
||
)
|
||
|
||
// ErrBudgetExceeded 预算超限错误(ReactLoop 据此判定软中止)
|
||
//
|
||
// 调用方应该捕获本错误后返回"已生成的部分内容",而不是直接报错给用户
|
||
var ErrBudgetExceeded = fmt.Errorf("token budget exceeded")
|
||
|
||
// TokenBudget Token 预算累加器(线程安全)
|
||
type TokenBudget struct {
|
||
mu sync.Mutex
|
||
used int // 已用 token 数(prompt + completion 累加)
|
||
limit int // 上限(0 表示不限制)
|
||
exceeded bool // 是否已超限
|
||
maxPerCall int // 单次调用最大输出 token 数(用于 CalcMaxTokensForCall)
|
||
}
|
||
|
||
// NewTokenBudget 创建预算器
|
||
//
|
||
// 参数:
|
||
// - limit - 单任务总 token 上限(0 表示不限)
|
||
// - maxPerCall- 单次 LLM 调用输出 token 上限(透传厂商 max_tokens)
|
||
func NewTokenBudget(limit, maxPerCall int) *TokenBudget {
|
||
if limit < 0 {
|
||
limit = 0
|
||
}
|
||
if maxPerCall < 0 {
|
||
maxPerCall = 0
|
||
}
|
||
return &TokenBudget{
|
||
limit: limit,
|
||
maxPerCall: maxPerCall,
|
||
}
|
||
}
|
||
|
||
// Consume 累加一次 LLM 调用的 token 用量
|
||
//
|
||
// 参数:
|
||
// - promptTokens - 本次输入 token 数
|
||
// - completionTokens - 本次输出 token 数
|
||
func (b *TokenBudget) Consume(promptTokens, completionTokens int) {
|
||
if b == nil {
|
||
return
|
||
}
|
||
b.mu.Lock()
|
||
defer b.mu.Unlock()
|
||
b.used += promptTokens + completionTokens
|
||
if b.limit > 0 && b.used > b.limit {
|
||
b.exceeded = true
|
||
}
|
||
}
|
||
|
||
// IsExceeded 是否已超限
|
||
func (b *TokenBudget) IsExceeded() bool {
|
||
if b == nil {
|
||
return false
|
||
}
|
||
b.mu.Lock()
|
||
defer b.mu.Unlock()
|
||
return b.exceeded
|
||
}
|
||
|
||
// Used 已用 token 数
|
||
func (b *TokenBudget) Used() int {
|
||
if b == nil {
|
||
return 0
|
||
}
|
||
b.mu.Lock()
|
||
defer b.mu.Unlock()
|
||
return b.used
|
||
}
|
||
|
||
// Remaining 剩余可用 token 数(limit=0 时返回 -1 表示不限)
|
||
func (b *TokenBudget) Remaining() int {
|
||
if b == nil {
|
||
return -1
|
||
}
|
||
b.mu.Lock()
|
||
defer b.mu.Unlock()
|
||
if b.limit == 0 {
|
||
return -1
|
||
}
|
||
r := b.limit - b.used
|
||
if r < 0 {
|
||
return 0
|
||
}
|
||
return r
|
||
}
|
||
|
||
// CalcMaxTokensForCall 计算下一次 LLM 调用应该传给厂商的 max_tokens 值
|
||
//
|
||
// 策略:取 min(配置的单次上限, 剩余预算)
|
||
// - 配置的单次上限:ai_token_max_per_call(如 2048)
|
||
// - 剩余预算:limit - used(如剩余 1000,则这次最多 1000,否则必然超)
|
||
//
|
||
// 注意:剩余预算只算 completion 部分(厂商 max_tokens 限制的就是 completion),
|
||
// 但这里无法预知下一次的 prompt token,所以是粗略估算(保守取值,宁可少生成)
|
||
func (b *TokenBudget) CalcMaxTokensForCall(cfgMaxPerCall int) int {
|
||
if b == nil {
|
||
return cfgMaxPerCall
|
||
}
|
||
b.mu.Lock()
|
||
defer b.mu.Unlock()
|
||
|
||
// 没设上限 → 用 cfg 值
|
||
if b.limit == 0 {
|
||
if cfgMaxPerCall > 0 {
|
||
return cfgMaxPerCall
|
||
}
|
||
return 2048 // 兜底
|
||
}
|
||
|
||
// 有上限:取 min(cfgMaxPerCall, 剩余)
|
||
upper := cfgMaxPerCall
|
||
if upper <= 0 {
|
||
upper = b.maxPerCall
|
||
}
|
||
if upper <= 0 {
|
||
upper = 2048
|
||
}
|
||
remain := b.limit - b.used
|
||
if remain < upper {
|
||
if remain < 1 {
|
||
return 1 // 至少给 1,避免厂商报错"max_tokens 必须 > 0"
|
||
}
|
||
return remain
|
||
}
|
||
return upper
|
||
}
|
||
|
||
// Snapshot 取一份只读快照(用于日志/返回给上层)
|
||
func (b *TokenBudget) Snapshot() BudgetSnapshot {
|
||
if b == nil {
|
||
return BudgetSnapshot{}
|
||
}
|
||
b.mu.Lock()
|
||
defer b.mu.Unlock()
|
||
return BudgetSnapshot{
|
||
Used: b.used,
|
||
Limit: b.limit,
|
||
Exceeded: b.exceeded,
|
||
}
|
||
}
|
||
|
||
// BudgetSnapshot 预算快照(只读)
|
||
type BudgetSnapshot struct {
|
||
Used int `json:"used"` // 已用 token
|
||
Limit int `json:"limit"` // 上限(0=不限)
|
||
Exceeded bool `json:"exceeded"` // 是否超限
|
||
}
|