Skip to content

Commit 68fc64a

Browse files
romanlutzCopilot
andcommitted
Merge repaired #2376 parent into #2377
Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 5d02c2d5-b499-4f78-a04d-03bffa750817
2 parents ae392e6 + 0285314 commit 68fc64a

65 files changed

Lines changed: 1349 additions & 561 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎pyproject.toml‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -222,6 +222,12 @@ include = ["pyrit/prompt_target/hugging_face/**"]
222222
[tool.ty.overrides.rules]
223223
invalid-argument-type = "ignore"
224224

225+
# Historical Alembic revisions are immutable; retain their defensive rowcount checks.
226+
[[tool.ty.overrides]]
227+
include = ["pyrit/memory/alembic/versions/1b3d5f7a9c2e_persist_scored_expectation.py"]
228+
[tool.ty.overrides.rules]
229+
redundant-condition-strict = "ignore"
230+
225231
# Unit tests intentionally exercise runtime validation failures, mutable test doubles,
226232
# and optional fields after fixtures have populated them. Keep library code strict while
227233
# avoiding hundreds of per-assertion ignores in tests.

‎pyrit/auth/azure_auth.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -206,9 +206,9 @@ async def async_token_provider() -> str: # pyrit-async-suffix-exempt
206206
str: The token string from the synchronous provider.
207207
"""
208208
result = api_key()
209-
if inspect.isawaitable(result):
210-
return await result # type: ignore[ty:invalid-return-type]
211-
return result
209+
if isinstance(result, str):
210+
return result
211+
return await result
212212

213213
return async_token_provider
214214

‎pyrit/backend/mappers/attack_mappers.py‎

Lines changed: 13 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -262,15 +262,23 @@ def _resolve_summary_timestamps(ar: AttackResult) -> tuple[datetime, datetime]:
262262
Returns:
263263
A ``(created_at, updated_at)`` tuple.
264264
"""
265-
created_str = ar.metadata.get("created_at")
265+
return _resolve_timestamps(created_str=ar.metadata.get("created_at"), timestamp=ar.timestamp)
266+
267+
268+
def _resolve_timestamps(*, created_str: str | None, timestamp: datetime | None) -> tuple[datetime, datetime]:
269+
"""
270+
Resolve display times, retaining fallbacks for unpersisted mutable results.
271+
272+
Returns:
273+
tuple[datetime, datetime]: Creation and last-update timestamps.
274+
"""
266275
if created_str:
267276
created_at = datetime.fromisoformat(created_str)
268-
elif ar.timestamp is not None:
269-
created_at = ar.timestamp
277+
elif timestamp is not None:
278+
created_at = timestamp
270279
else:
271280
created_at = datetime.now(UTC)
272-
updated_at = ar.timestamp if ar.timestamp is not None else created_at
273-
return created_at, updated_at
281+
return created_at, timestamp if timestamp is not None else created_at
274282

275283

276284
async def _summary_last_response_async(piece: MessagePiece | None) -> MessagePieceView | None:

‎pyrit/backend/services/configuration_file_service.py‎

Lines changed: 5 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
import os
99
import tempfile
1010
from collections.abc import AsyncGenerator
11-
from contextlib import asynccontextmanager
11+
from contextlib import ExitStack, asynccontextmanager
1212
from pathlib import Path
1313
from urllib.parse import urlparse
1414

@@ -103,9 +103,8 @@ def _validate_configuration_content(content: str) -> None:
103103

104104
def _replace_local_config_file(*, path: Path, content: str) -> None:
105105
"""Atomically replace a local configuration file."""
106-
temporary_path: Path | None = None
107-
try:
108-
path.parent.mkdir(parents=True, exist_ok=True)
106+
path.parent.mkdir(parents=True, exist_ok=True)
107+
with ExitStack() as cleanup:
109108
with tempfile.NamedTemporaryFile(
110109
mode="w",
111110
encoding="utf-8",
@@ -114,12 +113,10 @@ def _replace_local_config_file(*, path: Path, content: str) -> None:
114113
suffix=".tmp",
115114
delete=False,
116115
) as temporary_file:
117-
temporary_file.write(content)
118116
temporary_path = Path(temporary_file.name)
117+
cleanup.callback(temporary_path.unlink, missing_ok=True)
118+
temporary_file.write(content)
119119
os.replace(temporary_path, path)
120-
finally:
121-
if temporary_path is not None:
122-
temporary_path.unlink(missing_ok=True)
123120

124121

125122
class ConfigurationFileService:

‎pyrit/backend/services/scenario_progress_read_model.py‎

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -364,10 +364,19 @@ def calculate_progress_counts(
364364
@staticmethod
365365
def _result_order_key(attack_result: AttackResult) -> tuple[datetime, str]:
366366
"""Return a deterministic chronological key for one hydrated result attempt."""
367-
timestamp = attack_result.timestamp
368-
if not isinstance(timestamp, datetime):
369-
timestamp = datetime.min.replace(tzinfo=UTC)
370-
return timestamp, str(attack_result.attack_result_id)
367+
return ScenarioProgressReadModel._timestamp_order_key(attack_result.timestamp), str(
368+
attack_result.attack_result_id
369+
)
370+
371+
@staticmethod
372+
def _timestamp_order_key(timestamp: object) -> datetime:
373+
"""
374+
Normalize potentially malformed timestamps from mutable result objects.
375+
376+
Returns:
377+
datetime: The timestamp or a stable earliest-time fallback.
378+
"""
379+
return timestamp if isinstance(timestamp, datetime) else datetime.min.replace(tzinfo=UTC)
371380

372381
@staticmethod
373382
def total_retry_pressure(*, attempts_per_unit: Iterable[int], persisted_retries: Iterable[int]) -> int:

‎pyrit/backend/services/scenario_run_service.py‎

Lines changed: 18 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -893,10 +893,14 @@ def _schedule_terminalization_retry(self, *, active: _ActiveTask) -> None:
893893
self._handoff_retry_tasks.add(retry_task)
894894
retry_task.add_done_callback(self._handoff_retry_tasks.discard)
895895

896+
def _can_retry_active_run(self, *, scenario_result_id: str) -> bool:
897+
"""Return whether retry work may continue for the active run."""
898+
return not self._stopping and self._active_scenario_result_id == scenario_result_id
899+
896900
async def _retry_handoff_async(self, *, scenario_result_id: str) -> None:
897901
"""Retry scheduler handoff with bounded exponential delay until it succeeds or shutdown begins."""
898902
delay = _SCHEDULER_RETRY_INITIAL_SECONDS
899-
while not self._stopping and self._active_scenario_result_id == scenario_result_id:
903+
while self._can_retry_active_run(scenario_result_id=scenario_result_id):
900904
await asyncio.sleep(delay)
901905
try:
902906
await self._handoff_scheduler_async(scenario_result_id=scenario_result_id)
@@ -909,11 +913,11 @@ async def _retry_handoff_async(self, *, scenario_result_id: str) -> None:
909913
async def _retry_terminalization_async(self, *, active: _ActiveTask) -> None:
910914
"""Retry a failed cancellation transition, then perform the terminal handoff."""
911915
delay = _SCHEDULER_RETRY_INITIAL_SECONDS
912-
while not self._stopping and self._active_scenario_result_id == active.scenario_result_id:
916+
while self._can_retry_active_run(scenario_result_id=active.scenario_result_id):
913917
await asyncio.sleep(delay)
914918
try:
915919
async with self._scheduler_lock:
916-
if self._stopping or self._active_scenario_result_id != active.scenario_result_id:
920+
if not self._can_retry_active_run(scenario_result_id=active.scenario_result_id):
917921
return
918922
await asyncio.to_thread(
919923
self._memory.try_update_scenario_run_state,
@@ -1447,6 +1451,16 @@ def _load_started_at(*, scenario_result: ScenarioResult) -> datetime | None:
14471451
return None
14481452
return started_at if started_at.tzinfo is not None else None
14491453

1454+
@staticmethod
1455+
def _identifier_techniques(scenario_identifier: ScenarioIdentifier | None) -> list[str]:
1456+
"""
1457+
Read techniques when legacy persisted metadata has an identifier.
1458+
1459+
Returns:
1460+
list[str]: Stored techniques or an empty list.
1461+
"""
1462+
return list(scenario_identifier.techniques or []) if scenario_identifier is not None else []
1463+
14501464
@staticmethod
14511465
def _safe_run_metadata(
14521466
*,
@@ -1793,10 +1807,8 @@ def get_run_progress_from_storage(
17931807
target, datasets_used, scenario_parameters = self._safe_run_metadata(scenario_identifier=scenario_identifier)
17941808
if plan is not None:
17951809
techniques_used = list(dict.fromkeys(group.display_group for group in plan.atomic_groups))
1796-
elif scenario_identifier is not None:
1797-
techniques_used = list(scenario_identifier.techniques or [])
17981810
else:
1799-
techniques_used = []
1811+
techniques_used = self._identifier_techniques(scenario_identifier)
18001812
return ScenarioRunProgress(
18011813
run=ScenarioProgressHeader(
18021814
scenario_result_id=scenario_result_id,

‎pyrit/cli/_server_launcher.py‎

Lines changed: 17 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -825,18 +825,7 @@ async def start_async(
825825
_logger.info("Backend launcher PID: %d (logs: %s)", process.pid, self._log_path)
826826
await self._wait_for_readiness_async(plan=plan, host=host, startup_state=startup_state)
827827
finally:
828-
if startup_state.phase is not _StartupPhase.READY and startup_state.process is not None:
829-
if self._process is None:
830-
self._process = startup_state.process
831-
self._listener_pid = startup_state.process.pid
832-
self._port = port if startup_state.pid_record_written else None
833-
try:
834-
await startup_state.cleanup_async()
835-
finally:
836-
if startup_state.cleanup_succeeded:
837-
self._clear_process_state()
838-
elif startup_state.cleanup_succeeded is False:
839-
_logger.warning("Failed to stop backend launcher process %d", startup_state.process.pid)
828+
await self._cleanup_failed_startup_async(startup_state=startup_state, port=port)
840829

841830
if startup_state.timed_out:
842831
cleanup_message = (
@@ -856,6 +845,22 @@ async def start_async(
856845
)
857846
return plan.base_url
858847

848+
async def _cleanup_failed_startup_async(self, *, startup_state: _ServerStartupState, port: int) -> None:
849+
"""Clean up even when spawning raised before the launcher stored its process."""
850+
if startup_state.phase is _StartupPhase.READY or startup_state.process is None:
851+
return
852+
if self._process is None:
853+
self._process = startup_state.process
854+
self._listener_pid = startup_state.process.pid
855+
self._port = port if startup_state.pid_record_written else None
856+
try:
857+
await startup_state.cleanup_async()
858+
finally:
859+
if startup_state.cleanup_succeeded:
860+
self._clear_process_state()
861+
elif startup_state.cleanup_succeeded is False:
862+
_logger.warning("Failed to stop backend launcher process %d", startup_state.process.pid)
863+
859864
def _clear_process_state(self) -> None:
860865
"""Clear process state after the owned backend exits."""
861866
self._process = None

‎pyrit/common/apply_defaults.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -293,7 +293,7 @@ def wrapper(self: object, *args: object, **kwargs: object) -> T:
293293
f"Either pass the parameter explicitly or register a default using set_default_value()."
294294
)
295295
# If None was explicitly passed and parameter has REQUIRED_VALUE as default, also raise
296-
elif param_value is None:
296+
else:
297297
# Check if the parameter's default in the signature is REQUIRED_VALUE
298298
param_obj = sig.parameters.get(param_name)
299299
if param_obj and isinstance(param_obj.default, _RequiredValueSentinel):

‎pyrit/common/net_utility.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -53,7 +53,7 @@ def extract_url_parameters(url: str) -> dict[str, str]:
5353
parsed_url = urlparse(url)
5454
url_params = parse_qs(parsed_url.query, keep_blank_values=True)
5555
# Flatten params (parse_qs returns lists)
56-
return {k: v[0] if isinstance(v, list) and len(v) > 0 else "" for k, v in url_params.items()}
56+
return {k: v[0] if v else "" for k, v in url_params.items()}
5757

5858

5959
def remove_url_parameters(url: str) -> str:

‎pyrit/common/text_helper.py‎

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,17 @@
11
# Copyright (c) Microsoft Corporation.
22
# Licensed under the MIT license.
33

4-
from typing import IO, Any
4+
from typing import IO, Any, TypeGuard
5+
6+
7+
def is_non_empty_string(value: object) -> TypeGuard[str]:
8+
"""
9+
Check whether an untrusted value is a string containing non-whitespace text.
10+
11+
Returns:
12+
bool: Whether the value is a non-empty string.
13+
"""
14+
return isinstance(value, str) and bool(value.strip())
515

616

717
def read_txt(file: IO[Any]) -> list[dict[str, str]]:

0 commit comments

Comments
 (0)