games-development-ai/cheap-worker/tests/test_repair_middleware_spike.py
lili 061cf22c87 feat(cheap-worker): 续修 middleware 机制 spike 坐实(拦 finish+注入+续跑)(切片一 阶段一前置)
on_reasoning 洋葱拦 finish Msg + 独立跑 fake check + 没过压制+注入 UserMsg 续跑,
4/4 确定性单测绿(stub 模型、零 LLM/chrome)。机制成立,阶段一接真 run_gates。

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-01 05:22:49 -07:00

258 lines
13 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.

"""
test_repair_middleware_spike.py — 续修 middleware 机制 spike(配置控制面 §3.4 前置去风险)。
坐实 AgentScope 2.0.2 的 MiddlewareBase.on_reasoning 能:
① 拦到 agent 想 finish(_reasoning_impl 在无 tool call 时 yield 纯文本 Msg,已装源 _agent.py:875);
② 独立跑一个外部 check()(此处 fake、零 LLM/chrome,只验机制,不绑真门);
③ 没过 → 压制 finish(不 re-yield)+ 注入「修」(UserMsg 纯文本,agent.observe)→ agent 自动续跑一轮;
④ 过 / 达 max_repairs → 放行 finish,reply 循环收尾。
机制稳(全部对已装 2.0.2 源码逐条核实,非凭 brief 转述):
· finish Msg 由 _reasoning_impl 产出(_agent.py:857 先把内容存进 context,875 再 yield AssistantMsg),
穿过 on_reasoning 洋葱链(_agent.py:745)才到 reply 循环(_agent.py:614-620),故 middleware 能在它到达前吞掉;
· 吞掉后 _reasoning() 这轮不产 Msg → reply 循环不走 614 的 return 分支 → 落到工具批次(空,_batch_tool_calls 见
末条是 user 消息返 [],_agent.py:1108)→ cur_iter++(_agent.py:670)→ while 未封顶再进推理;
· 续跑必进推理的根因:注入走 agent.observe(UserMsg),_handle_incoming_messages 追加一条 role=user 的消息
(_agent.py:1100),使 _get_last_msg 返 None(它只认本 agent 的 assistant 末条,_agent.py:2263),于是
_check_next_action 返 ("reasoning", None)(_agent.py:2304-2306/2361)—— 这条是「续跑不落空」的承重保证;
· max_iters 封顶(ReActConfig.max_iters,_agent.py:595)兜底防跑飞,RepairMiddleware 另有 max_repairs 自封顶。
坑(须处理·非阻断):finish 文本在 _agent.py:857 已先存进 context,压制 yield 抹不掉 —— 要干净血缘配 context.pop()。
绿了即成机制回归;阶段一把 RepairMiddleware 提升进 tier2/gen-worker/worker/middleware.py、fake check 换真 run_gates 独立重跑。
跑:cheap-worker/.venv/bin/python cheap-worker/tests/test_repair_middleware_spike.py
"""
import asyncio
import sys
from pathlib import Path
from typing import Any, AsyncGenerator, Callable, Type
sys.path.insert(0, str(Path(__file__).resolve().parents[1])) # → cheap-worker/
import _bootstrap # noqa: E402,F401 仅加 sys.path(import 时不取 key、不触网)
from agentscope.agent import Agent, ReActConfig # noqa: E402
from agentscope.middleware import MiddlewareBase # noqa: E402
from agentscope.model import ChatModelBase, ChatResponse # noqa: E402
from agentscope.message import Msg, TextBlock, UserMsg # noqa: E402
from agentscope.tool import Toolkit # noqa: E402
from agentscope.credential import CredentialBase # noqa: E402
from pydantic import BaseModel # noqa: E402
# ── stub 模型(确定性,不打真 LLM)──
# 直接采用已装 2.0.2 的 tests/utils.py MockCredential/MockModel verbatim(构造签名 + _call_api 调用路径),
# 与 shipped middleware_test.py 逐字一致以杜绝 API 漂移;仅省去本 spike 不触发的结构化输出辅助方法
# (_call_api_with_structured_output 非 abstractmethod、reply 路径不走,省略不影响机制)。
class MockCredential(CredentialBase):
"""The mock credential class."""
@classmethod
def get_chat_model_class(cls) -> Type["ChatModelBase"]:
"""Return the mock model class."""
return MockModel
class MockModel(ChatModelBase):
"""A mock model for testing."""
class Parameters(BaseModel):
"""The parameters."""
def __init__(
self,
model: str = "mock-model",
stream: bool = True,
context_size: int = 1000,
mock_chat_responses: list | None = None,
mock_structured_response: Any = None,
) -> None:
"""Initialize the mock model."""
super().__init__(
credential=MockCredential(),
model=model,
stream=stream,
parameters=MockModel.Parameters(),
context_size=context_size,
)
self.mock_chat_responses = mock_chat_responses or []
self.mock_structured_response = mock_structured_response
self.cnt = 0
def set_responses(
self,
mock_responses: list[ChatResponse | list[ChatResponse]],
) -> None:
"""Set the mock responses."""
self.mock_chat_responses = mock_responses
if all(isinstance(_, ChatResponse) for _ in mock_responses):
self.stream = False
else:
self.stream = True
self.cnt = 0
async def _call_api(
self, # pylint: disable=unused-argument
*args: Any,
**kwargs: Any,
) -> ChatResponse | AsyncGenerator[ChatResponse, None]:
"""Mock the API call. 每次调用取下一条(模拟多轮推理)。"""
mock_responses = self.mock_chat_responses[self.cnt]
self.cnt += 1
if isinstance(mock_responses, list):
async def _stream() -> AsyncGenerator[ChatResponse, None]:
for response in mock_responses:
yield response
return _stream()
if isinstance(mock_responses, ChatResponse):
return mock_responses
raise AssertionError
# ── 续修 middleware(spike 核心验证对象;阶段一提升进 worker/middleware.py)──
class RepairMiddleware(MiddlewareBase):
"""on_reasoning 拦 finish → 独立跑 check → 没过则压制 + 注入续跑。check 契约:async (agent)->(passed, feedback)。"""
def __init__(self, check: Callable, max_repairs: int = 3, *, pop_finish_claim: bool = False) -> None:
self._check = check
self._max_repairs = max_repairs
self._pop_finish_claim = pop_finish_claim # True=续修前把「做完了」从 context 弹掉(干净血缘)
self.repairs = 0 # 供断言/预算读
async def on_reasoning(self, agent, input_kwargs, next_handler) -> AsyncGenerator:
async for item in next_handler(**input_kwargs): # 与 Tier2TraceMiddleware.on_reasoning 同款透传入参
if not isinstance(item, Msg):
yield item # 非 finish 的事件流(ModelCallStart/Text* 等)原样透传
continue
# —— 拦到 finish(纯文本 Msg)——
passed, feedback = await self._check(agent)
if passed or self.repairs >= self._max_repairs:
yield item # 放行 finish → reply 循环收尾 return
return
self.repairs += 1
if self._pop_finish_claim and agent.state.context and isinstance(agent.state.context[-1], Msg):
agent.state.context.pop() # 可选:弹掉刚存进的「做完了」声称(_agent.py:857 先存)
# 注入「修」:role=user 纯文本(禁 system/tool/thinking,否则 _handle_incoming_messages 抛 ValueError,_agent.py:1085)
await agent.observe(UserMsg(name="gate", content=f"以下门未通过,请修复后再交付:\n{feedback}"))
return # 吞掉 finish(不 yield)→ 本轮无 Msg → reply 循环 cur_iter++ 续跑
def _text(resp_or_msg) -> str:
"""从 Msg/final 安全取文本(2.0.2 用 get_text_content(),返回 str|None,故 or "")。"""
try:
return resp_or_msg.get_text_content() or ""
except Exception: # noqa: BLE001 —— spike 断言辅助,取文本失败按空
return ""
# ───────────────────────── 核心:拦一次 + 注入 + 续跑放行第二版 ─────────────────────────
def test_intercept_inject_and_continue():
"""模型两次都想 finish;check 第一次 False、第二次 True → 门被跑两次(压制一次)、放行第二版、注入消息在 context。"""
calls = {"n": 0}
async def check(agent):
calls["n"] += 1
return (calls["n"] >= 2, "" if calls["n"] >= 2 else "H_progress 门未过(无进展)")
model = MockModel()
model.set_responses([
ChatResponse(content=[TextBlock(text="第一版游戏,做完了")], is_last=True), # iter0 想 finish
ChatResponse(content=[TextBlock(text="已修复进展问题,再交付")], is_last=True), # iter1 想 finish
])
mw = RepairMiddleware(check, max_repairs=3)
agent = Agent(name="cheap_worker", system_prompt="你是便宜档生成 agent",
model=model, toolkit=Toolkit(), middlewares=[mw],
react_config=ReActConfig(max_iters=5)) # 封顶防跑飞
final = asyncio.run(agent.reply(UserMsg(name="user", content="生成一个点击得分游戏")))
assert calls["n"] == 2, f"门应被独立跑两次(压制一次),实际 {calls['n']}"
assert mw.repairs == 1, f"应续修一次,实际 {mw.repairs}"
assert "已修复" in _text(final), f"放行的应是第二版,实际:{_text(final)!r}"
assert any(isinstance(m, Msg) and getattr(m, 'role', None) == "user"
and "H_progress" in _text(m) for m in agent.state.context), "注入的修复消息应在 context"
# ───────────────────────── 边界:达 max_repairs 即放行(预算封顶,不无限续)─────────────────────────
def test_max_repairs_cap_releases():
"""check 恒 False;max_repairs=2 → 续修 2 次后放行(不再压制),reply 收尾、不无限跑。"""
async def check(agent):
return (False, "门恒不过(测封顶)")
model = MockModel()
model.set_responses([ChatResponse(content=[TextBlock(text=f"第{i}版")], is_last=True) for i in range(6)])
mw = RepairMiddleware(check, max_repairs=2)
agent = Agent(name="cw", system_prompt="p", model=model, toolkit=Toolkit(),
middlewares=[mw], react_config=ReActConfig(max_iters=10))
final = asyncio.run(agent.reply(UserMsg(name="user", content="生成")))
assert mw.repairs == 2, f"应恰好续修 max_repairs=2 次,实际 {mw.repairs}"
assert _text(final), "达封顶后应放行一个 finish、非空"
# ───────────────────────── 回落触发点①:空推理(无 Msg 无 tool call)优雅续跑 ─────────────────────────
# Agent1 提示:若 reply 循环对「reasoning 返回空」有额外早退分支(源码未见),spike 真跑才暴露 → 那就回落 control_plane。
# 本用例即 test_intercept_inject_and_continue 的隐含验证(压制那轮 = reasoning 无 Msg 落地);
# 若上面核心用例绿,则 628-670「无 Msg 自动 cur_iter++ 续跑」这条已被真跑坐实,回落触发点①排除。
# ───────────────────────── 回落触发点②:observe 中途 append 不炸 context/formatter ─────────────────────────
def test_observe_midreasoning_no_context_error():
"""续修注入(observe)在半程 context 上 append 后,下一轮推理不因 compress/formatter 报错。"""
calls = {"n": 0}
async def check(agent):
calls["n"] += 1
return (calls["n"] >= 2, "补一个资源环再交付")
model = MockModel()
model.set_responses([
ChatResponse(content=[TextBlock(text="薄循环,做完了")], is_last=True),
ChatResponse(content=[TextBlock(text="加了进货资源环,交付")], is_last=True),
])
mw = RepairMiddleware(check, max_repairs=3)
agent = Agent(name="cw", system_prompt="p", model=model, toolkit=Toolkit(),
middlewares=[mw], react_config=ReActConfig(max_iters=5))
final = asyncio.run(agent.reply(UserMsg(name="user", content="生成经营游戏"))) # 不抛即通过
assert "资源环" in _text(final)
# ───────────────────────── 干净血缘变体:pop 掉 finish 声称 ─────────────────────────
def test_pop_finish_claim_variant():
"""pop_finish_claim=True → 续修前弹掉「做完了」声称;context 里不残留第一版的 finish 文本。"""
calls = {"n": 0}
async def check(agent):
calls["n"] += 1
return (calls["n"] >= 2, "修")
model = MockModel()
model.set_responses([
ChatResponse(content=[TextBlock(text="UNIQUE_CLAIM_一版做完了")], is_last=True),
ChatResponse(content=[TextBlock(text="二版交付")], is_last=True),
])
mw = RepairMiddleware(check, max_repairs=3, pop_finish_claim=True)
agent = Agent(name="cw", system_prompt="p", model=model, toolkit=Toolkit(),
middlewares=[mw], react_config=ReActConfig(max_iters=5))
asyncio.run(agent.reply(UserMsg(name="user", content="生成")))
assert not any("UNIQUE_CLAIM" in _text(m) for m in agent.state.context), "弹掉后 context 不应残留第一版 finish 声称"
if __name__ == "__main__":
_fns = [v for k, v in sorted(globals().items()) if k.startswith("test_") and callable(v)]
_failed = 0
for _fn in _fns:
try:
_fn()
print(f" PASS {_fn.__name__}")
except Exception as e: # noqa: BLE001
_failed += 1
import traceback
print(f" FAIL {_fn.__name__}: {type(e).__name__}: {e}")
traceback.print_exc()
print(f"\n{len(_fns) - _failed}/{len(_fns)} passed")
sys.exit(1 if _failed else 0)