Files
sorrow404null a7412824e0 chore: import original project baseline
Import the pre-repair source tree as the history baseline.
Runtime data (data/), virtualenvs, bytecode caches and logs are
gitignored so local secrets and user state stay out of the repo.
2026-09-17 16:40:01 +08:00

342 lines
15 KiB
Python
Raw Permalink 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.
"""
core/agent/context.py
=====================
🌟 pi 上下文管理的 Python 1:1 移植 —— 令牌估算 + 输出预算钳制
对照 pi-main 源码:
packages/ai/src/api/simple-options.ts
- CONTEXT_SAFETY_TOKENS = 4096 (行 12)
- clampMaxTokensToContext() (行 12-34)
maxOutput = contextWindow - inputTokens - 4096
return min(model.maxTokens, maxOutput) (下限 1)
- clampOutputsToContext() (行 36-60)
多模型共享窗口时按比例分配剩余输出预算
packages/ai/src/api/token-utils.ts
- 每条消息估算: ceil(chars/4) + 4
🌟 已声明偏差(仅 1 处,其余全部 1:1):
pi 按 chars/4 估算(针对英文)。中文 1 字 ≈ 1 token,chars/4 会严重低估。
这里采用 CJK 感知估算:CJK 字符按 1 token/字,其余按 4 字符/token。
这直接影响「压缩触发时机」和「输出预算钳制」,必须更准才不会撞 vLLM 上限。
pi 的 usage 锚定机制 1:1 保留:若历史里存在有效 assistant usage
总估算 = 该 usage 的 totalTokens(服务商精确值)+ 其后消息的逐条估算。
"""
from __future__ import annotations
import math
import re
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Tuple
from .types import AgentMessage, ModelConfig
# ======================================================================
# 对照 simple-options.ts: export const CONTEXT_SAFETY_TOKENS = 4096
# 用途:上下文钳制时预留的安全余量(tokenizer 估算误差 + vLLM 精确计数的偏差)
# ======================================================================
CONTEXT_SAFETY_TOKENS = 4096
# 图片 token 估算(对照 pi 对 image part 的固定估算)
IMAGE_TOKENS = 1600
_CJK_RE = re.compile(
r"[\u2e80-\u2eff\u3040-\u30ff\u3400-\u4dbf\u4e00-\u9fff"
r"\uf900-\ufaff\uff00-\uffef]"
)
def _count_text_chars(value: Any) -> int:
"""统计内容里的字符数(对照 pi token-utils 的字符统计口径)"""
if isinstance(value, str):
return len(value)
if isinstance(value, list): # OpenAI 多模态 content 数组
n = 0
for part in value:
if isinstance(part, dict):
if part.get("type") == "text":
n += len(part.get("text", "") or "")
elif part.get("type") == "image_url":
n += 0 # 图片单独按 IMAGE_TOKENS 计
else:
n += len(str(part))
return n
return len(str(value or ""))
def _count_image_parts(content: Any) -> int:
if isinstance(content, list):
return sum(1 for p in content
if isinstance(p, dict) and p.get("type") == "image_url")
return 0
def estimate_message_tokens(message: AgentMessage) -> int:
"""
单条消息 token 估算(CJK 感知版)。
pi 口径: ceil(chars/4) + 4。本项目: CJK 字 ×1 + 其余 /4 + 4 + 图片×1600。
"""
raw = message.content
if isinstance(raw, str):
chars = len(raw)
cjk = len(_CJK_RE.findall(raw))
else:
# 数组 content:拼接所有文本部分
joined = []
for part in raw if isinstance(raw, list) else []:
if isinstance(part, dict) and part.get("type") == "text":
joined.append(part.get("text", "") or "")
text = "".join(joined)
chars = len(text)
cjk = len(_CJK_RE.findall(text))
tokens = cjk + math.ceil((chars - cjk) / 4) + 4
tokens += _count_image_parts(raw) * IMAGE_TOKENS
# 思考内容与工具参数也占上下文(assistant 的 reasoning / tool_calls
if message.reasoning:
rcjk = len(_CJK_RE.findall(message.reasoning))
tokens += rcjk + math.ceil((len(message.reasoning) - rcjk) / 4)
for tc in message.tool_calls:
arg_text = _json_str(tc.arguments)
tokens += math.ceil(len(arg_text) / 4)
# toolResult 的 content 走 content 字段(str 或 list
return tokens
def _json_str(obj: Any) -> str:
import json
try:
return json.dumps(obj, ensure_ascii=False)
except Exception:
return str(obj)
def _estimate_text_tokens(text: str) -> int:
"""CJK 感知文本 token 估算(与 estimate_message_tokens 同公式,不含每条 +4 开销)"""
if not text:
return 0
cjk = len(_CJK_RE.findall(text))
return cjk + math.ceil((len(text) - cjk) / 4)
def _estimate_tools_tokens(tools) -> int:
"""
🆕 P2: 工具 schema token 估算 —— name/description/parameters 的 JSON 字符数(CJK 感知)。
对照 pi Context 分支: all tools = Σ estimate_text(json(tool))。
"""
total = 0
for t in tools or []:
schema = {
"name": getattr(t, "name", "") or "",
"description": getattr(t, "description", "") or "",
"parameters": getattr(t, "parameters", None),
}
total += _estimate_text_tokens(_json_str(schema))
return total
def estimate_context_tokens(messages: List[AgentMessage],
system_prompt: str = "",
tools: Optional[List[Any]] = None
) -> "ContextUsageEstimate":
"""
🌟 1:1 对照 pi estimateContextTokensusage 锚定):
- 找到最后一条带有效 usage 的 assistant 消息(🆕 P0: 含 timestamp 时效校验)
- 总 tokens = 该 usage 的精确值 + 其后消息的逐条估算
- 无 usage 时 = 全部消息逐条估算
🆕 P2: 可选 system_prompt / tools 参数(对照 pi Context 分支):
- 仅在无锚点分支计入 system 提示词 + 工具 schema 的估算
(有锚点时 usage 已是含 system+tools 的精确总值,不重复计)
逐条估算用 CJK 感知版 estimate_message_tokens(已声明偏差)。
"""
last_idx, last_usage = _find_last_usage(messages)
sys_tokens = _estimate_text_tokens(system_prompt) if system_prompt else 0
tool_tokens = _estimate_tools_tokens(tools) if tools else 0
if last_idx is None:
estimated = (sum(estimate_message_tokens(m) for m in messages)
+ sys_tokens + tool_tokens)
return ContextUsageEstimate(tokens=estimated, usage_tokens=0,
trailing_tokens=estimated, last_usage_index=None)
usage_tokens = calculate_context_tokens(last_usage)
trailing = sum(estimate_message_tokens(m) for m in messages[last_idx + 1:])
return ContextUsageEstimate(tokens=usage_tokens + trailing,
usage_tokens=usage_tokens,
trailing_tokens=trailing,
last_usage_index=last_idx)
@dataclass
class ContextUsageEstimate:
"""
🌟 1:1 对照 pi estimateContextTokens 的返回值
packages/agent/src/harness/compaction/compaction.ts ContextUsageEstimate:
tokens 总估算
usageTokens 最后一条有效 assistant usage 报告的精确 token 数
trailingTokens 该 usage 之后消息的估算值
lastUsageIndex 提供 usage 的消息下标(无则 None
"""
tokens: int = 0
usage_tokens: int = 0
trailing_tokens: int = 0
last_usage_index: Optional[int] = None
def calculate_context_tokens(usage: Dict[str, Any]) -> int:
"""pi calculateContextTokens: totalTokens || input+output+cacheRead+cacheWrite"""
if not usage:
return 0
total = int(usage.get("totalTokens", 0) or 0)
if total:
return total
return (int(usage.get("input", 0) or 0) + int(usage.get("output", 0) or 0)
+ int(usage.get("cacheRead", 0) or 0)
+ int(usage.get("cacheWrite", 0) or 0))
def _get_assistant_usage(msg: AgentMessage) -> Dict[str, Any]:
"""
pi getAssistantUsage:只认「有效」的 assistant usage——
stopReason 不是 aborted/error,且 calculateContextTokens > 0。
(🆕 P0: 锚点扫描还需通过 timestamp 时效校验,见 _find_last_usage
"""
if msg.role != "assistant":
return {}
if msg.stop_reason in ("aborted", "error"):
return {}
usage = msg.usage or {}
if usage and calculate_context_tokens(usage) > 0:
return usage
return {}
def _row_block_tool_flags(messages: List["AgentMessage"]) -> List[bool]:
"""
🆕 G1: 逐条标记「其 DB 行(同非空 db_msg_id 的连续块)是否含工具活动」。
含工具输出的行,其入库 usage 是「本轮最后一个子请求」的内存快照
(工具结果未受 4000 字回放截断),不等于下一请求(截断回放)→ 锚点失效。
轮中内存条目(db_msg_id="")一律 False:保持原规则
(轮中条目的内存载荷 == 下一子请求载荷,锚点有效)。
"""
n = len(messages)
flags = [False] * n
i = 0
while i < n:
rid = getattr(messages[i], "db_msg_id", "") or ""
if not rid:
i += 1
continue
j = i
has_tool = False
while j < n and (getattr(messages[j], "db_msg_id", "") or "") == rid:
mj = messages[j]
if mj.role == "toolResult" or (mj.role == "assistant" and mj.tool_calls):
has_tool = True
j += 1
for k in range(i, j):
flags[k] = has_tool
i = j
return flags
def _find_last_usage(messages: List["AgentMessage"]):
"""
🆕 P0: 1:1 对照 pi getLastAssistantUsageInfo(含锚点失效校验):
assistant 的 usage 要成为锚点,必须满足
① timestamp >= 它前面所有消息的最大 timestamp(即它不比任何前消息「更旧」)。
场景:压缩后 state = [新摘要消息(T_new)] + [保留的旧消息(T_old)]
保留窗里 assistant 的 usage 来自压缩前(如 60K 精确值)——不校验的话,
下一次 should_compact 会锚定到过期值 → 误触发二次压缩(摘要套摘要)。
摘要消息 timestamp 晚于全部保留消息 → 前缀游标抬高 → 旧 usage 自动失效;
压缩后第一个新 assistant 回复到达 → 新 usage 锚点自动恢复。
② 🆕 G1: 该 usage 所在 DB 行不含工具活动(toolResult / assistant.tool_calls)。
含工具输出的行,入库 usage = 本轮内存快照(工具结果未截断)≠ 下一请求
(4000 字截断回放)→ 系统性高估 → 误触发压缩(实测 113904 锚点 vs
真实下一请求 ~90k)。纯文本行 / 轮中内存条目的 usage 仍是精确值,保留锚定。
(消息无 timestamp 信息(全 0)时行为等同旧版「取最后一条有效」,完全向后兼容。)
"""
latest_prefix_ts = -1 # 对照 pi: Number.NEGATIVE_INFINITY
last_idx = None
last_usage: Dict[str, Any] = {}
row_tool_flags = _row_block_tool_flags(messages)
for i, msg in enumerate(messages):
if msg.role == "assistant":
applies = (msg.timestamp or 0) >= latest_prefix_ts
if (applies
and not row_tool_flags[i]
and msg.stop_reason not in ("aborted", "error")
and msg.usage
and calculate_context_tokens(msg.usage) > 0):
last_idx, last_usage = i, msg.usage
ts = msg.timestamp or 0
if ts > latest_prefix_ts:
latest_prefix_ts = ts
return last_idx, last_usage
def clamp_max_tokens_to_context(model: ModelConfig,
context: List[AgentMessage],
system_prompt: str = "",
tools: Optional[List[Any]] = None
) -> Optional[Tuple[int, int]]:
"""
🌟 对照 simple-options.ts:12-34 clampMaxTokensToContext1:1 公式):
input_estimate = estimateContextTokens(context) # 🆕 P2: 含 system+toolspi Context 实参)
if (input_estimate > model.contextWindow) return null // 上下文已溢出
maxOutput = contextWindow - inputTokens - CONTEXT_SAFETY_TOKENS
return max(1, min(model.maxTokens, maxOutput))
这就是 pi 约束公式的本体:
输入 + 输出 ≤ contextWindow 4096 ≤ vLLM max_model_len
🆕 P2: system_prompt / tools 仅在无 usage 锚点分支计入(首轮流式精确记账,
与 pi 传完整 Context 的行为一致);有锚点时 usage 已精确覆盖,不重复计。
返回 (max_tokens, input_tokens);输入已溢出窗口时返回 None(由恢复逻辑接管)。
"""
input_tokens = estimate_context_tokens(
context, system_prompt=system_prompt, tools=tools).tokens
if input_tokens > model.context_window:
return None
max_output = model.context_window - input_tokens - CONTEXT_SAFETY_TOKENS
return max(1, min(model.max_tokens, max_output)), input_tokens
def clamp_outputs_to_context(context: List[AgentMessage],
models: List[ModelConfig],
reserve_tokens: int = CONTEXT_SAFETY_TOKENS
) -> List[int]:
"""
对照 simple-options.ts:36-60 clampOutputsToContext
多个模型共享同一窗口时,把剩余输出预算按比例分配给各模型。
(haocode 单模型场景用不到,保留以求框架完整)
"""
input_tokens = estimate_context_tokens(context).tokens
available = max(0, sum(m.context_window for m in models) - input_tokens - reserve_tokens)
if not models:
return []
weights = [m.max_tokens for m in models]
total_w = sum(weights) or 1
alloc = [max(1, int(available * w / total_w)) for w in weights]
# 各自不超过自身 max_tokens
alloc = [min(a, m.max_tokens) for a, m in zip(alloc, models)]
return alloc
def should_compact(messages: List[AgentMessage], model: ModelConfig,
reserve_tokens: int = 16384,
system_prompt: str = "",
tools: Optional[List[Any]] = None) -> Tuple[bool, int]:
"""
🌟 1:1 对照 pi compaction.ts shouldCompactharness 原版公式):
contextTokens > contextWindow - reserveTokens
reserveTokens 默认 16384DEFAULT_COMPACTION_SETTINGS.reserveTokens),
即「为摘要提示词与输出预留 16K」。估算值取 usage 锚定估算。
🆕 G2: system_prompt/tools 透传 —— 无锚点分支(工具行快照锚点被 G1 失效后
更常走到)也要计入 system + 工具 schema,度量「下一请求真实载荷」,
与显示端(update_context_display)同口径。
"""
tokens = estimate_context_tokens(messages, system_prompt=system_prompt,
tools=tools).tokens
return (tokens > model.context_window - reserve_tokens), tokens