Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
65 changes: 63 additions & 2 deletions src/Mod/VibeCAD/VibeCADProvider.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
import base64
from contextlib import ExitStack, contextmanager
from dataclasses import dataclass
from email.utils import mktime_tz, parsedate_tz
import hashlib
import json
import multiprocessing
Expand Down Expand Up @@ -4732,6 +4733,43 @@ def _is_retryable_anthropic_stream_error(
return any(token in text for token in retry_tokens)


def _is_retryable_anthropic_request_error(
exc: BaseException, anthropic_module: Any
) -> bool:
status_error = getattr(anthropic_module, "APIStatusError", ())
if isinstance(exc, status_error):
directive = exc.response.headers.get("x-should-retry")
if directive in {"true", "false"}:
return directive == "true"
return exc.status_code in {408, 409, 429} or exc.status_code >= 500
return _is_retryable_anthropic_stream_error(exc, anthropic_module)


def _anthropic_request_retry_delay(exc: BaseException, attempt: int) -> float:
response = getattr(exc, "response", None)
if response is None:
return min(2.0, 0.25 * attempt)

headers = response.headers
delay = None
for name, scale in (("retry-after-ms", 0.001), ("retry-after", 1.0)):
try:
delay = float(headers[name]) * scale
break
except (KeyError, TypeError, ValueError):
continue
if delay is None and headers.get("retry-after"):
try:
parsed = parsedate_tz(headers["retry-after"])
if parsed is not None:
delay = mktime_tz(parsed) - time.time()
except (TypeError, ValueError, OverflowError):
pass
if delay is not None and 0 < delay <= 60:
return delay
return min(8.0, 0.5 * 2 ** min(attempt - 1, 4))


def _bounded_compaction_text(value: Any, limit: int) -> str:
text = str(value or "").strip()
if len(text) <= limit:
Expand Down Expand Up @@ -5889,6 +5927,10 @@ def build_tool_surface(
)
)

# Stream retries belong to the outer loop, including failures after
# headers. Keep metadata and separate compaction SDK retries intact.
client.max_retries = 0

request_kwargs: dict[str, Any] = {
"model": model,
"max_tokens": max_tokens,
Expand All @@ -5902,7 +5944,11 @@ def build_tool_surface(
"effort": _anthropic_adaptive_effort(reasoning_effort)
}

response_content_observed = False

def _stream_response(turn: int, attempt: int) -> Any:
nonlocal response_content_observed
response_content_observed = False
# The SDK rejects non-streaming requests that could exceed ten
# minutes (large max_tokens plus thinking budgets), so always
# stream and accumulate the final message.
Expand Down Expand Up @@ -5971,6 +6017,8 @@ def _stream_response(turn: int, attempt: int) -> Any:
event_count += 1
summary = _anthropic_stream_event_summary(stream_event)
stream_event_type = summary.get("stream_event_type")
if stream_event_type in {"content_block_start", "content_block_delta"}:
response_content_observed = True
delta_type = summary.get("delta_type")
text_delta = summary.get("text_delta")
if text_delta:
Expand Down Expand Up @@ -6046,15 +6094,22 @@ def _stream_response(turn: int, attempt: int) -> Any:
return stream.get_final_message()

def _stream_response_with_retries(turn: int) -> Any:
status_failures = 0
for attempt in range(1, ANTHROPIC_STREAM_MAX_ATTEMPTS + 1):
try:
return _stream_response(turn, attempt)
except anthropic.BadRequestError:
raise
except Exception as exc:
is_status_error = isinstance(
exc, getattr(anthropic, "APIStatusError", ())
)
if is_status_error:
status_failures += 1
if (
attempt >= ANTHROPIC_STREAM_MAX_ATTEMPTS
or not _is_retryable_anthropic_stream_error(exc, anthropic)
or status_failures >= client_kwargs["max_retries"] + 1
or not _is_retryable_anthropic_request_error(exc, anthropic)
):
raise
_send_child_progress(
Expand All @@ -6064,11 +6119,17 @@ def _stream_response_with_retries(turn: int) -> Any:
"turn": turn,
"attempt": attempt,
"next_attempt": attempt + 1,
"transport_attempt_count": attempt,
"response_content_observed": response_content_observed,
"max_attempts": ANTHROPIC_STREAM_MAX_ATTEMPTS,
"error": _short_provider_error(exc),
},
)
time.sleep(min(2.0, 0.25 * attempt))
time.sleep(
_anthropic_request_retry_delay(
exc, status_failures if is_status_error else attempt
)
)
raise RuntimeError("Anthropic stream retry loop exited unexpectedly.")

turn = 1
Expand Down
195 changes: 195 additions & 0 deletions src/Mod/VibeCAD/vibecad_tests/test_anthropic_retry_budget.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,195 @@
# SPDX-License-Identifier: LGPL-2.1-or-later

"""Count real SDK transport attempts without contacting Anthropic."""

from __future__ import annotations

import json

import pytest

import VibeCADProvider as provider

anthropic = pytest.importorskip("anthropic")
httpx = pytest.importorskip("httpx")


def _event(kind, **fields):
return f"event: {kind}\ndata: {json.dumps({'type': kind, **fields})}\n\n".encode()


def _response_events():
return [
_event("message_start", message={
"id": "msg_test", "type": "message", "role": "assistant",
"model": "test-model", "content": [], "stop_reason": None,
"stop_sequence": None, "usage": {"input_tokens": 1, "output_tokens": 0},
}),
_event("content_block_start", index=0,
content_block={"type": "text", "text": ""}),
_event("content_block_delta", index=0,
delta={"type": "text_delta", "text": "Recovered."}),
_event("content_block_stop", index=0),
_event("message_delta", delta={"stop_reason": "end_turn", "stop_sequence": None},
usage={"output_tokens": 1}),
_event("message_stop"),
]


@pytest.fixture
def run_transport(monkeypatch):
real_client = anthropic.Anthropic
clients = []

def run(outcomes, *, metadata_failures=0):
requests = []
metadata = []
sleeps = []
messages = []

class InterruptedStream(httpx.SyncByteStream):
def __iter__(self):
yield b"".join(_response_events()[:3])
raise httpx.ReadError("peer closed connection")

def handle(request):
if request.method == "GET":
metadata.append(request)
if len(metadata) <= metadata_failures:
return httpx.Response(503, json={
"type": "error", "error": {
"type": "overloaded_error", "message": "Temporary"
},
})
return httpx.Response(200, json={
"id": "test-model", "type": "model",
"display_name": "Test", "created_at": "2026-01-01T00:00:00Z",
"max_tokens": 8192,
})
requests.append(request)
outcome = outcomes[min(len(requests) - 1, len(outcomes) - 1)]
if outcome == "timeout":
raise httpx.ReadTimeout("read timed out", request=request)
if outcome == "connect":
raise httpx.ConnectError("connection failed", request=request)
if outcome == "broken":
return httpx.Response(200, headers={"content-type": "text/event-stream"},
stream=InterruptedStream())
if outcome == "ok":
return httpx.Response(200, headers={"content-type": "text/event-stream"},
content=b"".join(_response_events()))
status, headers = outcome if isinstance(outcome, tuple) else (outcome, {})
return httpx.Response(status, headers=headers, json={
"type": "error", "error": {"type": "api_error", "message": "Temporary"},
})

def make_client(**kwargs):
client = real_client(
**kwargs, http_client=httpx.Client(transport=httpx.MockTransport(handle))
)
clients.append(client)
return client

class Connection:
def send(self, message):
messages.append(message)

def close(self):
pass

monkeypatch.setattr(anthropic, "Anthropic", make_client)
monkeypatch.setattr(provider, "_validate_provider_wire_surface", lambda _: None)
monkeypatch.setattr(provider.time, "sleep", sleeps.append)
provider._anthropic_child_main(
Connection(), "Inspect the model.", {"provider_tool_schemas": []},
"test-model", "fake-key", None, 1.0, 1, False,
)
return requests, metadata, messages, sleeps

yield run
for client in clients:
client.close()


@pytest.mark.parametrize("failure", ["timeout", "connect", "broken", 408, 409, 429, 503])
def test_stream_total_transport_attempts_are_bounded(run_transport, failure):
requests, _, messages, _ = run_transport([failure])
expected_attempts = 3 if isinstance(failure, int) else 6
assert len(requests) == expected_attempts
assert messages[-1]["type"] == "error"
retries = [
message["event"] for message in messages
if message.get("event", {}).get("event") == "anthropic_stream_retrying"
]
assert [event["next_attempt"] for event in retries] == list(
range(2, expected_attempts + 1)
)
assert [event["transport_attempt_count"] for event in retries] == list(
range(1, expected_attempts)
)
assert all(event["response_content_observed"] == (failure == "broken")
for event in retries)
assert not any(message["type"] == "tool" for message in messages)


def test_stream_mixed_failures_can_recover_on_last_attempt(run_transport):
requests, _, messages, _ = run_transport(
["timeout", 503, "broken", 429, "connect", "ok"]
)
assert len(requests) == 6
assert messages[-1] == {
"type": "done", "final_output": "Recovered.", "raw": None
}


@pytest.mark.parametrize("status", [400, 401, 403, 422])
def test_stream_permanent_errors_are_not_retried(run_transport, status):
requests, _, messages, sleeps = run_transport([status])
assert len(requests) == 1
assert messages[-1]["type"] == "error"
assert sleeps == []


def test_stream_obeys_explicit_do_not_retry(run_transport):
requests, _, messages, sleeps = run_transport([(503, {"x-should-retry": "false"})])
assert len(requests) == 1
assert messages[-1]["type"] == "error"
assert sleeps == []


@pytest.mark.parametrize("headers,delay", [
({"retry-after": "2"}, 2.0),
({"retry-after-ms": "1250", "retry-after": "2"}, 1.25),
({"retry-after": "Thu, 01 Jan 1970 00:16:43 GMT"}, 3.0),
])
def test_stream_respects_server_retry_delay(run_transport, monkeypatch, headers, delay):
monkeypatch.setattr(provider.time, "time", lambda: 1000.0)
requests, _, messages, sleeps = run_transport([(429, headers), "ok"])
assert len(requests) == 2
assert messages[-1]["type"] == "done"
assert sleeps == [delay]


def test_metadata_retains_sdk_retries(run_transport):
requests, metadata, messages, _ = run_transport(["ok"], metadata_failures=2)
assert len(metadata) == 3
assert len(requests) == 1
assert messages[-1]["type"] == "done"


@pytest.mark.parametrize("value", ["invalid", "NaN", "Infinity", "-1", "61"])
def test_stream_invalid_retry_delay_uses_bounded_backoff(run_transport, value):
requests, _, messages, sleeps = run_transport(
[(429, {"retry-after": value}), "ok"]
)
assert len(requests) == 2
assert messages[-1]["type"] == "done"
assert sleeps == [0.5]


def test_stream_status_retry_limit_survives_mixed_failures(run_transport):
requests, _, messages, _ = run_transport(
[503, "timeout", 429, "broken", 503, "ok"]
)
assert len(requests) == 5
assert messages[-1]["type"] == "error"
49 changes: 49 additions & 0 deletions src/Mod/VibeCAD/vibecad_tests/test_provider_subprocess.py
Original file line number Diff line number Diff line change
Expand Up @@ -676,6 +676,7 @@ def compact(**kwargs):
assert len(compaction_calls) == 1
assert compaction_calls[0]["model"] == "memory-model"
assert compaction_calls[0]["max_tokens"] == 128_000
assert compaction_calls[0]["client_kwargs"]["max_retries"] == 2
assert model_requests == ["interactive-model", "memory-model"]
assert messages.requests[0]["max_tokens"] == 128_000
recovery_request = messages.requests[1]
Expand Down Expand Up @@ -2691,3 +2692,51 @@ def test_partdesign_does_not_inject_a_model_manifest_at_turn_start(
assert "partdesign" not in context
assert "vibescript" not in context
assert context["editable_sources"]["domain"] == "partdesign"


@pytest.mark.parametrize("stop", ["cancel", "deadline"])
def test_retry_wait_is_terminated_by_parent_stop(monkeypatch, stop) -> None:
now = [0.0]

class RetryWaitProcess(_ExitedProcess):
terminated = False

def is_alive(self):
return not self.terminated

def terminate(self):
self.terminated = True

class RetryWaitPipe(_DelayedPipeMessage):
notified = False

def poll(self, timeout):
assert not self.notified, "Parent continued waiting after stop"
return True

def recv(self):
self.notified = True
now[0] = 2.0
return {
"type": "progress",
"event": {"event": "anthropic_stream_retrying", "attempt": 1},
}

context = _FakeMultiprocessingContext()
context.process = RetryWaitProcess()
context.parent_conn = RetryWaitPipe()
monkeypatch.setattr(
provider, "_provider_multiprocessing_context", lambda **kwargs: context
)
monkeypatch.setattr(provider.time, "monotonic", lambda: now[0])
error = provider.ProviderUnavailable if stop == "cancel" else TimeoutError
with pytest.raises(error):
provider._run_provider_subprocess(
prompt="Inspect", context={}, tool_runner=None, model="test-model",
api_key=None, reasoning_effort=None,
timeout_seconds=1.0 if stop == "deadline" else None,
cancellation_check=lambda: stop == "cancel" and context.parent_conn.notified,
child_main=_unused_child, clear_inherited_modules=False,
)
assert context.process.terminated
assert context.parent_conn.closed
Loading