Files
Haocode/tests/test_compaction_persist.py
T
2026-09-17 16:30:02 +08:00

290 lines
16 KiB
Python
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.
# -*- coding: utf-8 -*-
"""
离屏验证:压缩持久化(链上标记点)
背景:压缩此前只改内存 agent.state.messages,下一轮提问从 DB 全量重建历史
→「压缩等于没有压缩」。现:压缩成功后在链上插标记行(role="compaction"
is_ignored=1),保留首条改指到标记;build_api_context 在最后一个标记处截断、
以摘要(user 消息)替代之前一切。历史行全部保留,前端切会话渲染不受影响。
"""
import os, sys, json, tempfile, types
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
os.environ.setdefault("QT_QPA_PLATFORM", "offscreen")
from core.db_manager import DBManager # noqa: E402
from core.agent.types import AgentMessage, AgentConfig, ModelConfig # noqa: E402
from core.agent.stream_fn import from_openai_messages, to_openai_messages # noqa: E402
from core.agent.compaction import prepare_compaction, CompactionSettings # noqa: E402
from core.agent.context import estimate_context_tokens # noqa: E402
from core.agent.recovery import AgentRunner, _cut_ids_of # noqa: E402
from ui.views.main_window import MainWindow # noqa: E402 纯方法 unbound 调用,不实例化
RESULTS = []
def check(name, cond, extra=""):
RESULTS.append((name, bool(cond)))
print(f"{'PASS' if cond else 'FAIL'} {name} {extra if not cond else ''}", flush=True)
TMP = os.path.join(tempfile.gettempdir(), f"haocode_compact_persist_{os.getpid()}.db")
if os.path.exists(TMP):
os.remove(TMP)
db = DBManager(TMP)
sess = db.create_session("压缩持久化测试")
sid = sess["id"]
def add(role, content, parent):
return db.add_message(session_id=sid, role=role, content=content, parent_id=parent)
# 链:r1(user 旧) → r2(assistant 旧) → r3(user 旧) → r4(assistant 保留) → r5(user 保留)
r1 = add("user", "旧问题1", None)
r2 = add("assistant", "旧回答1", r1["id"])
r3 = add("user", "旧问题2", r2["id"])
r4 = add("assistant", "保留回答", r3["id"])
r5 = add("user", "保留问题", r4["id"])
leaf_before = db.get_session_leaf(sid)
check("T0.初始叶子=最后一条", leaf_before == r5["id"])
MARK_SUMMARY = "## 摘要\n上方对话已摘要:用户问了两个旧问题。"
# ---------- T1:标记插入(原子,不动叶子) ----------
mid = db.insert_compaction_mark(sid, MARK_SUMMARY, r3["id"], r4["id"],
json.dumps({"path": "pre_prompt"}, ensure_ascii=False))
chain = db.get_message_chain(sid)
check("T1.标记行已插入链上", mid is not None and any(m["id"] == mid for m in chain))
check("T1.链=6行且顺序正确",
[m["id"] for m in chain] == [r1["id"], r2["id"], r3["id"], mid, r4["id"], r5["id"]],
[m["id"] for m in chain])
check("T1.叶子指针不变", db.get_session_leaf(sid) == leaf_before)
check("T1.保留首条改指标记", chain[4]["parent_id"] == mid)
mark_row = next(m for m in chain if m["id"] == mid)
check("T1.标记 role/is_ignored", mark_row["role"] == "compaction" and mark_row["is_ignored"] == 1)
# ---------- T5:分支一致性(无幻影分支) ----------
check("T5.旧N唯一子=标记", [m["id"] for m in db.get_branch_info(r3["id"])] == [mid])
check("T5.标记唯一子=保留首条", [m["id"] for m in db.get_branch_info(mid)] == [r4["id"]])
# ---------- T2build_api_context 截断 ----------
class _Dummy:
db = db
api = MainWindow.build_api_context(_Dummy(), sid)
api_text = json.dumps([m.get("content") for m in api], ensure_ascii=False)
check("T2.首条=摘要user消息(带_kind)",
api and api[0]["role"] == "user" and api[0]["content"] == MARK_SUMMARY
and api[0].get("_kind") == "compaction_summary", api[:1])
check("T2.保留行仍在", "保留回答" in api_text and "保留问题" in api_text)
check("T2.切点前行消失",
"旧问题1" not in api_text and "旧回答1" not in api_text and "旧问题2" not in api_text)
check("T2.共3条(摘要+2保留)", len(api) == 3, len(api))
# ---------- T3kind 传播 → 二次压缩走迭代摘要 ----------
msgs = from_openai_messages(api)
check("T3.摘要 kind 还原", msgs and msgs[0].kind == "compaction_summary",
[m.kind for m in msgs])
prep = prepare_compaction(msgs, CompactionSettings(reserve_tokens=1000,
keep_recent_tokens=2))
check("T3.走 previous_summary 迭代路径",
prep is not None and prep.previous_summary == MARK_SUMMARY)
# ---------- T4:双标记取最后一个 ----------
mid2 = db.insert_compaction_mark(sid, "摘要2", r4["id"], r5["id"])
api2 = MainWindow.build_api_context(_Dummy(), sid)
api2_text = json.dumps([m.get("content") for m in api2], ensure_ascii=False)
check("T4.取最后一个标记", api2[0]["content"] == "摘要2" and len(api2) == 2, len(api2))
check("T4.只留第二个切点之后", "保留问题" in api2_text and "保留回答" not in api2_text)
# ---------- T6_cut_ids_of 切点计算(单元) ----------
A = AgentMessage(role="user", content="x", db_msg_id="a")
B = AgentMessage(role="assistant", content="y", db_msg_id="")
C = AgentMessage(role="toolResult", content="z", db_msg_id="")
S = AgentMessage(role="user", content="SUM", kind="compaction_summary")
check("T6.尾巴首条无DB行→first为空(UI跳过插标记)",
_cut_ids_of([A, B, C], [S, B, C]) == ("a", ""))
B2 = AgentMessage(role="assistant", content="y", db_msg_id="b")
check("T6.正常切点", _cut_ids_of([A, B2, C], [S, B2, C]) == ("a", "b"))
D = AgentMessage(role="assistant", content="d", db_msg_id="d")
check("T6.尾巴不在旧列表→不插", _cut_ids_of([A, B2, C], [S, D]) == ("", ""))
check("T6.无之前消息→无切点", _cut_ids_of([B2, C], [S, B2, C]) == ("", ""))
# T6c:⚠️ 同 id 回放序列(timeline 多条目共享一行 id)—— 曾导致 DB 自环卡死
Pm = AgentMessage(role="user", content="p", db_msg_id="p")
A1 = AgentMessage(role="assistant", content="a1", db_msg_id="a")
A2 = AgentMessage(role="assistant", content="a2", db_msg_id="a")
A3 = AgentMessage(role="toolResult", content="a3", db_msg_id="a")
Qm = AgentMessage(role="user", content="q", db_msg_id="q")
cut = _cut_ids_of([Pm, A1, A2, A3, Qm], [S, A2, A3, Qm])
check("T6c.同id回放切点回退到行头(cut_before≠first_retained)",
cut == ("p", "a"), cut)
cut2 = _cut_ids_of([Pm, A1, A2, A3, Qm], [S, A3, Qm])
check("T6c.尾巴起点更深也回退到行头", cut2 == ("p", "a"), cut2)
# ---------- T10insert_compaction_mark 自环守卫 ----------
bad = db.insert_compaction_mark(sid, "S", r5["id"], r5["id"])
check("T10.cut==retained 拒绝插入(防自环)", bad is None)
# ---------- T11get_message_chain 环守卫(手工造 2 环) ----------
sid2 = db.create_session("环守卫")["id"]
x1 = db.add_message(session_id=sid2, role="user", content="x1", parent_id=None)
x2 = db.add_message(session_id=sid2, role="assistant", content="x2", parent_id=x1["id"])
with db.get_connection() as conn:
conn.execute("UPDATE messages SET parent_id=? WHERE id=?", (x1["id"], x2["id"]))
conn.execute("UPDATE messages SET parent_id=? WHERE id=?", (x2["id"], x1["id"]))
conn.commit()
import time as _t
t0 = _t.time()
chain11 = db.get_message_chain(sid2)
dur = _t.time() - t0
check("T11.手工2环不死循环(<2s返回)", dur < 2 and 1 <= len(chain11) <= 3, (dur, len(chain11)))
# ---------- T7pre_prompt_compaction 事件 payload 端到端 ----------
model = ModelConfig(provider="t", name="t", context_window=500,
max_tokens=100, api_key="k", base_url="http://x")
cfg = AgentConfig(model=model, system_prompt="s", tools=[], tool_context={},
compaction_reserve=100, compaction_keep_recent=100)
from core.agent.agent import Agent # noqa: E402
agent = Agent(cfg)
old_msgs = [
AgentMessage(role="user", content="旧旧问题", db_msg_id="m1"),
AgentMessage(role="assistant", content="旧旧回答", db_msg_id="m2"),
AgentMessage(role="user", content="旧上下文 " * 200, db_msg_id="m3"),
AgentMessage(role="assistant", content="新回答", db_msg_id="m4"),
AgentMessage(role="user", content="新问题", db_msg_id="m5"),
]
agent.state.messages = old_msgs
runner = AgentRunner(agent, summarize_fn=lambda p, s, mt: "摘要X")
ok = runner.pre_prompt_compaction()
check("T7.压缩被触发", ok is True)
ev = runner.compaction_events[-1]
new_msgs = agent.state.messages
first = new_msgs[1]
idx = next(i for i, m in enumerate(old_msgs) if m is first)
check("T7.payload cut_before_id=切点前DB行",
ev.get("cut_before_id") == old_msgs[idx - 1].db_msg_id, ev)
check("T7.payload first_retained_id=尾巴首条DB行",
ev.get("first_retained_id") == first.db_msg_id, ev)
check("T7.新消息=摘要+尾巴", new_msgs[0].kind == "compaction_summary"
and new_msgs[1:] == old_msgs[idx:])
# ---------- T8_key 不泄漏到 API ----------
api_out = to_openai_messages(old_msgs)
check("T8.下划线内部字段不进API",
all(not any(k.startswith("_") for k in m.keys()) for m in api_out))
# ---------- T9UI 渲染过滤排除 compaction(静态) ----------
_mw_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
"ui", "views", "main_window.py")
with open(_mw_path, "r", encoding="utf-8") as f:
_src = f.read()
# 注意:探针与循环变量名解耦(源码列表推导式用 `m` 或 `msg` 均可),
# 只锁死“渲染过滤排除 system/compaction”这一语义。
check("T9.渲染过滤含 compaction", '["role"] not in ("system", "compaction")' in _src)
# ---------- T12Fix A — 摘要条目 timestamp → P0 失效过期锚点 ----------
# 场景(长程会话 92.7k→11k):保留行里 assistant 的入库 usage 是压缩前快照;
# 若摘要条目不带 timestamp,显示/压缩判定会锚到过期值(92.7k),而真实
# 下一轮输入只有摘要+保留行(~11k)。
sid3 = db.create_session("过期锚点失效")["id"]
u1 = db.add_message(session_id=sid3, role="user", content="u" * 2000, parent_id=None)
a1 = db.add_message(session_id=sid3, role="assistant", content="A" * 40000,
parent_id=u1["id"],
usage=json.dumps({"input": 92700, "output": 500,
"cacheRead": 0, "cacheWrite": 0}))
u2 = db.add_message(session_id=sid3, role="user", content="v" * 2000, parent_id=a1["id"])
db.add_message(session_id=sid3, role="assistant", content="B" * 40000, parent_id=u2["id"],
usage=json.dumps({"input": 92700, "output": 500,
"cacheRead": 0, "cacheWrite": 0}))
# 把全部行时间戳拨到 3000s 前(模拟它们早于压缩发生)
with db.get_connection() as conn:
conn.execute("UPDATE messages SET created_at=? WHERE session_id=?",
(int(_t.time()) - 3000, sid3))
conn.commit()
api_ns = MainWindow.build_api_context(_Dummy(), sid3)
est_ns = estimate_context_tokens(from_openai_messages(api_ns))
check("T12.对照:无标记→锚定入库usage", est_ns.usage_tokens > 0, est_ns.usage_tokens)
mid3 = db.insert_compaction_mark(sid3, "摘要", a1["id"], u2["id"])
api3 = MainWindow.build_api_context(_Dummy(), sid3)
check("T12.摘要条目带timestamp且晚于保留行",
api3 and api3[0].get("timestamp", 0) > 0
and all(api3[0]["timestamp"] > m.get("timestamp", 0) for m in api3[1:]))
est3 = estimate_context_tokens(from_openai_messages(api3))
check("T12.有标记→过期锚点失效(usage_tokens=0)", est3.usage_tokens == 0,
est3.usage_tokens)
check("T12.估算=全量公式(远离过期93.2k)", est3.tokens < 40000, est3.tokens)
# ---------- T13Fix B + G1 — 工具行显示不锚点(G1 行块规则单元化) ----------
# 场景(任务流会话 40.5k→80k):工具行的入库 usage 是本轮最后一个子请求的
# 快照(工具输出未截断),下一轮真实输入=4k 截断回放。🆕 G1 后不再需要
# 手动清锚:锚点选择器自动失效含工具行的 usage → 估算=全量公式。
sid4 = db.create_session("工具行不锚点")["id"]
p1 = db.add_message(session_id=sid4, role="user", content="prompt", parent_id=None)
tl = json.dumps([
{"t": "text", "text": "code" * 50000},
{"t": "tool", "id": "c1", "name": "bash", "args": "{}", "result": "ok"},
{"t": "text", "text": "done"},
], ensure_ascii=False)
db.add_message(session_id=sid4, role="assistant", content="x", parent_id=p1["id"],
timeline=tl,
usage=json.dumps({"input": 32517, "output": 7999,
"cacheRead": 0, "cacheWrite": 0}))
api4 = MainWindow.build_api_context(_Dummy(), sid4)
check("T13.回放含tool条目(规则条件成立)",
any(m.get("role") == "tool" for m in api4))
_est4 = estimate_context_tokens(from_openai_messages(api4))
check("T13.G1:工具行入库usage自动失效锚点(无需手动清锚)",
_est4.usage_tokens == 0, _est4.usage_tokens)
check("T13.估算=全量公式(>旧锚点40516,度量截断回放)",
_est4.tokens > 40516, _est4.tokens)
# G1 精确性:仅失效「含工具活动的行」的 usage,其后的纯文本行锚点保留
sid4b = db.create_session("工具行后纯文本行锚点保留")["id"]
p1b = db.add_message(session_id=sid4b, role="user", content="prompt", parent_id=None)
t1b = db.add_message(session_id=sid4b, role="assistant", content="x", parent_id=p1b["id"],
timeline=tl,
usage=json.dumps({"input": 32517, "output": 7999,
"cacheRead": 0, "cacheWrite": 0}))
u1b = db.add_message(session_id=sid4b, role="user", content="q2", parent_id=t1b["id"])
_a2b = db.add_message(session_id=sid4b, role="assistant", content="y", parent_id=u1b["id"],
usage=json.dumps({"input": 90000, "output": 100,
"cacheRead": 0, "cacheWrite": 0}))
api4b = MainWindow.build_api_context(_Dummy(), sid4b)
_est4b = estimate_context_tokens(from_openai_messages(api4b))
check("T13.G1:工具行(旧)usage失效、其后纯文本行(新)usage保留锚点",
_est4b.usage_tokens > 0 and _est4b.usage_tokens == 90100,
_est4b.usage_tokens)
# 纯文本会话:规则条件不成立 → 保留精确锚点(provider 实测值更准)
sid5 = db.create_session("文本行锚点")["id"]
p2 = db.add_message(session_id=sid5, role="user", content="hi", parent_id=None)
db.add_message(session_id=sid5, role="assistant", content="hello", parent_id=p2["id"],
usage=json.dumps({"input": 100, "output": 50,
"cacheRead": 0, "cacheWrite": 0}))
api5 = MainWindow.build_api_context(_Dummy(), sid5)
check("T13.纯文本行→无tool条目(保留锚点)",
not any(m.get("role") == "tool" for m in api5)
and estimate_context_tokens(from_openai_messages(api5)).usage_tokens > 0)
# ---------- T14G2 — should_compact 透传 system/tools(无锚点分支度量下一请求) ----------
from core.agent.context import should_compact
_mc14 = ModelConfig(name="t", context_window=100000)
_msgs14 = from_openai_messages(api4) # 工具行 → G1 无锚点
_sh_no, _tok_no = should_compact(_msgs14, _mc14, 16384)
_sh_yes, _tok_yes = should_compact(_msgs14, _mc14, 16384,
system_prompt="S" * 2000,
tools=[types.SimpleNamespace(
name="bash", description="d" * 1000,
parameters={})])
check("T14.传入system/tools后估算增大(无锚点分支)",
_tok_yes > _tok_no, (_tok_no, _tok_yes))
# 阈值边缘:同样内容,带上 system/tools 后跨过阈值 → 触发判定变化
_mc_edge = ModelConfig(name="t2", context_window=_tok_no + 16384 + 5)
_sh_edge_no, _ = should_compact(_msgs14, _mc_edge, 16384)
_sh_edge_yes, _ = should_compact(_msgs14, _mc_edge, 16384,
system_prompt="S" * 2000,
tools=[types.SimpleNamespace(
name="bash", description="d" * 1000,
parameters={})])
check("T14.阈值边缘:计入system/tools才触发(防低估漏压缩)",
_sh_edge_no is False and _sh_edge_yes is True,
(_sh_edge_no, _sh_edge_yes))
failed = [n for n, okk in RESULTS if not okk]
print(f"\n===== {len(RESULTS) - len(failed)}/{len(RESULTS)} PASS =====", flush=True)
print("ALL PASS" if not failed else f"FAILED: {failed}", flush=True)
sys.exit(0 if not failed else 1)