From 061cf22c878cdb38b8c18da156fa1442e07a11d0 Mon Sep 17 00:00:00 2001 From: lili Date: Wed, 1 Jul 2026 05:22:49 -0700 Subject: [PATCH] =?UTF-8?q?feat(cheap-worker):=20=E7=BB=AD=E4=BF=AE=20midd?= =?UTF-8?q?leware=20=E6=9C=BA=E5=88=B6=20spike=20=E5=9D=90=E5=AE=9E(?= =?UTF-8?q?=E6=8B=A6=20finish+=E6=B3=A8=E5=85=A5+=E7=BB=AD=E8=B7=91)(?= =?UTF-8?q?=E5=88=87=E7=89=87=E4=B8=80=20=E9=98=B6=E6=AE=B5=E4=B8=80?= =?UTF-8?q?=E5=89=8D=E7=BD=AE)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit on_reasoning 洋葱拦 finish Msg + 独立跑 fake check + 没过压制+注入 UserMsg 续跑, 4/4 确定性单测绿(stub 模型、零 LLM/chrome)。机制成立,阶段一接真 run_gates。 Co-Authored-By: Claude Opus 4.8 (1M context) --- .../tests/test_repair_middleware_spike.py | 257 ++++++++++++++++++ 1 file changed, 257 insertions(+) create mode 100644 cheap-worker/tests/test_repair_middleware_spike.py diff --git a/cheap-worker/tests/test_repair_middleware_spike.py b/cheap-worker/tests/test_repair_middleware_spike.py new file mode 100644 index 00000000..c884cdf8 --- /dev/null +++ b/cheap-worker/tests/test_repair_middleware_spike.py @@ -0,0 +1,257 @@ +""" +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)