mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-25 05:56:18 +00:00
fix(llm): fence circuit probe settlement (#5602)
Co-authored-by: CorgiBoyG <CorgiBoyG@users.noreply.github.com>
This commit is contained in:
parent
479d2f10c8
commit
69286297fd
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
|
||||
@ -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"
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user