Files
Video-Tree-TRM5/tests/unit/test_agent_loop.py
T

392 lines
14 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.
"""core/agent/loop.py 单元测试。
算法保真 #11 — AgentLoop 推理循环引擎。
9 个测试覆盖: 终止、预算、无效工具、解析错误、JSON 修复、
thinking 捕获、token 累加、call_id 透传、pluggy hook。
"""
from __future__ import annotations
import json
import ssl
from typing import TYPE_CHECKING, Any
from unittest.mock import AsyncMock
import pytest
from core.agent.loop import AgentLoop
from core.agent.protocols import hookimpl
from core.types import LLMResponse
if TYPE_CHECKING:
from core.agent.types import LoopResult, Step
# ── 测试基础设施 ──────────────────────────────────────────────
class _StubDispatcher:
"""测试用工具调度器。"""
async def dispatch(
self, tool_name: str, args: dict[str, Any], *, context: dict[str, Any]
) -> str:
if tool_name == "submit_answer":
return "答案已提交"
if tool_name == "search_tree":
return "搜索结果: 找到节点 L2-3"
raise ValueError(f"未知工具: {tool_name}")
def _make_response(content: str, thinking: str = "") -> LLMResponse:
"""构造测试用 LLMResponse。"""
return LLMResponse(
content=content,
thinking=thinking,
model="test-model",
provider="test",
prompt_tokens=10,
completion_tokens=10,
latency_ms=100,
ttft_ms=50.0,
max_inter_token_ms=10.0,
cache_hit=False,
call_id="test-call-id",
)
def _submit_json(answer: str = "42") -> str:
"""构造 submit_answer 的 JSON 响应。"""
return json.dumps(
{
"reflect": {"observation": "找到答案"},
"plan": {"next_step": "提交"},
"action": {"tool": "submit_answer", "args": {"answer": answer}},
}
)
def _search_json() -> str:
"""构造 search_tree 的 JSON 响应。"""
return json.dumps(
{
"reflect": {"observation": "需要搜索"},
"plan": {"next_step": "搜索"},
"action": {"tool": "search_tree", "args": {"query": "test"}},
}
)
def _invalid_tool_json() -> str:
"""构造无效工具的 JSON 响应。"""
return json.dumps(
{
"reflect": {},
"plan": {},
"action": {"tool": "unknown_tool", "args": {}},
}
)
# ── 测试用例 ──────────────────────────────────────────────────
class TestAgentLoop:
"""AgentLoop 推理循环引擎测试。"""
@pytest.mark.asyncio
async def test_submit_answer_terminates_loop(self) -> None:
"""submit_answer 终止循环 → finished, result=args, steps_used=1。"""
llm = AsyncMock()
llm.chat.return_value = _make_response(_submit_json())
loop = AgentLoop(llm=llm, max_steps=10)
result = await loop.run("system", "user", _StubDispatcher())
assert result.stop_reason == "finished"
assert result.result == {"answer": "42"}
assert result.steps_used == 1
assert len(result.steps) == 1
@pytest.mark.asyncio
async def test_budget_exceeded(self) -> None:
"""max_steps=3 用完 → budget_exceeded, steps_used=3。"""
llm = AsyncMock()
llm.chat.return_value = _make_response(_search_json())
loop = AgentLoop(llm=llm, max_steps=3)
result = await loop.run("system", "user", _StubDispatcher())
assert result.stop_reason == "budget_exceeded"
assert result.steps_used == 3
@pytest.mark.asyncio
async def test_invalid_tool_not_counted_as_step(self) -> None:
"""无效工具(ValueError)不计步 → steps_used=1。"""
llm = AsyncMock()
llm.chat.side_effect = [
_make_response(_invalid_tool_json()),
_make_response(_submit_json()),
]
loop = AgentLoop(llm=llm, max_steps=10)
result = await loop.run("system", "user", _StubDispatcher())
assert result.stop_reason == "finished"
assert result.steps_used == 1
@pytest.mark.asyncio
async def test_parse_error_after_max_retries(self) -> None:
"""非 JSON 内容连续失败 → parse_error, steps_used=0。"""
llm = AsyncMock()
llm.chat.return_value = _make_response("这不是JSON内容")
loop = AgentLoop(llm=llm, max_steps=10, max_retries=3)
result = await loop.run("system", "user", _StubDispatcher())
assert result.stop_reason == "parse_error"
assert result.steps_used == 0
@pytest.mark.asyncio
async def test_json_repair_handles_malformed(self) -> None:
"""轻微 JSON 缺陷(缺少闭合花括号)被 json_repair 修复。"""
malformed = (
'{"reflect": {}, "plan": {}, '
'"action": {"tool": "submit_answer", "args": {"answer": "42"}}'
)
llm = AsyncMock()
llm.chat.return_value = _make_response(malformed)
loop = AgentLoop(llm=llm, max_steps=10)
result = await loop.run("system", "user", _StubDispatcher())
assert result.stop_reason == "finished"
assert result.result == {"answer": "42"}
@pytest.mark.asyncio
async def test_thinking_content_captured_in_step(self) -> None:
"""LLMResponse.thinking → Step.thought。"""
llm = AsyncMock()
llm.chat.return_value = _make_response(_submit_json(), thinking="深度思考过程")
loop = AgentLoop(llm=llm, max_steps=10)
result = await loop.run("system", "user", _StubDispatcher())
assert result.steps[0].thought == "深度思考过程"
@pytest.mark.asyncio
async def test_token_usage_accumulated(self) -> None:
"""多步 token 累加: 3 次调用 × 10 tokens = 30。"""
llm = AsyncMock()
llm.chat.side_effect = [
_make_response(_search_json()),
_make_response(_search_json()),
_make_response(_submit_json()),
]
loop = AgentLoop(llm=llm, max_steps=10)
result = await loop.run("system", "user", _StubDispatcher())
assert result.token_usage["prompt_tokens"] == 30
assert result.token_usage["completion_tokens"] == 30
@pytest.mark.asyncio
async def test_call_id_propagated_to_step(self) -> None:
"""LLMResponse.call_id → Step.call_id。"""
llm = AsyncMock()
llm.chat.return_value = _make_response(_submit_json())
loop = AgentLoop(llm=llm, max_steps=10)
result = await loop.run("system", "user", _StubDispatcher())
assert result.steps[0].call_id == "test-call-id"
@pytest.mark.asyncio
async def test_pluggy_hooks_called(self) -> None:
"""TrackingPlugin 验证 before_step/after_tool/after_step/on_finish 全部触发。"""
class TrackingPlugin:
"""记录 hook 调用事件的测试插件。"""
def __init__(self) -> None:
self.events: list[str] = []
@hookimpl
async def before_step(self, iteration: int, messages: list[dict[str, Any]]) -> None:
self.events.append(f"before_step:{iteration}")
@hookimpl
async def after_tool(self, iteration: int, step: Step) -> str | None:
self.events.append(f"after_tool:{iteration}")
return None
@hookimpl
async def after_step(self, iteration: int, messages: list[dict[str, Any]]) -> None:
self.events.append(f"after_step:{iteration}")
@hookimpl
async def on_finish(self, result: LoopResult) -> None:
self.events.append(f"on_finish:{result.stop_reason}")
tracker = TrackingPlugin()
llm = AsyncMock()
llm.chat.return_value = _make_response(_submit_json())
loop = AgentLoop(llm=llm, max_steps=10)
result = await loop.run("system", "user", _StubDispatcher(), plugins=[tracker])
assert result.stop_reason == "finished"
assert "before_step:0" in tracker.events
assert "after_tool:0" in tracker.events
assert "after_step:0" in tracker.events
assert "on_finish:finished" in tracker.events
# ── A1 解析容错测试(Spec-1)──────────────────────────────────
# 生产真实样本结构:尾部围栏残留 + action.args 平铺(开头围栏场景由
# test_leading_json_fence 单独覆盖)
_REAL_FLAT_FENCED = """{
"plan": {
"goal": "从三个L1根节点开始建立全局认知",
"tool": "view_node",
"reason": "三个L1节点覆盖整个视频"
},
"action": {
"tool": "view_node",
"node_id": "J5Npf2xJpag_L1_000",
"question": "What is the overall topic of this video?"
}
}
```"""
class TestParseNormalization:
"""deepseek 输出变体(args 平铺 + ```json 围栏)归一化。"""
def _parse(self, content: str) -> tuple[str, dict, dict, str, dict, str] | None:
"""构造 AgentLoop 并解析给定 content,返回 _parse_response 结果。"""
loop = AgentLoop(llm=AsyncMock(), max_steps=10)
return loop._parse_response(_make_response(content))
def test_flat_args_with_trailing_fence(self) -> None:
"""生产样本:action 平铺 node_id/question + 尾部围栏。"""
parsed = self._parse(_REAL_FLAT_FENCED)
assert parsed is not None
action = parsed[4]
assert action["tool"] == "view_node"
assert action["args"] == {
"node_id": "J5Npf2xJpag_L1_000",
"question": "What is the overall topic of this video?",
}
def test_leading_json_fence(self) -> None:
"""开头 ```json 围栏 + 尾部围栏包裹的标准结构可正常解析。"""
content = '```json\n{"reflect": {}, "plan": {}, "action": {"tool": "submit_answer", "args": {"answer": "A"}}}\n```'
parsed = self._parse(content)
assert parsed is not None
assert parsed[4]["args"] == {"answer": "A"}
def test_nested_args_unchanged(self) -> None:
"""标准嵌套结构不受归一化影响。"""
parsed = self._parse(_submit_json("B"))
assert parsed is not None
assert parsed[4] == {"tool": "submit_answer", "args": {"answer": "B"}}
def test_action_missing_tool_still_rejected(self) -> None:
"""action 缺 tool 键的结构仍被拒绝。"""
content = json.dumps({"reflect": {}, "plan": {}, "action": {"node_id": "x"}})
assert self._parse(content) is None
def test_argless_action_still_rejected(self) -> None:
"""有 tool、无 args、无平铺参数键 → 不得收拢为空 args,必须拒绝。"""
content = json.dumps({"reflect": {}, "plan": {}, "action": {"tool": "submit_answer"}})
assert self._parse(content) is None
def test_empty_content_still_rejected(self) -> None:
"""空 content 照旧拒绝,归一化不改变该边界。"""
assert self._parse("") is None
# ── A2 步级重试测试(Spec-1)──────────────────────────────────
class TestStepLevelRetry:
"""LLM 瞬时异常的步级重试:可重试元组 / 退避 / fail-fast。"""
def _make_loop(self, chat_side_effects: list[object]) -> AgentLoop:
"""构造 chat 按序抛异常/返回响应的 AgentLoop。"""
llm = AsyncMock()
llm.chat = AsyncMock(side_effect=chat_side_effects)
return AgentLoop(llm=llm, max_steps=10)
@pytest.mark.asyncio
async def test_transient_error_retried_then_succeeds(
self, monkeypatch: pytest.MonkeyPatch
) -> None:
"""SSL/超时瞬时异常按 20s/40s 退避重试,第三次成功 → finished。"""
delays: list[float] = []
async def _fake_sleep(seconds: float) -> None:
delays.append(seconds)
monkeypatch.setattr("core.agent.loop.asyncio.sleep", _fake_sleep)
loop = self._make_loop(
[
ssl.SSLError("SSLV3_ALERT_BAD_RECORD_MAC"),
TimeoutError("watchdog"),
_make_response(_submit_json()),
]
)
result = await loop.run("sys", "user", _StubDispatcher())
assert result.stop_reason == "finished"
assert delays == [20.0, 40.0]
@pytest.mark.asyncio
async def test_retry_exhausted_terminates_with_error(
self, monkeypatch: pytest.MonkeyPatch
) -> None:
"""重试预算(首次 + 2 次)耗尽仍失败 → stop_reason=error。"""
async def _fake_sleep(seconds: float) -> None:
pass
monkeypatch.setattr("core.agent.loop.asyncio.sleep", _fake_sleep)
loop = self._make_loop([TimeoutError("t1"), TimeoutError("t2"), TimeoutError("t3")])
result = await loop.run("sys", "user", _StubDispatcher())
assert result.stop_reason == "error"
assert loop._llm.chat.await_count == 3 # 首次 + 2 次重试
@pytest.mark.asyncio
async def test_non_retryable_fails_fast(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""非可重试异常不退避不重发,首次即终止 → stop_reason=error。"""
sleep_mock = AsyncMock()
monkeypatch.setattr("core.agent.loop.asyncio.sleep", sleep_mock)
loop = self._make_loop([RuntimeError("programming bug")])
result = await loop.run("sys", "user", _StubDispatcher())
assert result.stop_reason == "error"
sleep_mock.assert_not_awaited()
assert loop._llm.chat.await_count == 1
@pytest.mark.asyncio
async def test_exhaustion_raises_last_original_exception(
self, monkeypatch: pytest.MonkeyPatch
) -> None:
"""耗尽语义窄测试:_call_llm_with_step_retry 耗尽后原样抛出最后一次异常。"""
async def _fake_sleep(seconds: float) -> None:
pass
monkeypatch.setattr("core.agent.loop.asyncio.sleep", _fake_sleep)
last_error = TimeoutError("t3-last")
loop = self._make_loop([TimeoutError("t1"), TimeoutError("t2"), last_error])
token_usage = {"prompt_tokens": 0, "completion_tokens": 0}
with pytest.raises(TimeoutError) as exc_info:
await loop._call_llm_with_step_retry([], token_usage, session_id=None)
assert exc_info.value is last_error
def test_empty_retry_delays_rejected_at_init(self) -> None:
"""空 step_retry_delays 在构造时 fail-fast 抛 ValueError。"""
with pytest.raises(ValueError, match="step_retry_delays"):
AgentLoop(llm=AsyncMock(), max_steps=10, step_retry_delays=())