from __future__ import annotations import pytest from gcmd_classifier.config import ModelSettings from gcmd_classifier.errors import ( ModelRetriesExhaustedError, NonRetryableModelError, RetryableModelError, ) from gcmd_classifier.llm import ( FakeModelClient, ModelRequest, ModelStage, RetryPolicy, generate_with_retries, ) from gcmd_classifier.llm.schemas import TopicResponse def _request() -> ModelRequest[TopicResponse]: return ModelRequest.from_settings( stage=ModelStage.TOPIC, prompt="prompt", response_schema=TopicResponse, settings=ModelSettings(max_retries=2), ) def test_retryable_error_succeeds_after_retry() -> None: client = FakeModelClient( [ RetryableModelError("temporary"), {"selected": [], "ambiguous_alternatives": [], "no_selection_reason": "None."}, ] ) response = generate_with_retries(client, _request(), RetryPolicy(max_retries=1)) assert response.retry_count == 1 assert len(client.requests) == 2 def test_retryable_error_fails_after_max_retries() -> None: client = FakeModelClient([RetryableModelError("one"), RetryableModelError("two")]) with pytest.raises(ModelRetriesExhaustedError) as exc_info: generate_with_retries(client, _request(), RetryPolicy(max_retries=1)) assert exc_info.value.retry_count == 1 assert len(client.requests) == 2 def test_non_retryable_error_is_not_retried() -> None: client = FakeModelClient([NonRetryableModelError("bad request")]) with pytest.raises(NonRetryableModelError): generate_with_retries(client, _request(), RetryPolicy(max_retries=3)) assert len(client.requests) == 1 def test_retry_count_is_recorded_as_zero_without_retry() -> None: client = FakeModelClient([{"selected": [], "ambiguous_alternatives": []}]) response = generate_with_retries(client, _request(), RetryPolicy(max_retries=3)) assert response.retry_count == 0 assert len(client.requests) == 1