From 509dd2158c73bc7f4e0d8f2120314e9a2f999f4a Mon Sep 17 00:00:00 2001 From: Brian Sparker Date: Wed, 12 Aug 2026 20:16:46 -0700 Subject: [PATCH] chore: consolidate nightly hardening PRs into one reviewed change set --- README.md | 7 +- examples/configs/local_model.yaml | 3 +- promptlens/cli.py | 18 ++- promptlens/models/config.py | 109 +++++++++++++++++- promptlens/models/test_case.py | 14 +-- promptlens/models/tools.py | 5 +- promptlens/providers/factory.py | 23 +++- promptlens/providers/http.py | 39 ++++++- promptlens/utils/retry.py | 23 +++- tests/test_cli_config_loading.py | 30 +++++ tests/test_config_validation_hardening.py | 29 +++++ .../test_http_provider_empty_content_guard.py | 33 ++++++ tests/test_http_provider_response_parsing.py | 38 ++++++ ...est_provider_config_endpoint_validation.py | 19 +++ tests/test_provider_factory_hardening.py | 65 +++++++++++ tests/test_retry_backoff_jitter.py | 60 ++++++++++ tests/test_retry_hardening.py | 39 +++++++ tests/test_run_config_model_uniqueness.py | 37 ++++++ 18 files changed, 561 insertions(+), 30 deletions(-) create mode 100644 tests/test_cli_config_loading.py create mode 100644 tests/test_config_validation_hardening.py create mode 100644 tests/test_http_provider_empty_content_guard.py create mode 100644 tests/test_provider_config_endpoint_validation.py create mode 100644 tests/test_provider_factory_hardening.py create mode 100644 tests/test_retry_backoff_jitter.py create mode 100644 tests/test_retry_hardening.py create mode 100644 tests/test_run_config_model_uniqueness.py diff --git a/README.md b/README.md index 1435635..21f975e 100644 --- a/README.md +++ b/README.md @@ -248,10 +248,9 @@ models: - name: "Local Llama" provider: http model: llama3.1:8b + endpoint: "http://localhost:11434/api/generate" temperature: 0.7 max_tokens: 1024 - additional_params: - endpoint: "http://localhost:11434/api/generate" ``` **Setup:** @@ -270,10 +269,12 @@ models: - name: "Display Name" # Human-readable name provider: anthropic # anthropic, openai, google, http model: model-identifier # Model ID + endpoint: "http://..." # Optional, for HTTP provider temperature: 0.7 # 0.0-1.0 max_tokens: 1024 # Maximum output tokens additional_params: # Provider-specific params - endpoint: "http://..." # For HTTP provider + # endpoint under additional_params is still accepted for backwards compatibility + custom_option: "value" ``` ### Judge diff --git a/examples/configs/local_model.yaml b/examples/configs/local_model.yaml index e642597..448d722 100644 --- a/examples/configs/local_model.yaml +++ b/examples/configs/local_model.yaml @@ -13,10 +13,9 @@ models: - name: "Local Llama 3.1 8B" provider: http model: llama3.1:8b + endpoint: "http://localhost:11434/api/generate" temperature: 0.7 max_tokens: 1024 - additional_params: - endpoint: "http://localhost:11434/api/generate" judge: provider: anthropic diff --git a/promptlens/cli.py b/promptlens/cli.py index 235963a..f840d9a 100644 --- a/promptlens/cli.py +++ b/promptlens/cli.py @@ -21,6 +21,21 @@ from promptlens.models.config import RunConfig from promptlens.runners.runner import Runner + +def _load_config_data(config_path: str) -> dict: + """Load and validate top-level config structure from YAML.""" + with open(config_path, "r") as f: + config_data = yaml.safe_load(f) + + if config_data is None: + raise ValueError("Configuration file is empty") + if not isinstance(config_data, dict): + raise ValueError( + f"Configuration must be a YAML object at top level, got {type(config_data).__name__}" + ) + + return config_data + # Load environment variables load_dotenv() @@ -99,8 +114,7 @@ def run( try: # Load config console.print(f"\n[cyan]Loading configuration from {config}...[/cyan]") - with open(config, "r") as f: - config_data = yaml.safe_load(f) + config_data = _load_config_data(config) # Override with CLI options if golden_set: diff --git a/promptlens/models/config.py b/promptlens/models/config.py index 7c85798..4c0a117 100644 --- a/promptlens/models/config.py +++ b/promptlens/models/config.py @@ -1,8 +1,9 @@ """Configuration data models.""" from typing import Any, Dict, List, Optional +from urllib.parse import urlparse -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator class ProviderConfig(BaseModel): @@ -28,6 +29,20 @@ class ProviderConfig(BaseModel): endpoint: Optional[str] = None additional_params: Dict[str, Any] = Field(default_factory=dict) + @field_validator("endpoint") + @classmethod + def validate_endpoint(cls, value: Optional[str]) -> Optional[str]: + """Validate endpoint URL when provided.""" + if value is None: + return value + + parsed = urlparse(value) + if parsed.scheme not in {"http", "https"}: + raise ValueError("endpoint must use http or https scheme") + if not parsed.netloc: + raise ValueError("endpoint must include a host") + return value + class ModelConfig(BaseModel): """Configuration for a model to test. @@ -36,6 +51,8 @@ class ModelConfig(BaseModel): name: Display name for the model provider: Provider name model: Model identifier + endpoint: Optional endpoint URL (for HTTP/local providers) + timeout: Optional request timeout in seconds temperature: Sampling temperature max_tokens: Maximum tokens to generate additional_params: Provider-specific parameters @@ -44,10 +61,26 @@ class ModelConfig(BaseModel): name: str provider: str model: str + endpoint: Optional[str] = None + timeout: Optional[int] = None temperature: float = 0.7 max_tokens: int = 1024 additional_params: Dict[str, Any] = Field(default_factory=dict) + @field_validator("temperature") + @classmethod + def validate_temperature(cls, value: float) -> float: + if not 0.0 <= value <= 2.0: + raise ValueError("temperature must be between 0.0 and 2.0") + return value + + @field_validator("max_tokens") + @classmethod + def validate_max_tokens(cls, value: int) -> int: + if value <= 0: + raise ValueError("max_tokens must be greater than 0") + return value + class JudgeConfig(BaseModel): """Configuration for the judge. @@ -82,6 +115,34 @@ class ExecutionConfig(BaseModel): retry_delay_seconds: float = 1.0 timeout_seconds: int = 60 + @field_validator("parallel_requests") + @classmethod + def validate_parallel_requests(cls, value: int) -> int: + if value <= 0: + raise ValueError("parallel_requests must be greater than 0") + return value + + @field_validator("retry_attempts") + @classmethod + def validate_retry_attempts(cls, value: int) -> int: + if value < 0: + raise ValueError("retry_attempts must be greater than or equal to 0") + return value + + @field_validator("retry_delay_seconds") + @classmethod + def validate_retry_delay_seconds(cls, value: float) -> float: + if value < 0: + raise ValueError("retry_delay_seconds must be greater than or equal to 0") + return value + + @field_validator("timeout_seconds") + @classmethod + def validate_timeout_seconds(cls, value: int) -> int: + if value <= 0: + raise ValueError("timeout_seconds must be greater than 0") + return value + class OutputConfig(BaseModel): """Configuration for output settings. @@ -96,6 +157,18 @@ class OutputConfig(BaseModel): formats: List[str] = Field(default_factory=lambda: ["html", "json"]) run_name: Optional[str] = None + @field_validator("formats") + @classmethod + def validate_formats(cls, value: List[str]) -> List[str]: + allowed = {"html", "json", "csv", "md"} + normalized = [fmt.lower() for fmt in value] + invalid = sorted({fmt for fmt in normalized if fmt not in allowed}) + if invalid: + raise ValueError(f"unsupported output format(s): {', '.join(invalid)}") + if not normalized: + raise ValueError("output formats must contain at least one format") + return normalized + class RunConfig(BaseModel): """Complete run configuration. @@ -114,9 +187,35 @@ class RunConfig(BaseModel): execution: ExecutionConfig = Field(default_factory=ExecutionConfig) output: OutputConfig = Field(default_factory=OutputConfig) - class Config: - """Pydantic config.""" - json_schema_extra = { + @model_validator(mode="after") + def validate_models(self) -> "RunConfig": + if not self.models: + raise ValueError("models must contain at least one model configuration") + return self + + @field_validator("models") + @classmethod + def validate_models_unique(cls, models: List[ModelConfig]) -> List[ModelConfig]: + """Ensure model display names are unique to avoid ambiguous reports.""" + seen = set() + duplicates = set() + + for model in models: + normalized = model.name.strip().lower() + if normalized in seen: + duplicates.add(model.name) + else: + seen.add(normalized) + + if duplicates: + duplicate_list = ", ".join(sorted(duplicates)) + raise ValueError( + f"Model names must be unique (case-insensitive). Duplicates: {duplicate_list}" + ) + + return models + + model_config = ConfigDict(json_schema_extra={ "example": { "golden_set": "./examples/golden_sets/customer_support.yaml", "models": [ @@ -142,4 +241,4 @@ class Config: "formats": ["html", "json"], }, } - } + }) diff --git a/promptlens/models/test_case.py b/promptlens/models/test_case.py index 931208f..a95bb39 100644 --- a/promptlens/models/test_case.py +++ b/promptlens/models/test_case.py @@ -2,7 +2,7 @@ from typing import Any, Dict, List, Optional -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field from promptlens.models.tools import ToolDefinition, ExpectedToolCall @@ -50,9 +50,7 @@ class TestCase(BaseModel): description="Whether to actually execute tools (default: False, evaluation only)" ) - class Config: - """Pydantic config.""" - json_schema_extra = { + model_config = ConfigDict(json_schema_extra={ "example": { "id": "cs-001", "query": "How do I reset my password?", @@ -60,7 +58,7 @@ class Config: "category": "account_management", "tags": ["password", "account"], } - } + }) class GoldenSet(BaseModel): @@ -80,9 +78,7 @@ class GoldenSet(BaseModel): test_cases: List[TestCase] metadata: Dict[str, Any] = Field(default_factory=dict) - class Config: - """Pydantic config.""" - json_schema_extra = { + model_config = ConfigDict(json_schema_extra={ "example": { "name": "Customer Support Tests", "description": "Test cases for customer support chatbot", @@ -97,4 +93,4 @@ class Config: } ], } - } + }) diff --git a/promptlens/models/tools.py b/promptlens/models/tools.py index bd7d52b..011447f 100644 --- a/promptlens/models/tools.py +++ b/promptlens/models/tools.py @@ -8,7 +8,7 @@ """ from typing import Any, Dict, List, Optional -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field class ToolParameter(BaseModel): @@ -24,8 +24,7 @@ class ToolParameter(BaseModel): properties: Optional[Dict[str, "ToolParameter"]] = Field(None, description="For object types, nested properties") items: Optional["ToolParameter"] = Field(None, description="For array types, the item schema") - class Config: - extra = "allow" # Allow additional JSON Schema fields + model_config = ConfigDict(extra="allow") # Allow additional JSON Schema fields class ToolDefinition(BaseModel): diff --git a/promptlens/providers/factory.py b/promptlens/providers/factory.py index 4f78907..3f1954f 100644 --- a/promptlens/providers/factory.py +++ b/promptlens/providers/factory.py @@ -32,22 +32,39 @@ def get_provider(model_config: ModelConfig) -> BaseProvider: Raises: ValueError: If provider is not supported """ - provider_name = model_config.provider.lower() + provider_name = model_config.provider.strip().lower() + + if not provider_name: + available = ", ".join(sorted(PROVIDER_REGISTRY.keys())) + raise ValueError( + "Provider name cannot be empty. " + f"Available providers: {available}" + ) if provider_name not in PROVIDER_REGISTRY: - available = ", ".join(PROVIDER_REGISTRY.keys()) + available = ", ".join(sorted(PROVIDER_REGISTRY.keys())) raise ValueError( f"Provider '{provider_name}' not supported. " f"Available providers: {available}" ) + # Copy params so we can normalize legacy keys without mutating input + additional_params = dict(model_config.additional_params) + + # Backwards compatibility: allow endpoint under additional_params for HTTP provider + endpoint = model_config.endpoint + if provider_name == "http" and endpoint is None: + endpoint = additional_params.pop("endpoint", None) + # Convert ModelConfig to ProviderConfig provider_config = ProviderConfig( name=provider_name, model=model_config.model, + endpoint=endpoint, + timeout=model_config.timeout or 60, temperature=model_config.temperature, max_tokens=model_config.max_tokens, - additional_params=model_config.additional_params, + additional_params=additional_params, ) # Get provider class and instantiate diff --git a/promptlens/providers/http.py b/promptlens/providers/http.py index f1828e2..4db3fc5 100644 --- a/promptlens/providers/http.py +++ b/promptlens/providers/http.py @@ -1,5 +1,6 @@ """Generic HTTP provider for local models (Ollama, LM Studio, etc.).""" +import json import logging from datetime import datetime from typing import Any, Dict, List, Optional @@ -59,6 +60,19 @@ def _extract_content(data: Dict[str, Any]) -> str: return "" + @staticmethod + def _extract_error_message(data: Dict[str, Any]) -> str: + """Extract an error message from common HTTP error response shapes.""" + for key in ("error", "message", "detail"): + value = data.get(key) + if isinstance(value, str) and value.strip(): + return value.strip() + if isinstance(value, dict): + nested_message = value.get("message") + if isinstance(nested_message, str) and nested_message.strip(): + return nested_message.strip() + return "" + def __init__(self, config: ProviderConfig) -> None: """Initialize the HTTP provider. @@ -119,11 +133,34 @@ async def _make_request() -> ModelResponse: timeout=aiohttp.ClientTimeout(total=self.config.timeout), ) as response: response.raise_for_status() - data = await response.json() + raw_text = await response.text() + + try: + data = json.loads(raw_text) + except json.JSONDecodeError as exc: + logger.error( + "HTTP provider returned non-JSON response from %s: %s", + self.endpoint, + raw_text[:500], + ) + raise ValueError( + "HTTP provider expected JSON response but received non-JSON content" + ) from exc # Extract content (try common response formats) content = self._extract_content(data) + if not content.strip(): + response_error = self._extract_error_message(data) + if response_error: + raise ValueError(f"HTTP endpoint returned error payload: {response_error}") + + top_level_keys = sorted(list(data.keys())) + raise ValueError( + "HTTP endpoint response did not contain supported content fields " + f"(keys: {top_level_keys})" + ) + # Local models typically don't provide token counts or cost return ModelResponse( content=content, diff --git a/promptlens/utils/retry.py b/promptlens/utils/retry.py index 7834517..dc96ffe 100644 --- a/promptlens/utils/retry.py +++ b/promptlens/utils/retry.py @@ -2,6 +2,7 @@ import asyncio import logging +import random from typing import Any, Callable, TypeVar logger = logging.getLogger(__name__) @@ -16,6 +17,7 @@ async def retry_with_exponential_backoff( backoff_factor: float = 2.0, max_delay: float = 60.0, retry_on: tuple[type[Exception], ...] = (Exception,), + jitter_ratio: float = 0.1, ) -> T: """Retry a function with exponential backoff. @@ -26,6 +28,7 @@ async def retry_with_exponential_backoff( backoff_factor: Factor to multiply delay by after each retry max_delay: Maximum delay between retries in seconds retry_on: Tuple of exception types to retry on + jitter_ratio: Random jitter as a fraction of delay (0.1 = +/-10%) Returns: Result from the function @@ -33,6 +36,18 @@ async def retry_with_exponential_backoff( Raises: Exception: The last exception if all retries fail """ + if max_attempts < 1: + raise ValueError("max_attempts must be >= 1") + + if initial_delay < 0: + raise ValueError("initial_delay must be >= 0") + + if backoff_factor <= 0: + raise ValueError("backoff_factor must be > 0") + + if max_delay <= 0: + raise ValueError("max_delay must be > 0") + delay = initial_delay last_exception = None @@ -46,11 +61,15 @@ async def retry_with_exponential_backoff( logger.error(f"All {max_attempts} attempts failed. Last error: {e}") raise + jitter_ratio = max(0.0, jitter_ratio) + jitter = random.uniform(-jitter_ratio, jitter_ratio) if jitter_ratio else 0.0 + sleep_for = max(0.0, delay * (1 + jitter)) + logger.warning( f"Attempt {attempt + 1}/{max_attempts} failed: {e}. " - f"Retrying in {delay:.1f}s..." + f"Retrying in {sleep_for:.2f}s..." ) - await asyncio.sleep(delay) + await asyncio.sleep(sleep_for) delay = min(delay * backoff_factor, max_delay) if last_exception: diff --git a/tests/test_cli_config_loading.py b/tests/test_cli_config_loading.py new file mode 100644 index 0000000..e313e6f --- /dev/null +++ b/tests/test_cli_config_loading.py @@ -0,0 +1,30 @@ +from pathlib import Path + +import pytest + +from promptlens.cli import _load_config_data + + +def test_load_config_data_rejects_empty_file(tmp_path: Path) -> None: + config = tmp_path / "config.yaml" + config.write_text("") + + with pytest.raises(ValueError, match="empty"): + _load_config_data(str(config)) + + +def test_load_config_data_rejects_non_mapping_top_level(tmp_path: Path) -> None: + config = tmp_path / "config.yaml" + config.write_text("- item\n- item2\n") + + with pytest.raises(ValueError, match="top level"): + _load_config_data(str(config)) + + +def test_load_config_data_accepts_mapping(tmp_path: Path) -> None: + config = tmp_path / "config.yaml" + config.write_text("golden_set: ./tests.yaml\nmodels: []\n") + + loaded = _load_config_data(str(config)) + + assert loaded["golden_set"] == "./tests.yaml" diff --git a/tests/test_config_validation_hardening.py b/tests/test_config_validation_hardening.py new file mode 100644 index 0000000..6b74334 --- /dev/null +++ b/tests/test_config_validation_hardening.py @@ -0,0 +1,29 @@ +import pytest +from pydantic import ValidationError + +from promptlens.models.config import ExecutionConfig, ModelConfig, OutputConfig, RunConfig + + +def test_model_config_temperature_bounds() -> None: + with pytest.raises(ValidationError): + ModelConfig(name="m", provider="openai", model="gpt", temperature=2.5) + + +def test_execution_config_parallel_requests_must_be_positive() -> None: + with pytest.raises(ValidationError): + ExecutionConfig(parallel_requests=0) + + +def test_output_formats_are_normalized_to_lowercase() -> None: + cfg = OutputConfig(formats=["HTML", "JSON"]) + assert cfg.formats == ["html", "json"] + + +def test_output_formats_reject_unsupported_values() -> None: + with pytest.raises(ValidationError): + OutputConfig(formats=["html", "xml"]) + + +def test_run_config_requires_at_least_one_model() -> None: + with pytest.raises(ValidationError): + RunConfig(golden_set="./tests.yaml", models=[]) diff --git a/tests/test_http_provider_empty_content_guard.py b/tests/test_http_provider_empty_content_guard.py new file mode 100644 index 0000000..cde27d1 --- /dev/null +++ b/tests/test_http_provider_empty_content_guard.py @@ -0,0 +1,33 @@ +from promptlens.models.config import ProviderConfig +from promptlens.providers.http import HTTPProvider + + +def _provider() -> HTTPProvider: + return HTTPProvider( + ProviderConfig( + name="http", + model="test-model", + endpoint="http://localhost:11434/api/generate", + ) + ) + + +def test_extract_error_message_from_nested_error() -> None: + provider = _provider() + + assert ( + provider._extract_error_message({"error": {"message": "model overloaded"}}) + == "model overloaded" + ) + + +def test_extract_error_message_from_top_level_detail() -> None: + provider = _provider() + + assert provider._extract_error_message({"detail": "service unavailable"}) == "service unavailable" + + +def test_extract_error_message_returns_empty_when_absent() -> None: + provider = _provider() + + assert provider._extract_error_message({"foo": "bar"}) == "" diff --git a/tests/test_http_provider_response_parsing.py b/tests/test_http_provider_response_parsing.py index 0949b6e..4ef58a2 100644 --- a/tests/test_http_provider_response_parsing.py +++ b/tests/test_http_provider_response_parsing.py @@ -1,3 +1,5 @@ +import pytest + from promptlens.models.config import ProviderConfig from promptlens.providers.http import HTTPProvider @@ -57,3 +59,39 @@ def test_extract_content_returns_empty_for_unknown_shape() -> None: provider = _provider() assert provider._extract_content({"foo": "bar"}) == "" + + +@pytest.mark.asyncio +async def test_http_provider_returns_clear_error_for_non_json_response() -> None: + provider = _provider() + + class MockResponse: + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb): + return False + + def raise_for_status(self) -> None: + return None + + async def text(self) -> str: + return "not-json" + + class MockSession: + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb): + return False + + def post(self, *args, **kwargs): + return MockResponse() + + from unittest.mock import patch + + with patch("promptlens.providers.http.aiohttp.ClientSession", return_value=MockSession()): + result = await provider.generate("hello") + + assert result.error is not None + assert "expected JSON response" in result.error diff --git a/tests/test_provider_config_endpoint_validation.py b/tests/test_provider_config_endpoint_validation.py new file mode 100644 index 0000000..304e641 --- /dev/null +++ b/tests/test_provider_config_endpoint_validation.py @@ -0,0 +1,19 @@ +import pytest +from pydantic import ValidationError + +from promptlens.models.config import ProviderConfig + + +def test_provider_config_accepts_http_endpoint() -> None: + config = ProviderConfig(name="http", model="m", endpoint="http://localhost:11434/api/generate") + assert config.endpoint == "http://localhost:11434/api/generate" + + +def test_provider_config_rejects_non_http_scheme() -> None: + with pytest.raises(ValidationError, match="endpoint must use http or https scheme"): + ProviderConfig(name="http", model="m", endpoint="file:///tmp/model.sock") + + +def test_provider_config_rejects_missing_host() -> None: + with pytest.raises(ValidationError, match="endpoint must include a host"): + ProviderConfig(name="http", model="m", endpoint="https:///api/generate") diff --git a/tests/test_provider_factory_hardening.py b/tests/test_provider_factory_hardening.py new file mode 100644 index 0000000..2f1787f --- /dev/null +++ b/tests/test_provider_factory_hardening.py @@ -0,0 +1,65 @@ +from promptlens.models.config import ModelConfig +from promptlens.providers.factory import get_provider +from promptlens.providers.http import HTTPProvider + + +def test_http_provider_supports_legacy_endpoint_in_additional_params() -> None: + provider = get_provider( + ModelConfig( + name="Local model", + provider="http", + model="llama3.1:8b", + additional_params={ + "endpoint": "http://localhost:11434/api/generate", + "temperature": 0.2, + }, + ) + ) + + assert isinstance(provider, HTTPProvider) + assert provider.endpoint == "http://localhost:11434/api/generate" + assert "endpoint" not in provider.config.additional_params + assert provider.config.additional_params["temperature"] == 0.2 + + +def test_http_provider_accepts_top_level_endpoint() -> None: + provider = get_provider( + ModelConfig( + name="Local model", + provider="http", + model="llama3.1:8b", + endpoint="http://localhost:11434/api/generate", + ) + ) + + assert isinstance(provider, HTTPProvider) + assert provider.endpoint == "http://localhost:11434/api/generate" + + +def test_get_provider_normalizes_whitespace_and_case(monkeypatch): + monkeypatch.setenv("OPENAI_API_KEY", "test-key") + model = ModelConfig( + name="OpenAI", + provider=" OpenAI ", + model="gpt-4o-mini", + ) + + provider = get_provider(model) + + assert provider.config.name == "openai" + + +def test_get_provider_rejects_blank_provider_name(): + model = ModelConfig( + name="Bad", + provider=" ", + model="gpt-4o-mini", + ) + + try: + get_provider(model) + assert False, "Expected ValueError for blank provider name" + except ValueError as exc: + msg = str(exc) + assert "cannot be empty" in msg + assert "Available providers:" in msg diff --git a/tests/test_retry_backoff_jitter.py b/tests/test_retry_backoff_jitter.py new file mode 100644 index 0000000..9d3b8c7 --- /dev/null +++ b/tests/test_retry_backoff_jitter.py @@ -0,0 +1,60 @@ +import pytest + +from promptlens.utils import retry as retry_module +from promptlens.utils.retry import retry_with_exponential_backoff + + +@pytest.mark.asyncio +async def test_retry_applies_jitter_to_sleep(monkeypatch): + calls = {"count": 0} + sleeps = [] + + async def flaky(): + calls["count"] += 1 + if calls["count"] < 3: + raise ValueError("boom") + return "ok" + + async def fake_sleep(delay): + sleeps.append(delay) + + monkeypatch.setattr(retry_module.asyncio, "sleep", fake_sleep) + monkeypatch.setattr(retry_module.random, "uniform", lambda a, b: 0.1) + + result = await retry_with_exponential_backoff( + flaky, + max_attempts=3, + initial_delay=1.0, + backoff_factor=2.0, + jitter_ratio=0.1, + ) + + assert result == "ok" + assert sleeps == [1.1, 2.2] + + +@pytest.mark.asyncio +async def test_retry_negative_jitter_ratio_is_clamped(monkeypatch): + calls = {"count": 0} + sleeps = [] + + async def flaky(): + calls["count"] += 1 + if calls["count"] < 2: + raise ValueError("boom") + return "ok" + + async def fake_sleep(delay): + sleeps.append(delay) + + monkeypatch.setattr(retry_module.asyncio, "sleep", fake_sleep) + + result = await retry_with_exponential_backoff( + flaky, + max_attempts=2, + initial_delay=1.0, + jitter_ratio=-0.5, + ) + + assert result == "ok" + assert sleeps == [1.0] diff --git a/tests/test_retry_hardening.py b/tests/test_retry_hardening.py new file mode 100644 index 0000000..2d20729 --- /dev/null +++ b/tests/test_retry_hardening.py @@ -0,0 +1,39 @@ +import pytest + +from promptlens.utils.retry import retry_with_exponential_backoff + + +@pytest.mark.asyncio +async def test_retry_rejects_non_positive_attempts() -> None: + async def always_fails() -> str: + raise RuntimeError("boom") + + with pytest.raises(ValueError, match="max_attempts"): + await retry_with_exponential_backoff(always_fails, max_attempts=0) + + +@pytest.mark.asyncio +async def test_retry_rejects_negative_initial_delay() -> None: + async def always_fails() -> str: + raise RuntimeError("boom") + + with pytest.raises(ValueError, match="initial_delay"): + await retry_with_exponential_backoff(always_fails, initial_delay=-0.1) + + +@pytest.mark.asyncio +async def test_retry_rejects_invalid_backoff_factor() -> None: + async def always_fails() -> str: + raise RuntimeError("boom") + + with pytest.raises(ValueError, match="backoff_factor"): + await retry_with_exponential_backoff(always_fails, backoff_factor=0) + + +@pytest.mark.asyncio +async def test_retry_rejects_non_positive_max_delay() -> None: + async def always_fails() -> str: + raise RuntimeError("boom") + + with pytest.raises(ValueError, match="max_delay"): + await retry_with_exponential_backoff(always_fails, max_delay=0) diff --git a/tests/test_run_config_model_uniqueness.py b/tests/test_run_config_model_uniqueness.py new file mode 100644 index 0000000..7ae4a28 --- /dev/null +++ b/tests/test_run_config_model_uniqueness.py @@ -0,0 +1,37 @@ +from pydantic import ValidationError + +from promptlens.models.config import RunConfig + + +def _base_config(models): + return { + "golden_set": "./examples/golden_sets/customer_support.yaml", + "models": models, + } + + +def test_run_config_rejects_duplicate_model_names_case_insensitive(): + config = _base_config( + [ + {"name": "Claude Fast", "provider": "anthropic", "model": "claude-3-5-sonnet-20241022"}, + {"name": "claude fast", "provider": "openai", "model": "gpt-4o-mini"}, + ] + ) + + try: + RunConfig(**config) + raise AssertionError("Expected ValidationError for duplicate model names") + except ValidationError as exc: + assert "Model names must be unique" in str(exc) + + +def test_run_config_allows_unique_model_names(): + config = _base_config( + [ + {"name": "Claude Fast", "provider": "anthropic", "model": "claude-3-5-sonnet-20241022"}, + {"name": "GPT Fast", "provider": "openai", "model": "gpt-4o-mini"}, + ] + ) + + run_config = RunConfig(**config) + assert len(run_config.models) == 2