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>
This commit is contained in:
parent
6d2f8789d7
commit
061cf22c87
257
cheap-worker/tests/test_repair_middleware_spike.py
Normal file
257
cheap-worker/tests/test_repair_middleware_spike.py
Normal file
@ -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)
|
||||
Loading…
x
Reference in New Issue
Block a user