Skip to content
Merged

78 #94

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
3 changes: 1 addition & 2 deletions src/promptum/providers/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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})"
)
7 changes: 2 additions & 5 deletions src/promptum/providers/openrouter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
12 changes: 4 additions & 8 deletions src/promptum/session/report.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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]
Expand Down Expand Up @@ -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)
7 changes: 7 additions & 0 deletions tests/benchmark/test_report_filtering.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
4 changes: 1 addition & 3 deletions tests/providers/test_openrouter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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!"
1 change: 1 addition & 0 deletions tests/validation/test_json_schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
1 change: 1 addition & 0 deletions tests/validation/test_regex.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down