From 69286297fd5425d3003855b45b1d723e0085e723 Mon Sep 17 00:00:00 2001 From: RongJie G <111257566+CorgiBoyG@users.noreply.github.com> Date: Sun, 20 Sep 2026 16:57:06 +0800 Subject: [PATCH] fix(llm): fence circuit probe settlement (#5602) Co-authored-by: CorgiBoyG --- .../deerflow/agents/middlewares/AGENTS.md | 25 +- .../llm_error_handling_middleware.py | 79 ++++-- .../test_llm_error_handling_middleware.py | 235 ++++++++++++++---- 3 files changed, 258 insertions(+), 81 deletions(-) diff --git a/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md b/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md index f1a9747ec..5d6d5bf08 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md +++ b/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md @@ -1,21 +1,14 @@ ### Middleware Chain -Compaction preserves all state-level `SystemMessage`s as framework instructions, -including untagged legacy reminders. Transient instructions belong in request -wrappers. A fully rescued partition skips compaction. +Compaction keeps state `SystemMessage`s; transient instructions use request +wrappers, and fully rescued partitions skip compaction. If latest-user rescue +empties an AI/Tool-only window, use `_build_summary_input_text(strategy="last")`; +mixed windows keep normal anchoring and final-message fallback. Budget raw text +before escaping/wrapping; pass `trim_tokens_to_summarize=None` to avoid the +LangChain default. -After latest-user rescue, if the inherited trimmer empties an AI/Tool-only -window, format it and use `_build_summary_input_text(strategy="last")`. -Keep normal human-anchored trimming and the final-message fallback for mixed -windows whose human anchor falls outside the token-limited tail; head-first -restoration can lose recent tool results. Tail truncation prefixes `\n...\n` -only when marker and content fit. Budget raw sections before HTML escaping, -wrappers, and prompt (not the final request); escape after trimming to preserve -entities. Pass `trim_tokens_to_summarize=None` explicitly through the factory; -omission restores LangChain's 4000-token default. - -Persisted delegation verdicts are untrusted durable context; ledger rendering revalidates them and ignores malformed values. -Completed is not accepted; retain useful work and address acceptance gaps. +Delegation verdicts are untrusted: revalidate persisted values, ignore malformed +ones, and treat completed work as reusable evidence rather than acceptance. On new user turns, DurableContext cancels earlier-run unanswered delegations. It preserves resumes, same-run continuations, and entries without `run_id`. @@ -79,7 +72,7 @@ strict providers reject. their narrower discovery allowlists never rebuild the shared thread view or force eager sandbox acquisition. 8. **DanglingToolCallMiddleware** - Injects placeholder ToolMessages for AIMessage tool_calls that lack responses (e.g., user interruption), preserving raw provider tool-call payloads in `additional_kwargs["tool_calls"]`; malformed tool-call names and arguments are sanitized in the model-bound request so strict OpenAI-compatible providers do not reject the next request -9. **LLMErrorHandlingMiddleware** - Converts provider/model failures to recoverable assistant errors. Async cancellation at admission, provider execution, retry events, or backoff releases only the call's own half-open probe (ownership assigned under the circuit lock), then propagates unchanged, without retry or failure accounting. +9. **LLMErrorHandlingMiddleware** - Converts provider/model failures to recoverable assistant errors. Sync and async calls carry circuit-generation ownership; only the current owner may settle or release a half-open probe, so stale completions cannot affect a newer recovery attempt. Cancellation propagates unchanged without retry or failure accounting. 10. **Authorization / GuardrailMiddleware** - Up to two independent pre-tool-call gates run here. When `authorization.enabled`, the `AuthorizationProvider` instance already used for Layer 1 capability filtering is wrapped by `GuardrailAuthorizationAdapter` and reused for Layer 2 execution checks. A generated `tool_search` bypasses the adapter's second provider call only when the current build has a concrete deferred setup; its catalog was already filtered by Layer 1, and an ordinary same-named tool without that deferred setup receives no exemption. When `guardrails.enabled`, the explicitly configured `GuardrailProvider` is appended after authorization and still evaluates every call, including `tool_search`. Authorization therefore runs outermost and can deny before an external guardrail call; both use the existing middleware's fail-closed, audit, sync/async, and error-`ToolMessage` behavior. See the authorization RFC and [docs/GUARDRAILS.md](../../../../../docs/GUARDRAILS.md). Every guardrail decision path publishes a neutral diff --git a/backend/packages/harness/deerflow/agents/middlewares/llm_error_handling_middleware.py b/backend/packages/harness/deerflow/agents/middlewares/llm_error_handling_middleware.py index 653a37409..401957b7e 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/llm_error_handling_middleware.py +++ b/backend/packages/harness/deerflow/agents/middlewares/llm_error_handling_middleware.py @@ -9,6 +9,7 @@ import threading import time from collections import deque from collections.abc import Awaitable, Callable +from dataclasses import dataclass from email.utils import parsedate_to_datetime from typing import Any, override @@ -34,6 +35,11 @@ _EMPTY_RESPONSE_RETRY_CONSUMED = object() _NON_CIRCUIT_FAILURE_REASONS = {"burst_rate", "empty_response"} +@dataclass(slots=True) +class _CircuitAdmission: + generation: int = -1 + + class EmptyModelResponseError(RuntimeError): """The model completed normally without producing persistent content.""" @@ -458,6 +464,7 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): self._circuit_failure_count = 0 self._circuit_open_until = 0.0 self._circuit_state = "closed" + self._circuit_generation = 0 self._circuit_probe_in_flight = False self._circuit_probe_token: object | None = None @@ -494,6 +501,7 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): if self._circuit_state == "open": if now < self._circuit_open_until: return True + self._circuit_generation += 1 self._circuit_state = "half_open" self._circuit_probe_in_flight = False self._circuit_probe_token = None @@ -503,24 +511,44 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): return True self._circuit_probe_in_flight = True self._circuit_probe_token = probe_token + if isinstance(probe_token, _CircuitAdmission): + probe_token.generation = self._circuit_generation return False + if isinstance(probe_token, _CircuitAdmission): + probe_token.generation = self._circuit_generation return False - def _record_success(self) -> None: + def _owns_current_circuit_generation(self, admission: _CircuitAdmission | None) -> bool: + if admission is None: + return True + if admission.generation != self._circuit_generation: + return False + if self._circuit_state == "half_open": + return self._circuit_probe_token is admission + return self._circuit_state == "closed" + + def _record_success(self, *, admission: _CircuitAdmission | None = None) -> None: with self._circuit_lock: + if not self._owns_current_circuit_generation(admission): + return if self._circuit_state != "closed" or self._circuit_failure_count > 0: logger.info("Circuit breaker reset (Closed). LLM service recovered.") + if self._circuit_state == "half_open": + self._circuit_generation += 1 self._circuit_failure_count = 0 self._circuit_open_until = 0.0 self._circuit_state = "closed" self._circuit_probe_in_flight = False self._circuit_probe_token = None - def _record_failure(self) -> None: + def _record_failure(self, *, admission: _CircuitAdmission | None = None) -> None: with self._circuit_lock: + if not self._owns_current_circuit_generation(admission): + return if self._circuit_state == "half_open": self._circuit_open_until = time.time() + self.circuit_recovery_timeout_sec + self._circuit_generation += 1 self._circuit_state = "open" self._circuit_probe_in_flight = False self._circuit_probe_token = None @@ -534,6 +562,7 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): if self._circuit_failure_count >= self.circuit_failure_threshold: self._circuit_open_until = time.time() + self.circuit_recovery_timeout_sec if self._circuit_state != "open": + self._circuit_generation += 1 self._circuit_state = "open" self._circuit_probe_in_flight = False self._circuit_probe_token = None @@ -543,18 +572,19 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): self.circuit_recovery_timeout_sec, ) - def _release_half_open_probe(self, *, probe_token: object | None = None) -> None: + def _release_half_open_probe(self, *, admission: _CircuitAdmission | None = None) -> None: """Release the in-flight half-open probe without recording a failure. Used when something other than a classified success/failure consumes the probe (a GraphBubbleUp control-flow signal, or a non-retriable error), so the circuit can admit - the next probe instead of fast-failing forever. Cancellation supplies an - admission token so an older call cannot release a different call's probe. + the next probe instead of fast-failing forever. The admission identity prevents an + older call from releasing a different call's probe. """ with self._circuit_lock: - if probe_token is not None and self._circuit_probe_token is not probe_token: + if not self._owns_current_circuit_generation(admission): return if self._circuit_state == "half_open": + self._circuit_generation += 1 self._circuit_probe_in_flight = False self._circuit_probe_token = None @@ -865,7 +895,8 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse], ) -> ModelCallResult: - if self._check_circuit(): + admission = _CircuitAdmission() + if self._check_circuit(probe_token=admission): return self._build_error_fallback_message( self._build_circuit_breaker_message(), error_type="CircuitBreakerOpen", @@ -879,11 +910,11 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): try: response = self._bounded_model_call_sync(request, handler) _raise_for_empty_response(response) - self._record_success() + self._record_success(admission=admission) return response except GraphBubbleUp: # Preserve LangGraph control-flow signals (interrupt/pause/resume). - self._release_half_open_probe() + self._release_half_open_probe(admission=admission) raise except Exception as exc: retriable, reason = self._classify_error(exc) @@ -901,7 +932,11 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): wait_ms, _extract_error_detail(exc), ) - self._emit_retry_event(attempt, wait_ms, reason, max_attempts=max_attempts) + try: + self._emit_retry_event(attempt, wait_ms, reason, max_attempts=max_attempts) + except GraphBubbleUp: + self._release_half_open_probe(admission=admission) + raise time.sleep(wait_ms / 1000) attempt += 1 continue @@ -912,10 +947,10 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): exc_info=exc, ) if retriable and reason not in _NON_CIRCUIT_FAILURE_REASONS: - self._record_failure() + self._record_failure(admission=admission) else: # These outcomes do not show that the provider is broadly unavailable. - self._release_half_open_probe() + self._release_half_open_probe(admission=admission) return self._build_user_fallback_message(exc, reason) @override @@ -924,8 +959,8 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): request: ModelRequest, handler: Callable[[ModelRequest], Awaitable[ModelResponse]], ) -> ModelCallResult: - probe_token = object() - if self._check_circuit(probe_token=probe_token): + admission = _CircuitAdmission() + if self._check_circuit(probe_token=admission): return self._build_error_fallback_message( self._build_circuit_breaker_message(), error_type="CircuitBreakerOpen", @@ -940,11 +975,11 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): try: response = await self._bounded_model_call(request, handler) _raise_for_empty_response(response) - self._record_success() + self._record_success(admission=admission) return response except GraphBubbleUp: # Preserve LangGraph control-flow signals (interrupt/pause/resume). - self._release_half_open_probe() + self._release_half_open_probe(admission=admission) raise except Exception as exc: retriable, reason = self._classify_error(exc) @@ -962,7 +997,11 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): wait_ms, _extract_error_detail(exc), ) - await self._aemit_retry_event(attempt, wait_ms, reason, max_attempts=max_attempts) + try: + await self._aemit_retry_event(attempt, wait_ms, reason, max_attempts=max_attempts) + except GraphBubbleUp: + self._release_half_open_probe(admission=admission) + raise await asyncio.sleep(wait_ms / 1000) attempt += 1 continue @@ -973,15 +1012,15 @@ class LLMErrorHandlingMiddleware(AgentMiddleware[AgentState]): exc_info=exc, ) if retriable and reason not in _NON_CIRCUIT_FAILURE_REASONS: - self._record_failure() + self._record_failure(admission=admission) else: # These outcomes do not show that the provider is broadly unavailable. - self._release_half_open_probe() + self._release_half_open_probe(admission=admission) return self._build_user_fallback_message(exc, reason) except asyncio.CancelledError: # Cancellation can arrive during admission, the provider call, retry # event delivery, or backoff. It is not a provider failure. - self._release_half_open_probe(probe_token=probe_token) + self._release_half_open_probe(admission=admission) raise diff --git a/backend/tests/test_llm_error_handling_middleware.py b/backend/tests/test_llm_error_handling_middleware.py index 42f2587d4..5c0e3178f 100644 --- a/backend/tests/test_llm_error_handling_middleware.py +++ b/backend/tests/test_llm_error_handling_middleware.py @@ -200,6 +200,151 @@ def test_async_model_call_marks_transient_retry_exhaustion_as_error_fallback( assert result.additional_kwargs["error_detail"] == "Connection error." +@pytest.mark.parametrize("stale_outcome", ["success", "failure", "non_retriable", "graph_bubble_up"]) +def test_async_stale_completion_cannot_settle_an_owned_half_open_probe(stale_outcome: str) -> None: + async def run() -> None: + middleware = _build_middleware(retry_max_attempts=1) + middleware.circuit_failure_threshold = 1 + middleware.circuit_recovery_timeout_sec = 0 + old_entered = asyncio.Event() + release_old = asyncio.Event() + probe_entered = asyncio.Event() + release_probe = asyncio.Event() + provider_calls: list[str] = [] + + async def handler(request) -> AIMessage: + name = request.name + provider_calls.append(name) + if name == "old": + old_entered.set() + await release_old.wait() + if stale_outcome == "failure": + raise FakeError("old request failed", status_code=503) + if stale_outcome == "non_retriable": + raise FakeError("insufficient_quota", status_code=429, code="insufficient_quota") + if stale_outcome == "graph_bubble_up": + raise GraphBubbleUp() + elif name == "open": + raise FakeError("provider unavailable", status_code=503) + elif name == "probe": + probe_entered.set() + await release_probe.wait() + return AIMessage(content=f"{name} success") + + old_task = asyncio.create_task(middleware.awrap_model_call(SimpleNamespace(name="old"), handler)) + probe_task: asyncio.Task | None = None + try: + await asyncio.wait_for(old_entered.wait(), timeout=1) + await middleware.awrap_model_call(SimpleNamespace(name="open"), handler) + assert middleware._circuit_state == "open" + + probe_task = asyncio.create_task(middleware.awrap_model_call(SimpleNamespace(name="probe"), handler)) + await asyncio.wait_for(probe_entered.wait(), timeout=1) + assert middleware._circuit_state == "half_open" + assert middleware._circuit_probe_in_flight is True + + release_old.set() + if stale_outcome == "graph_bubble_up": + with pytest.raises(GraphBubbleUp): + await asyncio.wait_for(old_task, timeout=1) + else: + await asyncio.wait_for(old_task, timeout=1) + + assert middleware._circuit_state == "half_open" + assert middleware._circuit_probe_in_flight is True + calls_before_extra = list(provider_calls) + extra = await middleware.awrap_model_call(SimpleNamespace(name="extra"), handler) + assert provider_calls == calls_before_extra + assert extra.additional_kwargs["error_type"] == "CircuitBreakerOpen" + finally: + release_old.set() + release_probe.set() + if not old_task.done(): + await old_task + if probe_task is not None: + await probe_task + + asyncio.run(run()) + + +@pytest.mark.parametrize("stale_outcome", ["success", "failure", "non_retriable", "graph_bubble_up"]) +def test_sync_stale_completion_cannot_settle_an_owned_half_open_probe(stale_outcome: str) -> None: + middleware = _build_middleware(retry_max_attempts=1) + middleware.circuit_failure_threshold = 1 + middleware.circuit_recovery_timeout_sec = 0 + old_entered = threading.Event() + release_old = threading.Event() + probe_entered = threading.Event() + release_probe = threading.Event() + provider_calls: list[str] = [] + provider_lock = threading.Lock() + failures: list[BaseException] = [] + + def handler(request) -> AIMessage: + name = request.name + with provider_lock: + provider_calls.append(name) + if name == "old": + old_entered.set() + assert release_old.wait(1) + if stale_outcome == "failure": + raise FakeError("old request failed", status_code=503) + if stale_outcome == "non_retriable": + raise FakeError("insufficient_quota", status_code=429, code="insufficient_quota") + if stale_outcome == "graph_bubble_up": + raise GraphBubbleUp() + elif name == "open": + raise FakeError("provider unavailable", status_code=503) + elif name == "probe": + probe_entered.set() + assert release_probe.wait(1) + return AIMessage(content=f"{name} success") + + def call(name: str) -> None: + try: + middleware.wrap_model_call(SimpleNamespace(name=name), handler) + except BaseException as exc: + failures.append(exc) + + old_thread = threading.Thread(target=call, args=("old",), daemon=True) + probe_thread = threading.Thread(target=call, args=("probe",), daemon=True) + try: + old_thread.start() + assert old_entered.wait(1) + middleware.wrap_model_call(SimpleNamespace(name="open"), handler) + assert middleware._circuit_state == "open" + + probe_thread.start() + assert probe_entered.wait(1) + assert middleware._circuit_state == "half_open" + assert middleware._circuit_probe_in_flight is True + + release_old.set() + old_thread.join(1) + assert not old_thread.is_alive() + if stale_outcome == "graph_bubble_up": + assert len(failures) == 1 + assert isinstance(failures[0], GraphBubbleUp) + else: + assert not failures + + assert middleware._circuit_state == "half_open" + assert middleware._circuit_probe_in_flight is True + with provider_lock: + calls_before_extra = list(provider_calls) + extra = middleware.wrap_model_call(SimpleNamespace(name="extra"), handler) + with provider_lock: + assert provider_calls == calls_before_extra + assert extra.additional_kwargs["error_type"] == "CircuitBreakerOpen" + finally: + release_old.set() + release_probe.set() + old_thread.join(1) + probe_thread.join(1) + if stale_outcome != "graph_bubble_up": + assert not failures + + def test_sync_model_call_uses_retry_after_header(monkeypatch: pytest.MonkeyPatch) -> None: middleware = _build_middleware(retry_max_attempts=2, retry_base_delay_ms=10, retry_cap_delay_ms=10) waits: list[float] = [] @@ -519,10 +664,7 @@ def test_empty_response_exhaustion_does_not_trip_circuit_breaker(monkeypatch: py def test_empty_response_exhaustion_releases_half_open_probe(monkeypatch: pytest.MonkeyPatch) -> None: middleware = _build_middleware(retry_max_attempts=2, retry_base_delay_ms=1, retry_cap_delay_ms=1) middleware._circuit_state = "half_open" - assert middleware._check_circuit() is False - assert middleware._circuit_probe_in_flight is True monkeypatch.setattr("time.sleep", lambda _delay: None) - monkeypatch.setattr(middleware, "_check_circuit", lambda: False) def empty_handler(_request) -> AIMessage: return AIMessage(content="", response_metadata={"finish_reason": "stop"}) @@ -693,6 +835,45 @@ def test_async_model_call_propagates_graph_bubble_up() -> None: asyncio.run(middleware.awrap_model_call(SimpleNamespace(), handler)) +def test_sync_retry_event_graph_bubble_up_releases_half_open_probe(monkeypatch: pytest.MonkeyPatch) -> None: + middleware = _build_middleware(retry_max_attempts=2) + middleware._circuit_state = "half_open" + + def unavailable(_request) -> AIMessage: + raise FakeError("Service unavailable", status_code=503) + + def interrupt_retry_event(*_args, **_kwargs) -> None: + raise GraphBubbleUp() + + monkeypatch.setattr(middleware, "_emit_retry_event", interrupt_retry_event) + + with pytest.raises(GraphBubbleUp): + middleware.wrap_model_call(SimpleNamespace(), unavailable) + + assert middleware._circuit_state == "half_open" + assert middleware._circuit_probe_in_flight is False + + +@pytest.mark.anyio +async def test_async_retry_event_graph_bubble_up_releases_half_open_probe(monkeypatch: pytest.MonkeyPatch) -> None: + middleware = _build_middleware(retry_max_attempts=2) + middleware._circuit_state = "half_open" + + async def unavailable(_request) -> AIMessage: + raise FakeError("Service unavailable", status_code=503) + + async def interrupt_retry_event(*_args, **_kwargs) -> None: + raise GraphBubbleUp() + + monkeypatch.setattr(middleware, "_aemit_retry_event", interrupt_retry_event) + + with pytest.raises(GraphBubbleUp): + await middleware.awrap_model_call(SimpleNamespace(), unavailable) + + assert middleware._circuit_state == "half_open" + assert middleware._circuit_probe_in_flight is False + + @pytest.mark.anyio @pytest.mark.parametrize("cancel_during", ["provider", "backoff", "queue"]) async def test_cancelled_recovery_probe_allows_next_model_call(cancel_during: str, monkeypatch: pytest.MonkeyPatch) -> None: @@ -811,25 +992,15 @@ def test_circuit_half_open_graph_bubble_up_resets_probe() -> None: """Verify that GraphBubbleUp in half_open state resets probe_in_flight.""" middleware = _build_middleware() - # Step 1: Manually set state to half_open and check_circuit() to set probe_in_flight=True middleware._circuit_state = "half_open" middleware._circuit_probe_in_flight = False - # Call _check_circuit() once to simulate the probe being allowed through - assert middleware._check_circuit() is False - assert middleware._circuit_probe_in_flight is True - # Step 2: Now trigger handler that raises GraphBubbleUp def handler(_request) -> AIMessage: raise GraphBubbleUp() - # Mock _check_circuit() to return False (since we already did the probe check) - import unittest.mock + with pytest.raises(GraphBubbleUp): + middleware.wrap_model_call(SimpleNamespace(), handler) - with unittest.mock.patch.object(middleware, "_check_circuit", return_value=False): - with pytest.raises(GraphBubbleUp): - middleware.wrap_model_call(SimpleNamespace(), handler) - - # Verify probe_in_flight was reset, state should remain half_open assert middleware._circuit_probe_in_flight is False assert middleware._circuit_state == "half_open" @@ -839,25 +1010,15 @@ async def test_async_circuit_half_open_graph_bubble_up_resets_probe() -> None: """Verify that GraphBubbleUp in half_open state resets probe_in_flight (async version).""" middleware = _build_middleware() - # Step 1: Manually set state to half_open and check_circuit() to set probe_in_flight=True middleware._circuit_state = "half_open" middleware._circuit_probe_in_flight = False - # Call _check_circuit() once to simulate the probe being allowed through - assert middleware._check_circuit() is False - assert middleware._circuit_probe_in_flight is True - # Step 2: Now trigger handler that raises GraphBubbleUp async def handler(_request) -> AIMessage: raise GraphBubbleUp() - # Mock _check_circuit() to return False (since we already did the probe check) - import unittest.mock + with pytest.raises(GraphBubbleUp): + await middleware.awrap_model_call(SimpleNamespace(), handler) - with unittest.mock.patch.object(middleware, "_check_circuit", return_value=False): - with pytest.raises(GraphBubbleUp): - await middleware.awrap_model_call(SimpleNamespace(), handler) - - # Verify probe_in_flight was reset, state should remain half_open assert middleware._circuit_probe_in_flight is False assert middleware._circuit_state == "half_open" @@ -872,25 +1033,15 @@ def test_circuit_half_open_non_retriable_error_resets_probe() -> None: fast-failed forever because no later call could ever run the handler to reach ``_record_success`` / ``_record_failure``. """ - import unittest.mock - middleware = _build_middleware() - # Enter half_open and let one probe through (probe_in_flight -> True). middleware._circuit_state = "half_open" middleware._circuit_probe_in_flight = False - assert middleware._check_circuit() is False - assert middleware._circuit_probe_in_flight is True def handler(_request) -> AIMessage: raise FakeError("insufficient_quota", status_code=429, code="insufficient_quota") - # _check_circuit already admitted the probe above; keep it False here so the - # top-of-call gate does not fast-fail before the handler runs. Force the - # error to classify as non-retriable regardless of heuristics. - with unittest.mock.patch.object(middleware, "_check_circuit", return_value=False): - with unittest.mock.patch.object(middleware, "_classify_error", return_value=(False, "quota")): - result = middleware.wrap_model_call(SimpleNamespace(), handler) + result = middleware.wrap_model_call(SimpleNamespace(), handler) # Non-retriable errors still surface a graceful fallback (not a raise) and # must NOT trip the breaker. @@ -906,21 +1057,15 @@ def test_circuit_half_open_non_retriable_error_resets_probe() -> None: @pytest.mark.anyio async def test_async_circuit_half_open_non_retriable_error_resets_probe() -> None: """Async mirror: a non-retriable error during a half-open probe releases it.""" - import unittest.mock - middleware = _build_middleware() middleware._circuit_state = "half_open" middleware._circuit_probe_in_flight = False - assert middleware._check_circuit() is False - assert middleware._circuit_probe_in_flight is True async def handler(_request) -> AIMessage: raise FakeError("insufficient_quota", status_code=429, code="insufficient_quota") - with unittest.mock.patch.object(middleware, "_check_circuit", return_value=False): - with unittest.mock.patch.object(middleware, "_classify_error", return_value=(False, "quota")): - result = await middleware.awrap_model_call(SimpleNamespace(), handler) + result = await middleware.awrap_model_call(SimpleNamespace(), handler) assert isinstance(result, AIMessage) assert middleware._circuit_state == "half_open"