diff --git a/backend/packages/harness/deerflow/agents/middlewares/loop_detection_middleware.py b/backend/packages/harness/deerflow/agents/middlewares/loop_detection_middleware.py index d9c1b0b15..9e1fc8571 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/loop_detection_middleware.py +++ b/backend/packages/harness/deerflow/agents/middlewares/loop_detection_middleware.py @@ -847,8 +847,26 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): self._pending_warning_touch_order.pop(pending_key, None) return warnings - def _augment_request(self, request: ModelRequest) -> ModelRequest: - """Append queued loop warnings (if any) to the outgoing message list. + def _restore_pending_warnings(self, runtime: Runtime, warnings: list[str]) -> None: + """Requeue warnings taken for a model call that raised. + + LLMErrorHandlingMiddleware sits outside this middleware and retries a + failed call by running this wrap again, so the retry must still find + the warning. It would not be queued again: it is already marked warned. + """ + if not warnings: + return + pending_key = self._pending_key(runtime) + with self._lock: + queued = self._pending_warnings[pending_key] + queued[:0] = [warning for warning in warnings if warning not in queued] + # Keep the restored warnings at the front; trim what came after them. + del queued[_MAX_PENDING_WARNINGS_PER_RUN:] + self._touch_pending_warning_key_locked(pending_key) + self._prune_pending_warning_state_locked(protected_key=pending_key) + + def _inject_warnings(self, request: ModelRequest, warnings: list[str]) -> ModelRequest: + """Append *warnings* to the outgoing message list. The warning is placed *after* every existing message, including the ToolMessage responses to the previous AIMessage(tool_calls). This @@ -857,7 +875,6 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): restriction (we use HumanMessage), and never mutates an existing AIMessage. """ - warnings = self._drain_pending_warnings(request.runtime) if not warnings: return request new_messages = [ @@ -872,7 +889,12 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse], ) -> ModelCallResult: - return handler(self._augment_request(request)) + warnings = self._drain_pending_warnings(request.runtime) + try: + return handler(self._inject_warnings(request, warnings)) + except Exception: + self._restore_pending_warnings(request.runtime, warnings) + raise @override async def awrap_model_call( @@ -880,7 +902,12 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): request: ModelRequest, handler: Callable[[ModelRequest], Awaitable[ModelResponse]], ) -> ModelCallResult: - return await handler(self._augment_request(request)) + warnings = self._drain_pending_warnings(request.runtime) + try: + return await handler(self._inject_warnings(request, warnings)) + except Exception: + self._restore_pending_warnings(request.runtime, warnings) + raise def reset(self, thread_id: str | None = None) -> None: """Clear tracking state. If thread_id given, clear only that thread.""" diff --git a/backend/packages/harness/deerflow/agents/middlewares/token_budget_middleware.py b/backend/packages/harness/deerflow/agents/middlewares/token_budget_middleware.py index 9a5154f83..06d5f46fd 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/token_budget_middleware.py +++ b/backend/packages/harness/deerflow/agents/middlewares/token_budget_middleware.py @@ -344,6 +344,20 @@ class TokenBudgetMiddleware(AgentMiddleware[AgentState]): warnings = self._pending_warnings.pop(run_id, None) return warnings or [] + def _restore_pending_warnings(self, runtime: Runtime, warnings: list[str]) -> None: + """Requeue warnings taken for a model call that raised. + + LLMErrorHandlingMiddleware sits outside this middleware and retries a + failed call by running this wrap again, so the retry must still find + the warning. It is not queued twice: ``_warned`` is already set. + """ + if not warnings: + return + run_id = self._get_run_id(runtime) + with self._lock: + queued = self._pending_warnings.setdefault(run_id, []) + queued[:0] = [warning for warning in warnings if warning not in queued] + def _inject_warnings(self, request: ModelRequest, warnings: list[str]) -> ModelRequest: if not warnings: return request @@ -357,14 +371,18 @@ class TokenBudgetMiddleware(AgentMiddleware[AgentState]): @override def wrap_model_call(self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse]) -> ModelCallResult: - warnings = self._drain_pending_warnings(request.runtime) - request = self._inject_warnings(request, warnings) - - return handler(request) + try: + return handler(self._inject_warnings(request, warnings)) + except Exception: + self._restore_pending_warnings(request.runtime, warnings) + raise @override async def awrap_model_call(self, request: ModelRequest, handler: Callable[[ModelRequest], Awaitable[ModelResponse]]) -> ModelCallResult: warnings = self._drain_pending_warnings(request.runtime) - request = self._inject_warnings(request, warnings) - return await handler(request) + try: + return await handler(self._inject_warnings(request, warnings)) + except Exception: + self._restore_pending_warnings(request.runtime, warnings) + raise diff --git a/backend/packages/harness/deerflow/agents/middlewares/tool_progress_middleware.py b/backend/packages/harness/deerflow/agents/middlewares/tool_progress_middleware.py index 7a0c13e36..6851a8440 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/tool_progress_middleware.py +++ b/backend/packages/harness/deerflow/agents/middlewares/tool_progress_middleware.py @@ -575,6 +575,23 @@ class ToolProgressMiddleware(AgentMiddleware[AgentState]): with self._lock: return self._pending.pop(key, []) + def _restore_pending(self, runtime: Runtime, hints: list[str]) -> None: + """Requeue hints taken for a model call that raised. + + LLMErrorHandlingMiddleware sits outside this middleware and retries a + failed call by running this wrap again, so the retry must still find + the hints. + """ + if not hints: + return + key = self._pending_key(runtime) + with self._lock: + if key[0] not in self._phase_states: + return + queue = self._pending[key] + queue[:0] = hints + del queue[_MAX_PENDING_PER_RUN:] + def _clear_stale_pending(self, runtime: Runtime) -> None: thread_id, current_run = self._pending_key(runtime) with self._lock: @@ -700,8 +717,7 @@ class ToolProgressMiddleware(AgentMiddleware[AgentState]): # ------------------------------------------------------------------ # wrap_model_call: drain pending hints and inject before model sees messages - def _augment_request(self, request: ModelRequest) -> ModelRequest: - hints = self._drain_pending(request.runtime) + def _inject_hints(self, request: ModelRequest, hints: list[str]) -> ModelRequest: if not hints: return request deduped = list(dict.fromkeys(hints)) @@ -722,7 +738,12 @@ class ToolProgressMiddleware(AgentMiddleware[AgentState]): request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse], ) -> ModelCallResult: - return handler(self._augment_request(request)) + hints = self._drain_pending(request.runtime) + try: + return handler(self._inject_hints(request, hints)) + except Exception: + self._restore_pending(request.runtime, hints) + raise @override async def awrap_model_call( @@ -730,7 +751,12 @@ class ToolProgressMiddleware(AgentMiddleware[AgentState]): request: ModelRequest, handler: Callable[[ModelRequest], Awaitable[ModelResponse]], ) -> ModelCallResult: - return await handler(self._augment_request(request)) + hints = self._drain_pending(request.runtime) + try: + return await handler(self._inject_hints(request, hints)) + except Exception: + self._restore_pending(request.runtime, hints) + raise # ------------------------------------------------------------------ # before_agent: clean up stale pending hints from previous runs diff --git a/backend/tests/test_loop_detection_middleware.py b/backend/tests/test_loop_detection_middleware.py index 827f26b8d..16d2a1ce9 100644 --- a/backend/tests/test_loop_detection_middleware.py +++ b/backend/tests/test_loop_detection_middleware.py @@ -1211,6 +1211,64 @@ class TestLoopDetectionAgentGraphIntegration: _scope_key("cached-thread", "run-2"), } + def test_loop_warning_survives_a_retried_model_call_in_real_agent_graph(self): + """LLMErrorHandlingMiddleware retries a failed call by running the inner wraps again; the retry must still carry the warning.""" + from deerflow.agents.middlewares.llm_error_handling_middleware import LLMErrorHandlingMiddleware + from deerflow.config.app_config import AppConfig, LlmCallConfig + from deerflow.config.sandbox_config import SandboxConfig + + class ProviderUnavailable(Exception): + def __init__(self) -> None: + super().__init__("503 Service Unavailable") + self.status_code = 503 + self.response = SimpleNamespace(status_code=503, headers={}) + + class FailsOnceOnWarning(_CapturingFakeMessagesListChatModel): + _failed: bool = PrivateAttr(default=False) + + def _generate(self, messages, stop=None, run_manager=None, **kwargs): + if not self._failed and any(isinstance(message, HumanMessage) and message.name == "loop_warning" for message in messages): + self._failed = True + self._seen_messages.append(list(messages)) + raise ProviderUnavailable() + return super()._generate(messages, stop=stop, run_manager=run_manager, **kwargs) + + @as_tool + def bash(command: str) -> str: + """Run a fake shell command.""" + return f"ran: {command}" + + repeated_calls = [[{"name": "bash", "id": f"call_ls_{i}", "args": {"command": "ls"}}] for i in range(3)] + model = FailsOnceOnWarning( + responses=[ + AIMessage(content="", tool_calls=repeated_calls[0]), + AIMessage(content="", tool_calls=repeated_calls[1]), + AIMessage(content="", tool_calls=repeated_calls[2]), + AIMessage(content="final answer"), + ], + ) + app_config = AppConfig( + sandbox=SandboxConfig(use="test"), + llm_call=LlmCallConfig(retry_max_attempts=3, retry_base_delay_ms=0, retry_cap_delay_ms=0), + ) + graph = create_agent( + model=model, + tools=[bash], + middleware=[LLMErrorHandlingMiddleware(app_config=app_config), LoopDetectionMiddleware(warn_threshold=3, hard_limit=10)], + ) + + result = graph.invoke( + {"messages": [("user", "inspect the directory")]}, + context={"thread_id": "retry-thread", "run_id": "retry-run"}, + config={"recursion_limit": 20}, + ) + + # Three tool-calling requests, then the failed attempt and its retry. + assert len(model.seen_messages) == 5 + has_warning = [any(isinstance(message, HumanMessage) and message.name == "loop_warning" for message in messages) for messages in model.seen_messages] + assert has_warning == [False, False, False, True, True] + assert result["messages"][-1].content == "final answer" + def test_loop_warning_is_transient_in_real_agent_graph(self): """after_model queues the warning; wrap_model_call injects it request-only.""" diff --git a/backend/tests/test_token_budget_middleware.py b/backend/tests/test_token_budget_middleware.py index f80ae708b..4791575b6 100644 --- a/backend/tests/test_token_budget_middleware.py +++ b/backend/tests/test_token_budget_middleware.py @@ -186,6 +186,28 @@ class TestTokenBudgetWarning: assert sent[2].name == "budget_warning" assert "TOKEN BUDGET WARNING" in sent[2].content + def test_warning_survives_a_failed_model_call(self): + """A call that raises is retried by LLMErrorHandlingMiddleware through this wrap; the warning must still be sent.""" + config = TokenBudgetConfig(max_tokens=100000, warn_threshold=0.8, enabled=True) + mw = TokenBudgetMiddleware.from_config(config) + runtime = _make_runtime() + mw._apply(_make_state_with_usage(total=85000), runtime) + + request = _make_request([AIMessage(content="hi")], runtime) + sent = [] + + def flaky_handler(req): + sent.append(req.messages) + if len(sent) == 1: + raise RuntimeError("503 Service Unavailable") + return MagicMock() + + with pytest.raises(RuntimeError): + mw.wrap_model_call(request, flaky_handler) + mw.wrap_model_call(request, flaky_handler) + + assert [any(getattr(message, "name", None) == "budget_warning" for message in messages) for messages in sent] == [True, True] + def test_warn_only_once_per_run(self): config = TokenBudgetConfig(max_tokens=100000, warn_threshold=0.8, enabled=True) mw = TokenBudgetMiddleware.from_config(config) diff --git a/backend/tests/test_tool_progress_middleware.py b/backend/tests/test_tool_progress_middleware.py index a56859426..68464a242 100644 --- a/backend/tests/test_tool_progress_middleware.py +++ b/backend/tests/test_tool_progress_middleware.py @@ -965,8 +965,8 @@ def test_no_hint_when_inject_assessment_disabled(): assert hints == [] -def test_augment_request_deduplicates_identical_hints(): - """L2: _augment_request must deduplicate identical hint strings via dict.fromkeys. +def test_inject_hints_deduplicates_identical_hints(): + """L2: _inject_hints must deduplicate identical hint strings via dict.fromkeys. If the same hint text appears multiple times in the queue (e.g. two successive no_results errors produce identical hint strings), only one copy should be @@ -1345,6 +1345,56 @@ async def test_awrap_tool_call_malformed_meta_passthrough(): assert mw._phase_states.get("t1", {}).get("web_search") is None +def test_hint_survives_a_failed_sync_model_call(): + mw = _make_mw(stagnation_threshold=2, warn_escalation_count=5) + rt = _make_runtime() + req = _make_tool_request(runtime=rt) + error_msg = _make_error_message() + mw.wrap_tool_call(req, lambda r: error_msg) + mw.wrap_tool_call(req, lambda r: error_msg) + + model_req = _make_model_request([], rt) + sent: list = [] + + def flaky_handler(r): + sent.append(r.messages) + if len(sent) == 1: + raise RuntimeError("503 Service Unavailable") + return MagicMock() + + with pytest.raises(RuntimeError): + mw.wrap_model_call(model_req, flaky_handler) + mw.wrap_model_call(model_req, flaky_handler) + + assert [any(isinstance(m, HumanMessage) and "PROGRESS HINT" in m.content for m in messages) for messages in sent] == [True, True] + + +@pytest.mark.anyio +async def test_hint_survives_a_failed_model_call(): + """A call that raises is retried by LLMErrorHandlingMiddleware through this wrap; the hint must still be sent.""" + mw = _make_mw(stagnation_threshold=2, warn_escalation_count=5) + rt = _make_runtime() + req = _make_tool_request(runtime=rt) + error_msg = _make_error_message() + mw.wrap_tool_call(req, lambda r: error_msg) + mw.wrap_tool_call(req, lambda r: error_msg) + + model_req = _make_model_request([], rt) + sent: list = [] + + async def flaky_handler(r): + sent.append(r.messages) + if len(sent) == 1: + raise RuntimeError("503 Service Unavailable") + return MagicMock() + + with pytest.raises(RuntimeError): + await mw.awrap_model_call(model_req, flaky_handler) + await mw.awrap_model_call(model_req, flaky_handler) + + assert [any(isinstance(m, HumanMessage) and "PROGRESS HINT" in m.content for m in messages) for messages in sent] == [True, True] + + @pytest.mark.anyio async def test_awrap_model_call_drains_and_injects_hints(): mw = _make_mw(stagnation_threshold=2, warn_escalation_count=5)