diff --git a/src/promptum/providers/exceptions.py b/src/promptum/providers/exceptions.py index 65d2230..3ccc96f 100644 --- a/src/promptum/providers/exceptions.py +++ b/src/promptum/providers/exceptions.py @@ -50,6 +50,5 @@ def __init__( self.last_response_body = last_response_body self.retry_delays = retry_delays super().__init__( - f"Request failed after {attempts} attempts" - f" (last status {last_status_code})" + f"Request failed after {attempts} attempts (last status {last_status_code})" ) diff --git a/src/promptum/providers/openrouter.py b/src/promptum/providers/openrouter.py index a20991c..883c2f9 100644 --- a/src/promptum/providers/openrouter.py +++ b/src/promptum/providers/openrouter.py @@ -78,10 +78,9 @@ async def generate( if conflicts: raise ValueError( f"Cannot override reserved payload fields: {', '.join(sorted(conflicts))}" - ) + ) payload.update(kwargs) - for attempt in range(config.max_attempts): start_time = time.perf_counter() try: @@ -133,9 +132,7 @@ async def generate( retry_delays.append(delay) await self._sleep(delay) else: - raise ProviderTransientError( - config.max_attempts, retry_delays - ) from e + raise ProviderTransientError(config.max_attempts, retry_delays) from e raise ProviderRetryExhaustedError( config.max_attempts, last_status_code, last_response_body, retry_delays diff --git a/src/promptum/session/report.py b/src/promptum/session/report.py index ee10cf4..756a485 100644 --- a/src/promptum/session/report.py +++ b/src/promptum/session/report.py @@ -15,8 +15,8 @@ def __post_init__(self): def get_summary(self) -> Summary: total = len(self.results) passed = sum(1 for r in self.results if r.passed) - execution_errors=self._count_execution_errors() - validation_failures=self._count_validation_failures() + execution_errors = self._count_execution_errors() + validation_failures = self._count_validation_failures() latencies = [r.metrics.latency_ms for r in self.results if r.metrics] total_cost = sum(r.metrics.cost_usd or 0 for r in self.results if r.metrics) @@ -42,7 +42,7 @@ def filter( tags: Sequence[str] | None = None, passed: bool | None = None, ) -> "Report": - filtered = list(self.results) + filtered = self.results if model is not None: filtered = [r for r in filtered if r.test_case.model == model] @@ -71,8 +71,4 @@ def _count_execution_errors(self) -> int: return sum(1 for r in self.results if r.execution_error is not None) def _count_validation_failures(self) -> int: - return sum( - 1 - for r in self.results - if not r.passed and r.execution_error is None - ) + return sum(1 for r in self.results if not r.passed and r.execution_error is None) diff --git a/tests/benchmark/test_report_filtering.py b/tests/benchmark/test_report_filtering.py index 1b12f88..cadbdc9 100644 --- a/tests/benchmark/test_report_filtering.py +++ b/tests/benchmark/test_report_filtering.py @@ -40,3 +40,10 @@ def test_report_group_by_model(sample_report: Report) -> None: assert "model2" in grouped assert len(grouped["model1"].results) == 2 assert len(grouped["model2"].results) == 1 + + +def test_report_no_filter(sample_report: Report) -> None: + unfiltered = sample_report.filter(model=None, tags=None, passed=None) + assert isinstance(unfiltered, Report) + assert len(unfiltered.results) == len(sample_report.results) + assert unfiltered == sample_report diff --git a/tests/providers/test_openrouter.py b/tests/providers/test_openrouter.py index 2d1ffd5..16b678d 100644 --- a/tests/providers/test_openrouter.py +++ b/tests/providers/test_openrouter.py @@ -367,8 +367,6 @@ async def test_generate_uses_per_call_retry_config_over_default( client._client.post = AsyncMock(side_effect=responses) client._sleep = AsyncMock() - content, _ = await client.generate( - prompt="hello", model="m", retry_config=per_call_config - ) + content, _ = await client.generate(prompt="hello", model="m", retry_config=per_call_config) assert content == "Hello, world!" diff --git a/tests/validation/test_json_schema.py b/tests/validation/test_json_schema.py index 272cfc4..d63dc0a 100644 --- a/tests/validation/test_json_schema.py +++ b/tests/validation/test_json_schema.py @@ -42,6 +42,7 @@ def test_json_schema_describe_includes_validator_info_and_required_keys() -> Non assert "a" in description assert "b" in description + def test_json_schema_describe_without_required_keys() -> None: validator = JsonSchema(required_keys=()) description = validator.describe() diff --git a/tests/validation/test_regex.py b/tests/validation/test_regex.py index fe3d4ce..6f065a7 100644 --- a/tests/validation/test_regex.py +++ b/tests/validation/test_regex.py @@ -32,6 +32,7 @@ def test_regex_describe_includes_validator_name_and_pattern() -> None: assert "Regex" in description assert r"\d+" in description + def test_regex_with_ignorecase_flag() -> None: validator = Regex(r"hello", flags=re.IGNORECASE)