From ba4f41bfa67a6f405e238b533f66a33f8f203ae0 Mon Sep 17 00:00:00 2001 From: Adrian Gavrila Date: Fri, 18 Sep 2026 14:28:17 -0400 Subject: [PATCH 1/7] Add GitHubCopilotTarget for single-turn local Copilot exchanges Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyproject.toml | 5 + pyrit/prompt_target/__init__.py | 2 + pyrit/prompt_target/github_copilot_target.py | 64 +++++++ .../target/test_github_copilot_target.py | 166 ++++++++++++++++++ uv.lock | 21 ++- 5 files changed, 257 insertions(+), 1 deletion(-) create mode 100644 pyrit/prompt_target/github_copilot_target.py create mode 100644 tests/unit/prompt_target/target/test_github_copilot_target.py diff --git a/pyproject.toml b/pyproject.toml index 9b618130f2..df0856cbf2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -133,12 +133,17 @@ litellm = [ "litellm>=1.84.0", ] +github-copilot = [ + "github-copilot-sdk>=1.0.11", +] + # all includes all functional dependencies excluding the ones from the "dev" dependency group all = [ "accelerate>=1.7.0", "azure-ai-ml>=1.32.0", "azure-cognitiveservices-speech>=1.44.0", "flask>=3.1.3", + "github-copilot-sdk>=1.0.11", "ipykernel>=6.29.5", "jupyter>=1.1.1", "litellm>=1.84.0", diff --git a/pyrit/prompt_target/__init__.py b/pyrit/prompt_target/__init__.py index 999c50834a..664c5185a9 100644 --- a/pyrit/prompt_target/__init__.py +++ b/pyrit/prompt_target/__init__.py @@ -31,6 +31,7 @@ from pyrit.prompt_target.common.target_requirements import CHAT_TARGET_REQUIREMENTS, TargetRequirements from pyrit.prompt_target.common.utils import limit_requests_per_minute from pyrit.prompt_target.gandalf_target import GandalfLevel, GandalfTarget + from pyrit.prompt_target.github_copilot_target import GitHubCopilotTarget from pyrit.prompt_target.http_target.http_target import HTTPTarget from pyrit.prompt_target.http_target.http_target_callback_functions import ( get_http_target_json_response_callback_function, @@ -66,6 +67,7 @@ "ConversationNormalizationPipeline": "pyrit.prompt_target.common.conversation_normalization_pipeline", "GandalfLevel": "pyrit.prompt_target.gandalf_target", "GandalfTarget": "pyrit.prompt_target.gandalf_target", + "GitHubCopilotTarget": "pyrit.prompt_target.github_copilot_target", "get_http_target_json_response_callback_function": "pyrit.prompt_target.http_target.http_target_callback_functions", "get_http_target_regex_matching_callback_function": ( "pyrit.prompt_target.http_target.http_target_callback_functions" diff --git a/pyrit/prompt_target/github_copilot_target.py b/pyrit/prompt_target/github_copilot_target.py new file mode 100644 index 0000000000..ad4c063b9a --- /dev/null +++ b/pyrit/prompt_target/github_copilot_target.py @@ -0,0 +1,64 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import logging + +from pyrit.models import Message, construct_response_from_request +from pyrit.prompt_target.common.prompt_target import PromptTarget + +logger = logging.getLogger(__name__) + + +class GitHubCopilotTarget(PromptTarget): + """Send single-turn text requests through the GitHub Copilot SDK.""" + + def __init__(self, *, model_name: str, retain_session: bool = False) -> None: + """ + Initialize the target with normal SDK login discovery. + + Args: + model_name (str): Explicit Copilot model ID. + retain_session (bool): Keep each Copilot session on disk instead of deleting it. Retained + session IDs are logged so they can be found later. Defaults to False. + """ + import copilot + + super().__init__(model_name=model_name) + self._sdk = copilot + self._retain_session = retain_session + + async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Message]) -> list[Message]: + request = normalized_conversation[-1].get_piece() + reply_text = await self._send_text_async(request.converted_value) + return [construct_response_from_request(request=request, response_text_pieces=[reply_text])] + + async def _send_text_async(self, prompt: str) -> str: + from copilot.generated.session_events import AssistantMessageData + + async with self._sdk.CopilotClient() as client: + # Keep the exchange local and free of ambient context so replies reflect the model, not the host machine. + session = await client.create_session( + model=self._model_name, + remote_session=self._sdk.RemoteSessionMode.OFF, + available_tools=[], + skip_custom_instructions=True, + instruction_directories=[], + enable_host_git_operations=False, + enable_config_discovery=False, + organization_custom_instructions="", + enable_on_demand_instruction_discovery=False, + ) + try: + reply = await session.send_and_wait(prompt) + if ( + reply is None + or not isinstance(reply.data, AssistantMessageData) + or not isinstance(reply.data.content, str) + ): + raise ValueError("Copilot did not return an assistant text reply.") + return reply.data.content + finally: + if self._retain_session: + logger.info("Retaining Copilot session %s as requested; delete it manually.", session.session_id) + else: + await client.delete_session(session.session_id) diff --git a/tests/unit/prompt_target/target/test_github_copilot_target.py b/tests/unit/prompt_target/target/test_github_copilot_target.py new file mode 100644 index 0000000000..053717ff7a --- /dev/null +++ b/tests/unit/prompt_target/target/test_github_copilot_target.py @@ -0,0 +1,166 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import logging +from datetime import UTC, datetime +from typing import TYPE_CHECKING, Any +from unittest.mock import AsyncMock, create_autospec, patch +from uuid import uuid4 + +import pytest + +from pyrit.models import Message, MessagePiece +from pyrit.prompt_normalizer import PromptNormalizer + +if TYPE_CHECKING: + from pyrit.memory import MemoryInterface + + +def _assistant_reply(text: str) -> Any: + """Build an SDK assistant-message event carrying ``text``.""" + from copilot.generated.session_events import AssistantMessageData, SessionEvent, SessionEventType + + return SessionEvent( + id=uuid4(), + timestamp=datetime.now(UTC), + type=SessionEventType.ASSISTANT_MESSAGE, + data=AssistantMessageData(content=text, message_id="sdk-reply"), + ) + + +def _mock_copilot_client(sdk: Any, *, reply: Any = None, error: Exception | None = None) -> Any: + """Build an autospec Copilot client whose session returns ``reply`` or raises ``error``. + + The target uses the client as an async context manager, so ``__aenter__`` must yield the + same mock the assertions inspect. ``session_id`` is set explicitly because the real + attribute is assigned in ``__init__`` and therefore absent from the class spec. + """ + session = create_autospec(sdk.CopilotSession, instance=True) + session.session_id = "sdk-session-id" + session.send_and_wait = AsyncMock(return_value=reply, side_effect=error) + client = create_autospec(sdk.CopilotClient, instance=True) + client.__aenter__.return_value = client + client.create_session.return_value = session + return client + + +@pytest.mark.usefixtures("patch_central_database") +async def test_normalizer_returns_and_stores_copilot_reply_async(sqlite_instance: MemoryInterface) -> None: + from pyrit.prompt_target import GitHubCopilotTarget + + sdk = pytest.importorskip("copilot") + + client = _mock_copilot_client(sdk, reply=_assistant_reply("HELLO")) + conversation_id = str(uuid4()) + + with patch.object(sdk, "CopilotClient", return_value=client): + target = GitHubCopilotTarget(model_name="gpt-5-mini") + response = await PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + conversation_id=conversation_id, + target=target, + ) + + assert isinstance(response, Message) + piece = response.get_piece() + assert (piece.role, piece.converted_value, piece.conversation_id, piece.response_error) == ( + "assistant", + "HELLO", + conversation_id, + "none", + ) + stored_messages = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) + assert [(message.get_piece().role, message.get_piece().converted_value) for message in stored_messages] == [ + ("user", "Reply exactly HELLO."), + ("assistant", "HELLO"), + ] + + +@pytest.mark.usefixtures("patch_central_database") +async def test_create_session_restricts_copilot_runtime_async() -> None: + from pyrit.prompt_target import GitHubCopilotTarget + + sdk = pytest.importorskip("copilot") + from copilot.generated.rpc import RemoteSessionMode + + restricted_configuration = { + "remote_session": RemoteSessionMode.OFF, + "available_tools": [], + "skip_custom_instructions": True, + "instruction_directories": [], + "enable_host_git_operations": False, + "enable_config_discovery": False, + "organization_custom_instructions": "", + "enable_on_demand_instruction_discovery": False, + } + client = _mock_copilot_client(sdk, reply=_assistant_reply("HELLO")) + + with patch.object(sdk, "CopilotClient", return_value=client): + await PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + conversation_id=str(uuid4()), + target=GitHubCopilotTarget(model_name="gpt-5-mini"), + ) + + sent_configuration = client.create_session.call_args.kwargs + assert {name: sent_configuration.get(name) for name in restricted_configuration} == restricted_configuration + + +@pytest.mark.usefixtures("patch_central_database") +async def test_owned_session_is_deleted_after_successful_send_async() -> None: + from pyrit.prompt_target import GitHubCopilotTarget + + sdk = pytest.importorskip("copilot") + + client = _mock_copilot_client(sdk, reply=_assistant_reply("HELLO")) + + with patch.object(sdk, "CopilotClient", return_value=client): + await PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + conversation_id=str(uuid4()), + target=GitHubCopilotTarget(model_name="gpt-5-mini"), + ) + + client.delete_session.assert_awaited_once_with("sdk-session-id") + client.__aexit__.assert_awaited_once() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_owned_session_is_deleted_when_send_fails_async() -> None: + from pyrit.prompt_target import GitHubCopilotTarget + + sdk = pytest.importorskip("copilot") + + client = _mock_copilot_client(sdk, error=RuntimeError("copilot runtime exploded")) + target = GitHubCopilotTarget(model_name="gpt-5-mini") + request = MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message() + + with patch.object(sdk, "CopilotClient", return_value=client): + with pytest.raises(RuntimeError, match="copilot runtime exploded"): + await target._send_prompt_to_target_async(normalized_conversation=[request]) + + client.delete_session.assert_awaited_once_with("sdk-session-id") + client.__aexit__.assert_awaited_once() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_retained_session_is_kept_and_reported_async(caplog: pytest.LogCaptureFixture) -> None: + from pyrit.prompt_target import GitHubCopilotTarget + + sdk = pytest.importorskip("copilot") + + client = _mock_copilot_client(sdk, reply=_assistant_reply("HELLO")) + target = GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=True) + + with patch.object(sdk, "CopilotClient", return_value=client), caplog.at_level(logging.INFO): + await PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + conversation_id=str(uuid4()), + target=target, + ) + + client.delete_session.assert_not_awaited() + client.__aexit__.assert_awaited_once() + assert "sdk-session-id" in caplog.text diff --git a/uv.lock b/uv.lock index e781166b28..6bb1cc2714 100644 --- a/uv.lock +++ b/uv.lock @@ -1805,6 +1805,19 @@ http = [ { name = "aiohttp" }, ] +[[package]] +name = "github-copilot-sdk" +version = "1.0.14" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "httpx" }, + { name = "pydantic" }, + { name = "python-dateutil" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/ed/e2/4d8c622545f55f06773548cf86ffe60f59e6a4cedcc412ff84fcac244b2b/github_copilot_sdk-1.0.14-py3-none-any.whl", hash = "sha256:9f4538cec7295f24507c650d794307396458a51f280a0fbbcd840da177109050", size = 614607, upload-time = "2026-09-16T05:11:14.572Z" }, +] + [[package]] name = "greenlet" version = "3.3.0" @@ -4791,6 +4804,7 @@ all = [ { name = "azure-ai-ml" }, { name = "azure-cognitiveservices-speech" }, { name = "flask" }, + { name = "github-copilot-sdk" }, { name = "ipykernel" }, { name = "jupyter" }, { name = "litellm" }, @@ -4812,6 +4826,9 @@ gcg = [ { name = "sentencepiece" }, { name = "torch" }, ] +github-copilot = [ + { name = "github-copilot-sdk" }, +] huggingface = [ { name = "sentencepiece" }, { name = "torch" }, @@ -4884,6 +4901,8 @@ requires-dist = [ { name = "fastapi", specifier = ">=0.133.0" }, { name = "flask", marker = "extra == 'all'", specifier = ">=3.1.3" }, { name = "flask", marker = "extra == 'playwright'", specifier = ">=3.1.3" }, + { name = "github-copilot-sdk", marker = "extra == 'all'", specifier = ">=1.0.11" }, + { name = "github-copilot-sdk", marker = "extra == 'github-copilot'", specifier = ">=1.0.11" }, { name = "httpx", extras = ["http2"], specifier = ">=0.27.2" }, { name = "ipykernel", marker = "extra == 'all'", specifier = ">=6.29.5" }, { name = "jinja2", specifier = ">=3.1.6" }, @@ -4931,7 +4950,7 @@ requires-dist = [ { name = "uvicorn", extras = ["standard"], specifier = ">=0.32.0" }, { name = "websockets", specifier = ">=14.0" }, ] -provides-extras = ["huggingface", "gcg", "playwright", "fairness-bias", "opencv", "speech", "litellm", "all"] +provides-extras = ["huggingface", "gcg", "playwright", "fairness-bias", "opencv", "speech", "litellm", "github-copilot", "all"] [package.metadata.requires-dev] dev = [ From 0b955248a6ab24456e56eecfeec3a5b7667d804d Mon Sep 17 00:00:00 2001 From: Adrian Gavrila Date: Tue, 22 Sep 2026 12:05:30 -0400 Subject: [PATCH 2/7] Complete GitHub Copilot text target Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/targets/github_copilot_target.md | 80 +++ doc/myst.yml | 1 + pyrit/prompt_target/github_copilot_target.py | 183 ++++- .../target/test_github_copilot_target.py | 679 ++++++++++++++++-- 4 files changed, 845 insertions(+), 98 deletions(-) create mode 100644 doc/code/targets/github_copilot_target.md diff --git a/doc/code/targets/github_copilot_target.md b/doc/code/targets/github_copilot_target.md new file mode 100644 index 0000000000..09d667e5b9 --- /dev/null +++ b/doc/code/targets/github_copilot_target.md @@ -0,0 +1,80 @@ +# GitHub Copilot SDK target + +`GitHubCopilotTarget` sends a fresh single-turn text exchange through the GitHub Copilot SDK. +Use Python 3.11 through 3.14 and install the optional extra: + +```console +python -m pip install "pyrit[github-copilot]" +``` + +Authenticate beforehand with an eligible Copilot login, or supply an SDK environment token. +The SDK checks `COPILOT_GITHUB_TOKEN`, then `GH_TOKEN`, then `GITHUB_TOKEN`; environment tokens +take precedence over saved login. Alternatively, pass a securely supplied token as +`github_token=token`, which takes precedence over discovery. Do not hard-code credentials. +The target does not initiate interactive sign-in. See [PyRIT setup](../setup/0_setup.md) +for framework initialization. + +## One text exchange + +Run this standalone script from your chosen existing local directory; no Git checkout is needed. +`Path.cwd()` explicitly selects that directory, and its resolved path is included in the target +identity. Running the script makes a real Copilot model request. + +```python +import asyncio +import logging +from pathlib import Path + +from pyrit.models import MessagePiece +from pyrit.prompt_normalizer import PromptNormalizer +from pyrit.prompt_target import GitHubCopilotTarget +from pyrit.setup import IN_MEMORY, initialize_pyrit_async + + +async def main_async() -> None: + await initialize_pyrit_async(memory_db_type=IN_MEMORY, silent=True) + + copilot_logger = logging.getLogger("pyrit.prompt_target.github_copilot_target") + copilot_logger.setLevel(logging.INFO) + copilot_logger.addHandler(logging.StreamHandler()) + copilot_logger.propagate = False + + target = GitHubCopilotTarget( + model_name="gpt-5-mini", + working_directory=Path.cwd(), + ) + response = await PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + target=target, + ) + print(response.get_piece().converted_value) + + +if __name__ == "__main__": + asyncio.run(main_async()) +``` + +The normalizer owns conversion and request/response persistence. Here, [IN_MEMORY](../memory/0_memory.md) +keeps PyRIT records only for the current process, independently of Copilot's local session data. + +## Retention and boundaries + +By default, each exchange releases owned SDK resources and deletes its owned local SDK session data. +Set `retain_session=True` to keep that data for diagnostics; client cleanup still runs. +Cleanup errors propagate, and crashes can leave residual data. Remote session export is explicitly +`OFF` in either mode: local retention does not enable export, and `OFF` does not mean offline inference. + +Capture the target's INFO logs for creation-attempt records linking PyRIT conversation IDs to requested +SDK session IDs, with SDK/runtime versions, protocol, retention and remote mode. These per-exchange +details are not automatically added to response metadata or target-identity exports. An attempt record +does not prove session allocation or a successful exchange; retained-session logs identify data kept +for diagnostics. + +This target supports fresh single-turn text only, not a public caller system prompt or native +continuation. It requests text-only runtime restrictions and retains restricted SDK base instructions; +other SDK/runtime context may remain. These controls are not an OS sandbox. + +`response_timeout_seconds` bounds dispatch and completion, excluding client/session startup, the +startup status lookup, PyRIT pacing and cleanup. A failed status lookup stops the exchange before +session creation or prompt dispatch. Errors, timeouts and cancellation propagate without automatic +replay or a rollback guarantee. diff --git a/doc/myst.yml b/doc/myst.yml index 78bc9147eb..90c1d3b22b 100644 --- a/doc/myst.yml +++ b/doc/myst.yml @@ -137,6 +137,7 @@ project: - file: code/targets/use_huggingface_chat_target.ipynb - file: code/targets/websocket_target.ipynb - file: code/targets/round_robin_target.ipynb + - file: code/targets/github_copilot_target.md - file: code/converters/0_converters.ipynb children: - file: code/converters/1_text_to_text_converters.ipynb diff --git a/pyrit/prompt_target/github_copilot_target.py b/pyrit/prompt_target/github_copilot_target.py index ad4c063b9a..acdbb38700 100644 --- a/pyrit/prompt_target/github_copilot_target.py +++ b/pyrit/prompt_target/github_copilot_target.py @@ -1,64 +1,199 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +import asyncio import logging +import math +from pathlib import Path +from uuid import uuid4 -from pyrit.models import Message, construct_response_from_request +from pyrit.models import ComponentIdentifier, Message, construct_response_from_request from pyrit.prompt_target.common.prompt_target import PromptTarget +from pyrit.prompt_target.common.utils import limit_requests_per_minute logger = logging.getLogger(__name__) class GitHubCopilotTarget(PromptTarget): - """Send single-turn text requests through the GitHub Copilot SDK.""" + """ + Send single-turn text requests through the GitHub Copilot SDK. - def __init__(self, *, model_name: str, retain_session: bool = False) -> None: + Capture INFO logs for session mapping and SDK/runtime version diagnostics. + """ + + def __init__( + self, + *, + model_name: str, + github_token: str | None = None, + working_directory: str | Path | None = None, + retain_session: bool = False, + response_timeout_seconds: float = 60.0, + max_requests_per_minute: int | None = None, + ) -> None: """ - Initialize the target with normal SDK login discovery. + Initialize the target with an explicit token or normal SDK login discovery. Args: model_name (str): Explicit Copilot model ID. + github_token (str | None): Nonblank GitHub token, forwarded unchanged with precedence over other + authentication methods. Defaults to None for SDK environment-token and saved-login discovery. + working_directory (str | Path | None): Existing local directory; no Git repository is required. + Supplied paths are resolved once against the current directory at construction and included + in saved target identifiers. Defaults to None for the SDK's current directory at each client start. retain_session (bool): Keep each Copilot session on disk instead of deleting it. Retained session IDs are logged so they can be found later. Defaults to False. + response_timeout_seconds (float): Shared time budget for dispatch and completion, in seconds. + Excludes client/session creation, startup status lookup, and cleanup. Defaults to 60. + max_requests_per_minute (int | None): PyRIT per-send pacing. Positive values delay each send by + 60 / value seconds before SDK client creation, outside the response deadline. + None or nonpositive values disable pacing. Defaults to None. + + Raises: + ValueError: If model_name or a supplied github_token is blank, or response_timeout_seconds + is not finite and positive, or a supplied working_directory is blank, missing, or not a directory. + OSError: If the working directory cannot be resolved or inspected. + RuntimeError: If the optional GitHub Copilot SDK is not installed. """ - import copilot + if not model_name.strip(): + raise ValueError("model_name must not be empty.") + if github_token is not None and not github_token.strip(): + raise ValueError("github_token must not be blank when supplied.") + if not math.isfinite(response_timeout_seconds) or response_timeout_seconds <= 0: + raise ValueError("response_timeout_seconds must be a finite positive number.") + + self._working_directory: str | None = None + if working_directory is not None: + if isinstance(working_directory, str) and not working_directory.strip(): + raise ValueError("working_directory must not be blank when supplied.") + resolved_directory = Path(working_directory).resolve() + if not resolved_directory.is_dir(): + raise ValueError("working_directory must be an existing directory.") + self._working_directory = str(resolved_directory) + + try: + import copilot + except ModuleNotFoundError as e: + raise RuntimeError("Could not import copilot. Install it with 'pip install pyrit[github-copilot]'.") from e - super().__init__(model_name=model_name) + super().__init__(model_name=model_name, max_requests_per_minute=max_requests_per_minute) self._sdk = copilot + self._github_token = github_token self._retain_session = retain_session + self._response_timeout_seconds = response_timeout_seconds + def _build_identifier(self) -> ComponentIdentifier: + """ + Build the identifier with the selected working directory. + + Returns: + ComponentIdentifier: The target identifier, including the resolved selected path when provided. + """ + return self._create_identifier(params={"working_directory": self._working_directory}) + + @limit_requests_per_minute async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Message]) -> list[Message]: request = normalized_conversation[-1].get_piece() - reply_text = await self._send_text_async(request.converted_value) + reply_text = await self._send_text_async( + prompt=request.converted_value, conversation_id=request.conversation_id + ) return [construct_response_from_request(request=request, response_text_pieces=[reply_text])] - async def _send_text_async(self, prompt: str) -> str: - from copilot.generated.session_events import AssistantMessageData - - async with self._sdk.CopilotClient() as client: - # Keep the exchange local and free of ambient context so replies reflect the model, not the host machine. - session = await client.create_session( - model=self._model_name, - remote_session=self._sdk.RemoteSessionMode.OFF, - available_tools=[], - skip_custom_instructions=True, - instruction_directories=[], - enable_host_git_operations=False, - enable_config_discovery=False, - organization_custom_instructions="", - enable_on_demand_instruction_discovery=False, + async def _send_text_async(self, *, prompt: str, conversation_id: str | None) -> str: + from copilot.generated.session_events import ( + AbortData, + AssistantMessageData, + SessionEvent, + SessionIdleData, + ToolExecutionStartData, + ) + + aborted = False + tool_execution_started = False + + def _record_turn_events(event: SessionEvent) -> None: + nonlocal aborted, tool_execution_started + if isinstance(event.data, AbortData) or (isinstance(event.data, SessionIdleData) and event.data.aborted): + aborted = True + if isinstance(event.data, ToolExecutionStartData): + tool_execution_started = True + + client = await asyncio.to_thread( + self._sdk.CopilotClient, github_token=self._github_token, working_directory=self._working_directory + ) + try: + await client.start() + status = await client.get_status() + session_id = str(uuid4()) + logger.info( + "Attempting Copilot session creation: pyrit_conversation_id=%s requested_sdk_session_id=%s " + "sdk_version=%s runtime_version=%s runtime_protocol_version=%s retain_session=%s remote_mode=OFF", + conversation_id, + session_id, + self._sdk.__version__, + status.version, + status.protocol_version, + self._retain_session, ) try: - reply = await session.send_and_wait(prompt) + # SDK runtime instructions are retained; other SDK-provided context may remain. + session = await client.create_session( + session_id=session_id, + model=self._model_name, + system_message={ + "mode": "customize", + "sections": { + "environment_context": {"action": "remove"}, + "custom_instructions": {"action": "remove"}, + }, + }, + remote_session=self._sdk.RemoteSessionMode.OFF, + available_tools=[], + skip_custom_instructions=True, + instruction_directories=[], + enable_host_git_operations=False, + enable_config_discovery=False, + organization_custom_instructions="", + enable_on_demand_instruction_discovery=False, + infinite_sessions={"enabled": False}, + memory={"enabled": False}, + enable_session_store=False, + enable_file_hooks=False, + ) + except (Exception, asyncio.CancelledError): + if await client.get_session_metadata(session_id) is not None: + if self._retain_session: + logger.info("Retaining Copilot session %s as requested; delete it manually.", session_id) + else: + await client.delete_session(session_id) + raise + try: + unsubscribe = session.on(_record_turn_events) + try: + reply = await asyncio.wait_for( + session.send_and_wait(prompt, timeout=self._response_timeout_seconds), + timeout=self._response_timeout_seconds, + ) + finally: + unsubscribe() + + if aborted: + raise RuntimeError("Copilot turn was aborted.") + if tool_execution_started: + raise RuntimeError("Copilot turn reported tool execution.") if ( reply is None + or reply.agent_id is not None or not isinstance(reply.data, AssistantMessageData) or not isinstance(reply.data.content, str) + or not reply.data.content ): - raise ValueError("Copilot did not return an assistant text reply.") + raise ValueError("Copilot did not return a non-empty root assistant text reply.") return reply.data.content finally: if self._retain_session: logger.info("Retaining Copilot session %s as requested; delete it manually.", session.session_id) else: await client.delete_session(session.session_id) + finally: + await client.stop() diff --git a/tests/unit/prompt_target/target/test_github_copilot_target.py b/tests/unit/prompt_target/target/test_github_copilot_target.py index 053717ff7a..49fd2f10e1 100644 --- a/tests/unit/prompt_target/target/test_github_copilot_target.py +++ b/tests/unit/prompt_target/target/test_github_copilot_target.py @@ -3,23 +3,54 @@ from __future__ import annotations +import asyncio import logging +import threading +from dataclasses import replace from datetime import UTC, datetime +from functools import partial from typing import TYPE_CHECKING, Any -from unittest.mock import AsyncMock, create_autospec, patch -from uuid import uuid4 +from unittest.mock import AsyncMock, NonCallableMagicMock, create_autospec, patch +from uuid import UUID, uuid4 import pytest from pyrit.models import Message, MessagePiece from pyrit.prompt_normalizer import PromptNormalizer +from pyrit.prompt_target import GitHubCopilotTarget if TYPE_CHECKING: + from collections.abc import Callable, Iterator + from pathlib import Path + + from copilot.generated.session_events import SessionEvent + from pyrit.memory import MemoryInterface +TARGET_LOGGER = "pyrit.prompt_target.github_copilot_target" + + +@pytest.fixture +def sdk() -> Any: + return pytest.importorskip("copilot") + + +@pytest.fixture +def client(sdk: Any) -> Iterator[NonCallableMagicMock]: + from copilot.client import GetStatusResponse + + session = create_autospec(sdk.CopilotSession, instance=True) + session.session_id = "sdk-session-id" + session.send_and_wait.return_value = _assistant_reply("HELLO") + client = create_autospec(sdk.CopilotClient, instance=True) + assert isinstance(client, NonCallableMagicMock) + client.create_session.return_value = session + client.get_status.return_value = GetStatusResponse(version="6.5.4", protocol_version=3) + with patch.object(sdk, "CopilotClient", return_value=client): + yield client -def _assistant_reply(text: str) -> Any: - """Build an SDK assistant-message event carrying ``text``.""" + +def _assistant_reply(text: str) -> SessionEvent: from copilot.generated.session_events import AssistantMessageData, SessionEvent, SessionEventType return SessionEvent( @@ -30,37 +61,42 @@ def _assistant_reply(text: str) -> Any: ) -def _mock_copilot_client(sdk: Any, *, reply: Any = None, error: Exception | None = None) -> Any: - """Build an autospec Copilot client whose session returns ``reply`` or raises ``error``. - - The target uses the client as an async context manager, so ``__aenter__`` must yield the - same mock the assertions inspect. ``session_id`` is set explicitly because the real - attribute is assigned in ``__init__`` and therefore absent from the class spec. - """ - session = create_autospec(sdk.CopilotSession, instance=True) - session.session_id = "sdk-session-id" - session.send_and_wait = AsyncMock(return_value=reply, side_effect=error) - client = create_autospec(sdk.CopilotClient, instance=True) - client.__aenter__.return_value = client - client.create_session.return_value = session - return client +def _mock_session_storage(*, client: NonCallableMagicMock, sessions: set[str]) -> None: + from copilot import SessionMetadata + async def get_session_metadata_async(session_id: str) -> SessionMetadata | None: + if session_id not in sessions: + return None + return SessionMetadata( + session_id=session_id, + start_time=datetime(2026, 1, 1, tzinfo=UTC), + modified_time=datetime(2026, 1, 1, tzinfo=UTC), + is_remote=False, + ) -@pytest.mark.usefixtures("patch_central_database") -async def test_normalizer_returns_and_stores_copilot_reply_async(sqlite_instance: MemoryInterface) -> None: - from pyrit.prompt_target import GitHubCopilotTarget + client.get_session_metadata.side_effect = get_session_metadata_async + client.delete_session.side_effect = sessions.remove - sdk = pytest.importorskip("copilot") - client = _mock_copilot_client(sdk, reply=_assistant_reply("HELLO")) +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("retain_session", [False, True], ids=["delete", "retain"]) +async def test_normalizer_round_trip_and_retention_async( + *, + sdk: Any, + client: NonCallableMagicMock, + sqlite_instance: MemoryInterface, + caplog: pytest.LogCaptureFixture, + retain_session: bool, +) -> None: conversation_id = str(uuid4()) - - with patch.object(sdk, "CopilotClient", return_value=client): - target = GitHubCopilotTarget(model_name="gpt-5-mini") + request = MessagePiece( + role="user", original_value="Original text before conversion.", converted_value="Reply exactly HELLO." + ).to_message() + with patch.object(sdk, "__version__", "9.8.7"), caplog.at_level(logging.INFO, logger=TARGET_LOGGER): response = await PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + message=request, conversation_id=conversation_id, - target=target, + target=GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=retain_session), ) assert isinstance(response, Message) @@ -71,21 +107,65 @@ async def test_normalizer_returns_and_stores_copilot_reply_async(sqlite_instance conversation_id, "none", ) - stored_messages = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) - assert [(message.get_piece().role, message.get_piece().converted_value) for message in stored_messages] == [ - ("user", "Reply exactly HELLO."), - ("assistant", "HELLO"), + stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) + assert [ + (m.get_piece().role, m.get_piece().original_value, m.get_piece().converted_value, m.get_piece().response_error) + for m in stored + ] == [ + ("user", "Original text before conversion.", "Reply exactly HELLO.", "none"), + ("assistant", "HELLO", "HELLO", "none"), ] + session = client.create_session.return_value + session.send_and_wait.assert_awaited_once_with("Reply exactly HELLO.", timeout=60.0) + session.on.return_value.assert_called_once_with() + client.start.assert_awaited_once() + client.get_status.assert_awaited_once() + client.create_session.assert_awaited_once() + client.get_session_metadata.assert_not_awaited() + client.stop.assert_awaited_once() + sdk.CopilotClient.assert_called_once_with(github_token=None, working_directory=None) + requested_id = client.create_session.await_args.kwargs["session_id"] + assert str(UUID(requested_id)) == requested_id + assert requested_id != "sdk-session-id" + records = [r.getMessage() for r in caplog.records if r.name == TARGET_LOGGER and r.levelno == logging.INFO] + assert len(records) == (2 if retain_session else 1) + assert records[0].startswith("Attempting Copilot session creation:") + for field in ( + f"pyrit_conversation_id={conversation_id}", + f"requested_sdk_session_id={requested_id}", + "sdk_version=9.8.7", + "runtime_version=6.5.4", + "runtime_protocol_version=3", + f"retain_session={retain_session}", + "remote_mode=OFF", + ): + assert field in records[0] + if retain_session: + client.delete_session.assert_not_awaited() + assert records[1] == "Retaining Copilot session sdk-session-id as requested; delete it manually." + else: + client.delete_session.assert_awaited_once_with("sdk-session-id") @pytest.mark.usefixtures("patch_central_database") -async def test_create_session_restricts_copilot_runtime_async() -> None: - from pyrit.prompt_target import GitHubCopilotTarget - - sdk = pytest.importorskip("copilot") +async def test_create_session_restricts_copilot_runtime_async(client: NonCallableMagicMock) -> None: from copilot.generated.rpc import RemoteSessionMode - restricted_configuration = { + await PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + target=GitHubCopilotTarget(model_name="gpt-5-mini"), + ) + configuration = dict(client.create_session.await_args.kwargs) + assert str(UUID(configuration.pop("session_id"))) + assert configuration == { + "model": "gpt-5-mini", + "system_message": { + "mode": "customize", + "sections": { + "environment_context": {"action": "remove"}, + "custom_instructions": {"action": "remove"}, + }, + }, "remote_session": RemoteSessionMode.OFF, "available_tools": [], "skip_custom_instructions": True, @@ -94,73 +174,524 @@ async def test_create_session_restricts_copilot_runtime_async() -> None: "enable_config_discovery": False, "organization_custom_instructions": "", "enable_on_demand_instruction_discovery": False, + "infinite_sessions": {"enabled": False}, + "memory": {"enabled": False}, + "enable_session_store": False, + "enable_file_hooks": False, } - client = _mock_copilot_client(sdk, reply=_assistant_reply("HELLO")) - with patch.object(sdk, "CopilotClient", return_value=client): + +@pytest.mark.usefixtures("patch_central_database", "sdk") +def test_target_advertises_single_turn_text_only_capabilities() -> None: + capabilities = GitHubCopilotTarget(model_name="gpt-4o").capabilities + assert capabilities.supports_multi_turn is False + assert capabilities.supports_system_prompt is False + assert capabilities.supports_multi_message_pieces is False + assert capabilities.input_modalities == frozenset({frozenset({"text"})}) + assert capabilities.output_modalities == frozenset({frozenset({"text"})}) + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("failure_stage", ["start", "status", "send", "stop"]) +async def test_normalizer_surfaces_lifecycle_failures_async( + *, + client: NonCallableMagicMock, + sqlite_instance: MemoryInterface, + failure_stage: str, +) -> None: + from copilot.client import StopError + + session = client.create_session.return_value + error = ( + ExceptionGroup("SDK shutdown failed", [StopError(message="Synthetic shutdown failure")]) + if failure_stage == "stop" + else RuntimeError(f"Synthetic {failure_stage} failure") + ) + operation = { + "start": client.start, + "status": client.get_status, + "send": session.send_and_wait, + "stop": client.stop, + }[failure_stage] + operation.side_effect = error + conversation_id = str(uuid4()) + with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as exc_info: await PromptNormalizer().send_prompt_async( message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), - conversation_id=str(uuid4()), + conversation_id=conversation_id, target=GitHubCopilotTarget(model_name="gpt-5-mini"), ) - sent_configuration = client.create_session.call_args.kwargs - assert {name: sent_configuration.get(name) for name in restricted_configuration} == restricted_configuration + assert exc_info.value.__cause__ is error + stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) + assert [(m.get_piece().role, m.get_piece().response_error) for m in stored] == [ + ("user", "none"), + ("assistant", "processing"), + ] + client.start.assert_awaited_once() + client.stop.assert_awaited_once() + client.get_session_metadata.assert_not_awaited() + if failure_stage == "start": + client.get_status.assert_not_awaited() + else: + client.get_status.assert_awaited_once() + if failure_stage in ("start", "status"): + client.create_session.assert_not_awaited() + session.send_and_wait.assert_not_awaited() + client.delete_session.assert_not_awaited() + else: + client.create_session.assert_awaited_once() + session.send_and_wait.assert_awaited_once() + client.delete_session.assert_awaited_once_with("sdk-session-id") + session.on.return_value.assert_called_once_with() + session.send.assert_not_awaited() @pytest.mark.usefixtures("patch_central_database") -async def test_owned_session_is_deleted_after_successful_send_async() -> None: - from pyrit.prompt_target import GitHubCopilotTarget - - sdk = pytest.importorskip("copilot") +async def test_normalizer_surfaces_dispatch_timeout_and_cleans_up_without_replay_async( + *, + sdk: Any, + client: NonCallableMagicMock, +) -> None: + session = client.create_session.return_value + + async def stall_send_async(*_args: Any, **_kwargs: Any) -> None: + await asyncio.Event().wait() + + session.send.side_effect = stall_send_async + session.send_and_wait.side_effect = partial(sdk.CopilotSession.send_and_wait, session) + with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as exc_info: + # A bare watchdog TimeoutError must not satisfy the normalizer-wrapped failure. + await asyncio.wait_for( + PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + target=GitHubCopilotTarget(model_name="gpt-5-mini", response_timeout_seconds=0.01), + ), + timeout=2.0, + ) + assert isinstance(exc_info.value.__cause__, TimeoutError) + session.send.assert_awaited_once() + client.delete_session.assert_awaited_once_with("sdk-session-id") + client.stop.assert_awaited_once() - client = _mock_copilot_client(sdk, reply=_assistant_reply("HELLO")) - with patch.object(sdk, "CopilotClient", return_value=client): +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("invalid_reply", ["empty", "subagent", "absent", "non-assistant"]) +async def test_normalizer_rejects_invalid_reply_async( + *, + client: NonCallableMagicMock, + sqlite_instance: MemoryInterface, + invalid_reply: str, +) -> None: + from copilot.generated.session_events import SessionEventType, SessionIdleData + + reply = _assistant_reply("HELLO") + session = client.create_session.return_value + session.send_and_wait.return_value = { + "empty": _assistant_reply(""), + "subagent": replace(reply, agent_id="sdk-subagent-id"), + "absent": None, + "non-assistant": replace(reply, type=SessionEventType.SESSION_IDLE, data=SessionIdleData(aborted=False)), + }[invalid_reply] + conversation_id = str(uuid4()) + with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as exc_info: await PromptNormalizer().send_prompt_async( message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), - conversation_id=str(uuid4()), + conversation_id=conversation_id, target=GitHubCopilotTarget(model_name="gpt-5-mini"), ) - + assert isinstance(exc_info.value.__cause__, ValueError) + stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) + assert [(m.get_piece().role, m.get_piece().response_error) for m in stored] == [ + ("user", "none"), + ("assistant", "processing"), + ] + session.send_and_wait.assert_awaited_once() + session.on.return_value.assert_called_once_with() client.delete_session.assert_awaited_once_with("sdk-session-id") - client.__aexit__.assert_awaited_once() + client.stop.assert_awaited_once() @pytest.mark.usefixtures("patch_central_database") -async def test_owned_session_is_deleted_when_send_fails_async() -> None: - from pyrit.prompt_target import GitHubCopilotTarget +@pytest.mark.parametrize( + ("event_stream", "expected_error"), + [("aborted-idle", "abort"), ("abort-then-idle", "abort"), ("tool-then-idle", "tool")], + ids=["aborted-idle", "abort-then-idle", "tool-then-idle"], +) +async def test_normalizer_rejects_unsafe_events_async( + *, + sdk: Any, + client: NonCallableMagicMock, + sqlite_instance: MemoryInterface, + event_stream: str, + expected_error: str, +) -> None: + from copilot.generated.session_events import AbortData, SessionEventType, SessionIdleData, ToolExecutionStartData + + session = client.create_session.return_value + reply = _assistant_reply("HELLO") + events = { + "aborted-idle": [ + reply, + replace(reply, id=uuid4(), type=SessionEventType.SESSION_IDLE, data=SessionIdleData(aborted=True)), + ], + "abort-then-idle": [ + reply, + replace( + reply, + id=uuid4(), + type=SessionEventType.ABORT, + data=AbortData.from_dict({"reason": "user_initiated"}), + ), + replace(reply, id=uuid4(), type=SessionEventType.SESSION_IDLE, data=SessionIdleData(aborted=None)), + ], + "tool-then-idle": [ + replace( + reply, + id=uuid4(), + type=SessionEventType.TOOL_EXECUTION_START, + data=ToolExecutionStartData( + tool_call_id="synthetic-tool-call", tool_name="benign_test_tool", arguments={"text": "HELLO"} + ), + ), + reply, + replace(reply, id=uuid4(), type=SessionEventType.SESSION_IDLE, data=SessionIdleData(aborted=False)), + ], + }[event_stream] + handlers: list[Callable[[SessionEvent], None]] = [] + + def subscribe(handler: Callable[[SessionEvent], None]) -> Callable[[], None]: + handlers.append(handler) + return partial(handlers.remove, handler) + + async def send_events_async(*_args: Any, **_kwargs: Any) -> str: + for event in events: + for handler in tuple(handlers): + handler(event) + return "sdk-request-id" + + session.on.side_effect = subscribe + session.send.side_effect = send_events_async + session.send_and_wait.side_effect = partial(sdk.CopilotSession.send_and_wait, session) + conversation_id = str(uuid4()) + with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as exc_info: + await PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + conversation_id=conversation_id, + target=GitHubCopilotTarget(model_name="gpt-5-mini", response_timeout_seconds=1.0), + ) + assert isinstance(exc_info.value.__cause__, RuntimeError) + assert expected_error in str(exc_info.value.__cause__).lower() + stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) + assert [(m.get_piece().role, m.get_piece().response_error) for m in stored] == [ + ("user", "none"), + ("assistant", "processing"), + ] + session.send.assert_awaited_once() + client.delete_session.assert_awaited_once_with("sdk-session-id") + client.stop.assert_awaited_once() + assert not handlers - sdk = pytest.importorskip("copilot") - client = _mock_copilot_client(sdk, error=RuntimeError("copilot runtime exploded")) - target = GitHubCopilotTarget(model_name="gpt-5-mini") - request = MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message() +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("retain_session", [False, True], ids=["delete", "retain"]) +async def test_normalizer_cleans_up_partial_creation_async( + *, + client: NonCallableMagicMock, + sqlite_instance: MemoryInterface, + caplog: pytest.LogCaptureFixture, + retain_session: bool, +) -> None: + session = client.create_session.return_value + sessions = {"unrelated-session-id"} + allocated_session_id = "" + creation_error = RuntimeError("Copilot post-create options update failed") + + async def create_session_async(*, session_id: str = "sdk-generated-session-id", **_kwargs: Any) -> None: + nonlocal allocated_session_id + allocated_session_id = session_id + sessions.add(session_id) + raise creation_error + + client.create_session.side_effect = create_session_async + _mock_session_storage(client=client, sessions=sessions) + conversation_id = str(uuid4()) + with caplog.at_level(logging.INFO, logger=TARGET_LOGGER): + with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as exc_info: + await PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + conversation_id=conversation_id, + target=GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=retain_session), + ) + + assert exc_info.value.__cause__ is creation_error + assert allocated_session_id == client.create_session.await_args.kwargs["session_id"] + assert str(UUID(allocated_session_id)) == allocated_session_id + assert allocated_session_id != "sdk-session-id" + client.create_session.assert_awaited_once() + client.get_session_metadata.assert_awaited_once_with(allocated_session_id) + session.send_and_wait.assert_not_awaited() + session.send.assert_not_awaited() + client.stop.assert_awaited_once() + stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) + assert [(m.get_piece().role, m.get_piece().response_error) for m in stored] == [ + ("user", "none"), + ("assistant", "processing"), + ] + retained_logs = [ + r.getMessage() + for r in caplog.records + if r.name == TARGET_LOGGER + and r.levelno == logging.INFO + and r.getMessage().startswith("Retaining Copilot session ") + ] + if retain_session: + client.delete_session.assert_not_awaited() + assert sessions == {"unrelated-session-id", allocated_session_id} + assert retained_logs == [f"Retaining Copilot session {allocated_session_id} as requested; delete it manually."] + else: + client.delete_session.assert_awaited_once_with(allocated_session_id) + assert sessions == {"unrelated-session-id"} + assert not retained_logs - with patch.object(sdk, "CopilotClient", return_value=client): - with pytest.raises(RuntimeError, match="copilot runtime exploded"): - await target._send_prompt_to_target_async(normalized_conversation=[request]) - client.delete_session.assert_awaited_once_with("sdk-session-id") - client.__aexit__.assert_awaited_once() +@pytest.mark.usefixtures("patch_central_database") +async def test_normalizer_deletes_owned_session_when_creation_is_cancelled_after_allocation_async( + client: NonCallableMagicMock, +) -> None: + session = client.create_session.return_value + sessions = {"unrelated-session-id"} + allocated = asyncio.Event() + + async def create_session_async(*, session_id: str = "sdk-generated-session-id", **_kwargs: Any) -> None: + sessions.add(session_id) + allocated.set() + await asyncio.Event().wait() + + client.create_session.side_effect = create_session_async + _mock_session_storage(client=client, sessions=sessions) + request_task = asyncio.create_task( + PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + target=GitHubCopilotTarget(model_name="gpt-5-mini"), + ) + ) + try: + await asyncio.wait_for(allocated.wait(), timeout=2.0) + request_task.cancel() + with pytest.raises(asyncio.CancelledError): + await request_task + finally: + if not request_task.done(): + request_task.cancel() + await asyncio.gather(request_task, return_exceptions=True) + + client.create_session.assert_awaited_once() + allocated_id = client.create_session.await_args.kwargs["session_id"] + client.get_session_metadata.assert_awaited_once_with(allocated_id) + client.delete_session.assert_awaited_once_with(allocated_id) + session.send_and_wait.assert_not_awaited() + session.send.assert_not_awaited() + client.stop.assert_awaited_once() + assert sessions == {"unrelated-session-id"} @pytest.mark.usefixtures("patch_central_database") -async def test_retained_session_is_kept_and_reported_async(caplog: pytest.LogCaptureFixture) -> None: - from pyrit.prompt_target import GitHubCopilotTarget +async def test_normalizer_stops_owned_client_when_startup_is_cancelled_async( + *, + sdk: Any, + client: NonCallableMagicMock, +) -> None: + session = client.create_session.return_value + owned_resources: set[str] = set() + allocated = asyncio.Event() + original_cancellation: asyncio.CancelledError | None = None + + async def start_async() -> None: + nonlocal original_cancellation + owned_resources.add("owned-runtime") + allocated.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError as error: + original_cancellation = error + raise + + async def stop_async() -> None: + owned_resources.remove("owned-runtime") + + client.start.side_effect = start_async + client.stop.side_effect = stop_async + request_task = asyncio.create_task( + PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + target=GitHubCopilotTarget(model_name="gpt-5-mini"), + ) + ) + try: + await asyncio.wait_for(allocated.wait(), timeout=2.0) + request_task.cancel() + with pytest.raises(asyncio.CancelledError) as exc_info: + await request_task + finally: + if not request_task.done(): + request_task.cancel() + await asyncio.gather(request_task, return_exceptions=True) + + assert exc_info.value is original_cancellation + sdk.CopilotClient.assert_called_once() + client.start.assert_awaited_once() + client.get_status.assert_not_awaited() + client.create_session.assert_not_awaited() + client.get_session_metadata.assert_not_awaited() + client.delete_session.assert_not_awaited() + session.send.assert_not_awaited() + session.send_and_wait.assert_not_awaited() + client.stop.assert_awaited_once() + assert owned_resources == set() - sdk = pytest.importorskip("copilot") - client = _mock_copilot_client(sdk, reply=_assistant_reply("HELLO")) - target = GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=True) +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + ("github_token", "use_working_directory", "max_requests_per_minute"), + [ + pytest.param(" dummy-github-token ", False, None, id="token-only"), + pytest.param(None, True, None, id="directory-only"), + pytest.param(None, False, 30, id="throttle-only"), + pytest.param(" dummy-github-token ", True, 30, id="all-options"), + ], +) +async def test_normalizer_forwards_options_without_exposing_token_async( + *, + sdk: Any, + client: NonCallableMagicMock, + caplog: pytest.LogCaptureFixture, + tmp_path: Path, + github_token: str | None, + use_working_directory: bool, + max_requests_per_minute: int | None, +) -> None: + with ( + caplog.at_level(logging.DEBUG, logger=TARGET_LOGGER), + patch.object(asyncio, "sleep", new_callable=AsyncMock) as mock_sleep, + ): + target = GitHubCopilotTarget( + model_name="gpt-5-mini", + github_token=github_token, + working_directory=tmp_path if use_working_directory else None, + max_requests_per_minute=max_requests_per_minute, + ) + response = await PromptNormalizer().send_prompt_async( + message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), target=target + ) + sdk.CopilotClient.assert_called_once_with( + github_token=github_token, working_directory=str(tmp_path) if use_working_directory else None + ) + assert response.get_piece().converted_value == "HELLO" + client.create_session.return_value.send_and_wait.assert_awaited_once() + if max_requests_per_minute is not None: + mock_sleep.assert_awaited_once_with(2.0) + assert target.get_identifier().params["max_requests_per_minute"] == 30 + else: + mock_sleep.assert_not_awaited() + if use_working_directory: + assert target.get_identifier().params["working_directory"] == str(tmp_path) + assert "dummy-github-token" not in target.get_identifier().model_dump_json() + assert "dummy-github-token" not in caplog.text + + +def test_init_without_copilot_sdk_reports_installation_guidance() -> None: + with patch.dict("sys.modules", {"copilot": None}): + with pytest.raises(RuntimeError, match=r"pip install pyrit\[github-copilot\]"): + GitHubCopilotTarget(model_name="gpt-5-mini") + + +@pytest.mark.parametrize( + ("overrides", "field"), + [ + pytest.param({"model_name": " "}, "model_name", id="blank-model"), + pytest.param({"github_token": " "}, "github_token", id="blank-token"), + pytest.param({"response_timeout_seconds": 0}, "response_timeout_seconds", id="zero-timeout"), + pytest.param({"response_timeout_seconds": -1}, "response_timeout_seconds", id="negative-timeout"), + pytest.param({"response_timeout_seconds": float("inf")}, "response_timeout_seconds", id="infinite-timeout"), + pytest.param({"response_timeout_seconds": float("nan")}, "response_timeout_seconds", id="nan-timeout"), + pytest.param({"working_directory": " "}, "working_directory", id="blank-directory"), + ], +) +def test_init_rejects_invalid_options_before_sdk_import(*, overrides: dict[str, Any], field: str) -> None: + with patch.dict("sys.modules", {"copilot": None}), pytest.raises(ValueError, match=field): + GitHubCopilotTarget(**{"model_name": "gpt-5-mini", **overrides}) + + +@pytest.mark.parametrize("path_kind", ["missing", "file"]) +def test_init_rejects_non_directory_before_sdk_import(*, tmp_path: Path, path_kind: str) -> None: + path = tmp_path / "not-a-directory" + if path_kind == "file": + path.write_text("local test fixture", encoding="utf-8") + with patch.dict("sys.modules", {"copilot": None}), pytest.raises(ValueError, match="working_directory"): + GitHubCopilotTarget(model_name="gpt-5-mini", working_directory=path) - with patch.object(sdk, "CopilotClient", return_value=client), caplog.at_level(logging.INFO): - await PromptNormalizer().send_prompt_async( + +@pytest.mark.usefixtures("patch_central_database") +async def test_normalizer_keeps_event_loop_responsive_during_client_construction_async( + *, sdk: Any, client: NonCallableMagicMock, sqlite_instance: MemoryInterface +) -> None: + loop = asyncio.get_running_loop() + constructor_entered = asyncio.Event() + constructor_finished = asyncio.Event() + release = threading.Event() + released_while_constructing = False + + def construct_client(*, github_token: str | None, working_directory: str | None) -> NonCallableMagicMock: + nonlocal released_while_constructing + try: + loop.call_soon_threadsafe(constructor_entered.set) + # The bound lets a blocked event loop escape without satisfying the responsiveness assertion. + released_while_constructing = release.wait(timeout=5.0) + return client + finally: + loop.call_soon_threadsafe(constructor_finished.set) + + async def release_constructor_async() -> None: + await constructor_entered.wait() + release.set() + + sdk.CopilotClient.side_effect = construct_client + conversation_id = str(uuid4()) + request_task = asyncio.create_task( + PromptNormalizer().send_prompt_async( message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), - conversation_id=str(uuid4()), - target=target, + conversation_id=conversation_id, + target=GitHubCopilotTarget(model_name="gpt-5-mini"), ) + ) + release_task = asyncio.create_task(release_constructor_async()) + try: + response, _ = await asyncio.wait_for(asyncio.gather(request_task, release_task), timeout=10.0) + finally: + release.set() + for task in (request_task, release_task): + if not task.done(): + task.cancel() + await asyncio.gather(request_task, release_task, return_exceptions=True) + await asyncio.wait_for(constructor_finished.wait(), timeout=5.0) - client.delete_session.assert_not_awaited() - client.__aexit__.assert_awaited_once() - assert "sdk-session-id" in caplog.text + assert isinstance(response, Message) + piece = response.get_piece() + assert (piece.role, piece.converted_value, piece.conversation_id, piece.response_error) == ( + "assistant", + "HELLO", + conversation_id, + "none", + ) + stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) + assert [(m.get_piece().role, m.get_piece().converted_value, m.get_piece().response_error) for m in stored] == [ + ("user", "Reply exactly HELLO.", "none"), + ("assistant", "HELLO", "none"), + ] + sdk.CopilotClient.assert_called_once_with(github_token=None, working_directory=None) + client.start.assert_awaited_once() + client.create_session.return_value.send_and_wait.assert_awaited_once_with("Reply exactly HELLO.", timeout=60.0) + client.stop.assert_awaited_once() + client.delete_session.assert_awaited_once_with("sdk-session-id") + assert released_while_constructing is True From fcc1c376693d4a31002d8eb9b898ec34a5a389e2 Mon Sep 17 00:00:00 2001 From: Adrian Gavrila Date: Wed, 23 Sep 2026 13:42:03 -0400 Subject: [PATCH 3/7] Support native GitHub Copilot conversations Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/targets/github_copilot_target.md | 100 +- pyrit/prompt_target/common/prompt_target.py | 9 +- pyrit/prompt_target/github_copilot_target.py | 390 ++++++-- .../target/test_github_copilot_target.py | 919 ++++++++++++++++-- .../target/test_prompt_target.py | 112 ++- 5 files changed, 1318 insertions(+), 212 deletions(-) diff --git a/doc/code/targets/github_copilot_target.md b/doc/code/targets/github_copilot_target.md index 09d667e5b9..ca4f43aaa7 100644 --- a/doc/code/targets/github_copilot_target.md +++ b/doc/code/targets/github_copilot_target.md @@ -1,6 +1,7 @@ # GitHub Copilot SDK target -`GitHubCopilotTarget` sends a fresh single-turn text exchange through the GitHub Copilot SDK. +`GitHubCopilotTarget` sends text conversations through the GitHub Copilot SDK, retaining native +conversation state across turns. Use Python 3.11 through 3.14 and install the optional extra: ```console @@ -14,18 +15,19 @@ take precedence over saved login. Alternatively, pass a securely supplied token The target does not initiate interactive sign-in. See [PyRIT setup](../setup/0_setup.md) for framework initialization. -## One text exchange +## Two-turn native conversation Run this standalone script from your chosen existing local directory; no Git checkout is needed. `Path.cwd()` explicitly selects that directory, and its resolved path is included in the target -identity. Running the script makes a real Copilot model request. +identity. Running the script makes **two real Copilot model requests**. ```python import asyncio import logging from pathlib import Path +from uuid import uuid4 -from pyrit.models import MessagePiece +from pyrit.models import Message from pyrit.prompt_normalizer import PromptNormalizer from pyrit.prompt_target import GitHubCopilotTarget from pyrit.setup import IN_MEMORY, initialize_pyrit_async @@ -43,11 +45,23 @@ async def main_async() -> None: model_name="gpt-5-mini", working_directory=Path.cwd(), ) - response = await PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), - target=target, - ) - print(response.get_piece().converted_value) + normalizer = PromptNormalizer() + try: + conversation_id = str(uuid4()) + target.set_system_prompt(system_prompt="Answer concisely.", conversation_id=conversation_id) + + for prompt in ( + "Remember the codeword ORCHID for this conversation.", + "What codeword did I ask you to remember?", + ): + response = await normalizer.send_prompt_async( + message=Message.from_prompt(prompt=prompt, role="user"), + conversation_id=conversation_id, + target=target, + ) + print(response.get_piece().converted_value) + finally: + await target.cleanup_target_async() if __name__ == "__main__": @@ -56,25 +70,59 @@ if __name__ == "__main__": The normalizer owns conversion and request/response persistence. Here, [IN_MEMORY](../memory/0_memory.md) keeps PyRIT records only for the current process, independently of Copilot's local session data. +The target lazily shares one client across isolated native conversations, with one native session per +PyRIT conversation; this shares SDK runtime/authentication, not process or security isolation. Keep +the target alive for all sends and join workflow tasks before calling `cleanup_target_async()`. ## Retention and boundaries -By default, each exchange releases owned SDK resources and deletes its owned local SDK session data. -Set `retain_session=True` to keep that data for diagnostics; client cleanup still runs. -Cleanup errors propagate, and crashes can leave residual data. Remote session export is explicitly -`OFF` in either mode: local retention does not enable export, and `OFF` does not mean offline inference. - -Capture the target's INFO logs for creation-attempt records linking PyRIT conversation IDs to requested -SDK session IDs, with SDK/runtime versions, protocol, retention and remote mode. These per-exchange -details are not automatically added to response metadata or target-identity exports. An attempt record -does not prove session allocation or a successful exchange; retained-session logs identify data kept -for diagnostics. - -This target supports fresh single-turn text only, not a public caller system prompt or native -continuation. It requests text-only runtime restrictions and retains restricted SDK base instructions; -other SDK/runtime context may remain. These controls are not an OS sandbox. +`cleanup_target_async()` is caller-owned, terminal cleanup. It drains active target work, rejects +new or queued target work, releases owned resources, and surfaces cleanup failures. It does not +gather caller tasks that may themselves be waiting on cleanup. By default, `retain_session=False` +deletes owned native session data. With `retain_session=True`, native session data is retained +but live session resources are still disconnected and the shared client is stopped. PyRIT memory +records remain independently retained. Crashes can still leave residual data. Remote session export +is explicitly `OFF` in either mode: local retention does not enable export, and `OFF` does not mean +offline inference. + +`reset_conversation_async()` releases one established conversation without stopping unrelated +conversations or the shared client. A reset conversation is terminal: later sends with that +conversation ID fail rather than silently reopening native history. Unknown IDs and repeated +resets are safe no-ops. To start a new conversation, use a new PyRIT conversation ID; this target +does not resume, import, replay, fork, or reconcile native history. Edits to PyRIT memory do not +update the native Copilot context. + +Capture the target's INFO logs for session-creation attempt records linking PyRIT conversation IDs +to requested SDK session IDs, with SDK/runtime versions, protocol, retention and remote mode. These +records are emitted before each native session creation attempt, not on every turn, and are not +automatically added to response metadata or target-identity exports. An attempt record does not +prove session allocation or a successful exchange; retained-session logs identify data kept for +diagnostics. + +Each turn sends only the newest normalized user content; prior turns remain in the native session. +The target supports native continuation for text conversations and does not support editable history. + +If a caller supplies an initial system prompt before the first user turn, the target uses SDK +replacement mode with that exact text. Replacement replaces the SDK's default system message and +its guardrails/security restrictions. The target still explicitly requests remote `OFF`, an empty +tool list, disabled configuration and instruction discovery, disabled host Git operations, session +store, memory, and file hooks. Without a supplied system prompt, the target keeps the existing +customize mode and removes only `environment_context` and `custom_instructions`; other SDK/runtime +context may remain. These controls are not an OS sandbox. Later public system-prompt changes are +rejected. + +Finalize `model_name` before the first identity-capturing operation: `get_identifier()`, system +prompt setup, normalizer registration, scenario or registry use, target mapping, or direct native +execution. After identity capture, a different model requires a new target; setting the same model +again is harmless. Native SDK model switching is not used. `response_timeout_seconds` bounds dispatch and completion, excluding client/session startup, the -startup status lookup, PyRIT pacing and cleanup. A failed status lookup stops the exchange before -session creation or prompt dispatch. Errors, timeouts and cancellation propagate without automatic -replay or a rollback guarantee. +startup status lookup, PyRIT pacing and cleanup. The synchronous SDK client constructor is moved +off the event loop, but cancellation cannot abort synchronous work already running in its worker. +A failed status lookup stops the exchange before session creation or prompt dispatch. + +SDK-observed ambiguous timeout, cancellation, abort, unsafe event, or invalid native outcomes retire +only the affected conversation. Later sends to that ID fail, while a fresh ID can proceed; there is +no automatic replay. Errors, timeouts, and cancellation propagate without a crash-proof cleanup +guarantee. Downstream converter or PyRIT persistence failures occur outside the target's native +observation boundary and do not provide a safe-continuation guarantee. diff --git a/pyrit/prompt_target/common/prompt_target.py b/pyrit/prompt_target/common/prompt_target.py index 790abcbba4..f3bb4457d5 100644 --- a/pyrit/prompt_target/common/prompt_target.py +++ b/pyrit/prompt_target/common/prompt_target.py @@ -334,13 +334,16 @@ def set_system_prompt( conversation_id (str): The conversation id to attach the prompt to. Raises: - ValueError: If the target does not support multi-turn or editable history. + ValueError: If the target does not support multi-turn conversations, or + supports neither editable history nor native system prompts. RuntimeError: If the conversation already has messages. """ - if not self.capabilities.supports_multi_turn or not self.capabilities.supports_editable_history: + if not self.capabilities.supports_multi_turn or not ( + self.capabilities.supports_editable_history or self.capabilities.supports_system_prompt + ): raise ValueError( f"Target {type(self).__name__} does not support setting a system prompt. " - "It must support both multi-turn conversations and editable history." + "It must support multi-turn conversations and either editable history or native system prompts." ) messages = self._memory.get_conversation_messages(conversation_id=conversation_id) diff --git a/pyrit/prompt_target/github_copilot_target.py b/pyrit/prompt_target/github_copilot_target.py index acdbb38700..a5f70c1410 100644 --- a/pyrit/prompt_target/github_copilot_target.py +++ b/pyrit/prompt_target/github_copilot_target.py @@ -4,23 +4,41 @@ import asyncio import logging import math +from dataclasses import dataclass, field from pathlib import Path +from typing import TYPE_CHECKING from uuid import uuid4 from pyrit.models import ComponentIdentifier, Message, construct_response_from_request from pyrit.prompt_target.common.prompt_target import PromptTarget +from pyrit.prompt_target.common.target_capabilities import TargetCapabilities +from pyrit.prompt_target.common.target_configuration import TargetConfiguration from pyrit.prompt_target.common.utils import limit_requests_per_minute logger = logging.getLogger(__name__) +if TYPE_CHECKING: + from copilot import CopilotClient, CopilotSession, GetStatusResponse, SystemMessageConfig + + +@dataclass +class _ConversationState: + lock: asyncio.Lock = field(default_factory=asyncio.Lock) + session: "CopilotSession | None" = None + retired: bool = False + class GitHubCopilotTarget(PromptTarget): """ - Send single-turn text requests through the GitHub Copilot SDK. + Send text requests through the GitHub Copilot SDK. Capture INFO logs for session mapping and SDK/runtime version diagnostics. """ + _DEFAULT_CONFIGURATION: TargetConfiguration = TargetConfiguration( + capabilities=TargetCapabilities(supports_multi_turn=True, supports_system_prompt=True) + ) + def __init__( self, *, @@ -81,6 +99,15 @@ def __init__( self._github_token = github_token self._retain_session = retain_session self._response_timeout_seconds = response_timeout_seconds + self._client: CopilotClient | None = None + self._runtime_status: GetStatusResponse | None = None + self._client_start_lock = asyncio.Lock() + self._lifecycle_lock = asyncio.Lock() + self._cleanup_task: asyncio.Task[None] | None = None + self._conversations: dict[str, _ConversationState] = {} + self._active_target_operations = 0 + self._active_operations_drained = asyncio.Event() + self._active_operations_drained.set() def _build_identifier(self) -> ComponentIdentifier: """ @@ -91,15 +118,68 @@ def _build_identifier(self) -> ComponentIdentifier: """ return self._create_identifier(params={"working_directory": self._working_directory}) + def set_model_name(self, *, model_name: str) -> None: + """ + Set the model before identity capture; afterward, create a new target to change it. + + Args: + model_name (str): The nonblank Copilot model ID. + + Raises: + ValueError: If model_name is blank. + RuntimeError: If the target identity has already been captured and model_name differs. + """ + if not model_name.strip(): + raise ValueError("model_name must not be empty.") + if self._identifier is not None and model_name != self._model_name: + raise RuntimeError("model_name is frozen after identity capture; create a new target to change it.") + super().set_model_name(model_name=model_name) + @limit_requests_per_minute async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Message]) -> list[Message]: - request = normalized_conversation[-1].get_piece() - reply_text = await self._send_text_async( - prompt=request.converted_value, conversation_id=request.conversation_id - ) - return [construct_response_from_request(request=request, response_text_pieces=[reply_text])] + self.get_identifier() + async with self._lifecycle_lock: + if self._cleanup_task is not None: + raise RuntimeError("GitHubCopilotTarget has been cleaned up and cannot send more prompts.") + self._active_target_operations += 1 + self._active_operations_drained.clear() + try: + request = normalized_conversation[-1].get_piece() + conversation_id = request.conversation_id or "" + conversation = self._conversations.setdefault(conversation_id, _ConversationState()) + async with conversation.lock: + async with self._lifecycle_lock: + if self._cleanup_task is not None: + raise RuntimeError("GitHubCopilotTarget has been cleaned up and cannot send more prompts.") + initial_system_prompt: str | None = None + if normalized_conversation[0].api_role == "system": + initial_system_prompt = "\n\n".join( + piece.converted_value for piece in normalized_conversation[0].message_pieces + ) + session = await self._get_or_create_session_async( + conversation_id=conversation_id, + initial_system_prompt=initial_system_prompt, + ) + try: + reply_text = await self._send_text_async( + session=session, + prompt=request.converted_value, + ) + except BaseException: + await asyncio.shield( + self._retire_conversation_async( + conversation=conversation, + ) + ) + raise + return [construct_response_from_request(request=request, response_text_pieces=[reply_text])] + finally: + async with self._lifecycle_lock: + self._active_target_operations -= 1 + if self._active_target_operations == 0: + self._active_operations_drained.set() - async def _send_text_async(self, *, prompt: str, conversation_id: str | None) -> str: + async def _send_text_async(self, *, session: "CopilotSession", prompt: str) -> str: from copilot.generated.session_events import ( AbortData, AssistantMessageData, @@ -118,82 +198,232 @@ def _record_turn_events(event: SessionEvent) -> None: if isinstance(event.data, ToolExecutionStartData): tool_execution_started = True - client = await asyncio.to_thread( - self._sdk.CopilotClient, github_token=self._github_token, working_directory=self._working_directory + unsubscribe = session.on(_record_turn_events) + try: + reply = await asyncio.wait_for( + session.send_and_wait(prompt, timeout=self._response_timeout_seconds), + timeout=self._response_timeout_seconds, + ) + finally: + unsubscribe() + + if aborted: + raise RuntimeError("Copilot turn was aborted.") + if tool_execution_started: + raise RuntimeError("Copilot turn reported tool execution.") + if ( + reply is None + or reply.agent_id is not None + or not isinstance(reply.data, AssistantMessageData) + or not isinstance(reply.data.content, str) + or not reply.data.content + ): + raise ValueError("Copilot did not return a non-empty root assistant text reply.") + return reply.data.content + + async def cleanup_target_async(self) -> None: + """ + Stop accepting target work, drain active operations, and release owned SDK resources. + + Cleanup is terminal and idempotent. Retained sessions are preserved while the shared + client is always stopped. Cleanup attempts every owned resource before surfacing failures. + """ + async with self._lifecycle_lock: + if self._cleanup_task is None: + self._cleanup_task = asyncio.create_task(self._cleanup_owned_resources_async()) + cleanup_task = self._cleanup_task + await asyncio.shield(cleanup_task) + + async def reset_conversation_async(self, *, conversation_id: str) -> None: + """ + Release one established Copilot conversation without stopping the shared client. + + Unknown or already released conversations are no-ops. Established conversations + become retired before release, so a later send cannot silently create a new native + history for the same PyRIT conversation ID. + + Args: + conversation_id (str): The PyRIT conversation ID to release. + """ + async with self._lifecycle_lock: + conversation = self._conversations.get(conversation_id) + if conversation is None or (conversation.retired and conversation.session is None): + return + cleanup_task = self._cleanup_task + if cleanup_task is None: + self._active_target_operations += 1 + self._active_operations_drained.clear() + + if cleanup_task is not None: + await asyncio.shield(cleanup_task) + return + + try: + async with conversation.lock: + await self._retire_conversation_async(conversation=conversation) + finally: + async with self._lifecycle_lock: + self._active_target_operations -= 1 + if self._active_target_operations == 0: + self._active_operations_drained.set() + + async def _get_or_create_session_async( + self, + *, + conversation_id: str, + initial_system_prompt: str | None, + ) -> "CopilotSession": + async with self._lifecycle_lock: + conversation = self._conversations[conversation_id] + if conversation.retired: + raise RuntimeError( + f"Copilot conversation {conversation_id} was retired and cannot accept further sends; " + "use a new conversation ID." + ) + existing_session = conversation.session + if existing_session is not None: + return existing_session + + client = await self._get_or_start_client_async() + session_id = str(uuid4()) + status = self._runtime_status + if status is None: + raise RuntimeError("Copilot runtime status is unavailable after client startup.") + logger.info( + "Attempting Copilot session creation: pyrit_conversation_id=%s requested_sdk_session_id=%s " + "sdk_version=%s runtime_version=%s runtime_protocol_version=%s retain_session=%s remote_mode=OFF", + conversation_id, + session_id, + self._sdk.__version__, + status.version, + status.protocol_version, + self._retain_session, ) + try: - await client.start() - status = await client.get_status() - session_id = str(uuid4()) - logger.info( - "Attempting Copilot session creation: pyrit_conversation_id=%s requested_sdk_session_id=%s " - "sdk_version=%s runtime_version=%s runtime_protocol_version=%s retain_session=%s remote_mode=OFF", - conversation_id, - session_id, - self._sdk.__version__, - status.version, - status.protocol_version, - self._retain_session, + system_message: SystemMessageConfig + if initial_system_prompt is None: + system_message = { + "mode": "customize", + "sections": { + "environment_context": {"action": "remove"}, + "custom_instructions": {"action": "remove"}, + }, + } + else: + system_message = {"mode": "replace", "content": initial_system_prompt} + session = await client.create_session( + session_id=session_id, + model=self._model_name, + system_message=system_message, + remote_session=self._sdk.RemoteSessionMode.OFF, + available_tools=[], + skip_custom_instructions=True, + instruction_directories=[], + enable_host_git_operations=False, + enable_config_discovery=False, + organization_custom_instructions="", + enable_on_demand_instruction_discovery=False, + infinite_sessions={"enabled": False}, + memory={"enabled": False}, + enable_session_store=False, + enable_file_hooks=False, + ) + async with self._lifecycle_lock: + conversation.session = session + return session + except BaseException: + await self._cleanup_allocated_session_async(client=client, session_id=session_id) + raise + + async def _get_or_start_client_async(self) -> "CopilotClient": + async with self._client_start_lock: + if self._client is not None: + return self._client + + client = await asyncio.to_thread( + self._sdk.CopilotClient, + github_token=self._github_token, + working_directory=self._working_directory, ) try: - # SDK runtime instructions are retained; other SDK-provided context may remain. - session = await client.create_session( - session_id=session_id, - model=self._model_name, - system_message={ - "mode": "customize", - "sections": { - "environment_context": {"action": "remove"}, - "custom_instructions": {"action": "remove"}, - }, - }, - remote_session=self._sdk.RemoteSessionMode.OFF, - available_tools=[], - skip_custom_instructions=True, - instruction_directories=[], - enable_host_git_operations=False, - enable_config_discovery=False, - organization_custom_instructions="", - enable_on_demand_instruction_discovery=False, - infinite_sessions={"enabled": False}, - memory={"enabled": False}, - enable_session_store=False, - enable_file_hooks=False, - ) - except (Exception, asyncio.CancelledError): - if await client.get_session_metadata(session_id) is not None: - if self._retain_session: - logger.info("Retaining Copilot session %s as requested; delete it manually.", session_id) - else: - await client.delete_session(session_id) + await client.start() + self._runtime_status = await client.get_status() + except BaseException as error: + try: + await client.stop() + except BaseException as cleanup_error: + if isinstance(error, asyncio.CancelledError): + raise error from cleanup_error + raise BaseExceptionGroup( + "Copilot client startup and cleanup failed", + [error, cleanup_error], + ) from error raise - try: - unsubscribe = session.on(_record_turn_events) + + async with self._lifecycle_lock: + self._client = client + return client + + async def _retire_conversation_async(self, *, conversation: _ConversationState) -> None: + async with self._lifecycle_lock: + session = conversation.session + client = self._client + if session is None or client is None: + return + conversation.retired = True + await self._release_session_async(client=client, conversation=conversation) + + async def _cleanup_allocated_session_async(self, *, client: "CopilotClient", session_id: str) -> None: + if await client.get_session_metadata(session_id) is None: + return + if self._retain_session: + logger.info("Retaining Copilot session %s as requested; delete it manually.", session_id) + else: + await client.delete_session(session_id) + + async def _release_session_async( + self, + *, + client: "CopilotClient", + conversation: _ConversationState, + ) -> None: + session = conversation.session + if session is None: + return + session_id = session.session_id + if self._retain_session: + await session.disconnect() + logger.info("Retaining Copilot session %s as requested; delete it manually.", session_id) + else: + await client.delete_session(session_id) + async with self._lifecycle_lock: + if conversation.session is session: + conversation.session = None + + async def _cleanup_owned_resources_async(self) -> None: + await self._active_operations_drained.wait() + async with self._lifecycle_lock: + conversations = [ + conversation for conversation in self._conversations.values() if conversation.session is not None + ] + client = self._client + self._client = None + self._runtime_status = None + + errors: list[BaseException] = [] + if client is not None: + for conversation in conversations: try: - reply = await asyncio.wait_for( - session.send_and_wait(prompt, timeout=self._response_timeout_seconds), - timeout=self._response_timeout_seconds, - ) - finally: - unsubscribe() - - if aborted: - raise RuntimeError("Copilot turn was aborted.") - if tool_execution_started: - raise RuntimeError("Copilot turn reported tool execution.") - if ( - reply is None - or reply.agent_id is not None - or not isinstance(reply.data, AssistantMessageData) - or not isinstance(reply.data.content, str) - or not reply.data.content - ): - raise ValueError("Copilot did not return a non-empty root assistant text reply.") - return reply.data.content - finally: - if self._retain_session: - logger.info("Retaining Copilot session %s as requested; delete it manually.", session.session_id) - else: - await client.delete_session(session.session_id) - finally: - await client.stop() + await self._release_session_async(client=client, conversation=conversation) + except BaseException as error: + errors.append(error) + try: + await client.stop() + except BaseException as error: + errors.append(error) + + if len(errors) == 1: + raise errors[0] + if errors: + raise BaseExceptionGroup("Copilot target cleanup failed", errors) diff --git a/tests/unit/prompt_target/target/test_github_copilot_target.py b/tests/unit/prompt_target/target/test_github_copilot_target.py index 49fd2f10e1..414125b558 100644 --- a/tests/unit/prompt_target/target/test_github_copilot_target.py +++ b/tests/unit/prompt_target/target/test_github_copilot_target.py @@ -6,11 +6,12 @@ import asyncio import logging import threading +from contextlib import suppress from dataclasses import replace from datetime import UTC, datetime from functools import partial from typing import TYPE_CHECKING, Any -from unittest.mock import AsyncMock, NonCallableMagicMock, create_autospec, patch +from unittest.mock import AsyncMock, NonCallableMagicMock, call, create_autospec, patch from uuid import UUID, uuid4 import pytest @@ -78,6 +79,108 @@ async def get_session_metadata_async(session_id: str) -> SessionMetadata | None: client.delete_session.side_effect = sessions.remove +def _make_sdk_session( + *, + sdk: Any, + session_id: str, +) -> NonCallableMagicMock: + session = create_autospec(sdk.CopilotSession, instance=True) + assert isinstance(session, NonCallableMagicMock) + session.session_id = session_id + return session + + +def _user_message( + *, + original_value: str, + converted_value: str | None = None, + conversation_id: str | None = None, +) -> Message: + return MessagePiece( + role="user", + conversation_id=conversation_id, + original_value=original_value, + converted_value=original_value if converted_value is None else converted_value, + ).to_message() + + +async def _send_normalized_async( + *, + target: GitHubCopilotTarget, + original_value: str, + converted_value: str | None = None, + conversation_id: str | None = None, +) -> Message: + return await PromptNormalizer().send_prompt_async( + message=_user_message( + original_value=original_value, + converted_value=converted_value, + ), + conversation_id=conversation_id, + target=target, + ) + + +def _message_state(*, memory: MemoryInterface, conversation_id: str) -> list[tuple[str, str, str, str | None, str]]: + return [ + ( + message.get_piece().role, + message.get_piece().original_value, + message.get_piece().converted_value, + message.get_piece().conversation_id, + message.get_piece().response_error, + ) + for message in memory.get_conversation_messages(conversation_id=conversation_id) + ] + + +def _message_roles_and_errors(*, memory: MemoryInterface, conversation_id: str) -> list[tuple[str, str]]: + return [ + (message.get_piece().role, message.get_piece().response_error) + for message in memory.get_conversation_messages(conversation_id=conversation_id) + ] + + +def _message_values_and_errors(*, memory: MemoryInterface, conversation_id: str) -> list[tuple[str, str, str]]: + return [ + (message.get_piece().role, message.get_piece().converted_value, message.get_piece().response_error) + for message in memory.get_conversation_messages(conversation_id=conversation_id) + ] + + +def _expected_session_configuration(*, system_message: dict[str, Any]) -> dict[str, Any]: + from copilot.generated.rpc import RemoteSessionMode + + return { + "model": "gpt-5-mini", + "system_message": system_message, + "remote_session": RemoteSessionMode.OFF, + "available_tools": [], + "skip_custom_instructions": True, + "instruction_directories": [], + "enable_host_git_operations": False, + "enable_config_discovery": False, + "organization_custom_instructions": "", + "enable_on_demand_instruction_discovery": False, + "infinite_sessions": {"enabled": False}, + "memory": {"enabled": False}, + "enable_session_store": False, + "enable_file_hooks": False, + } + + +async def _cancel_tasks_async(*tasks: asyncio.Task[Any] | None) -> None: + for task in tasks: + if task is not None and not task.done(): + task.cancel() + await asyncio.gather(*[task for task in tasks if task is not None], return_exceptions=True) + + +def _assert_no_resource_release(client: NonCallableMagicMock) -> None: + client.delete_session.assert_not_awaited() + client.stop.assert_not_awaited() + + @pytest.mark.usefixtures("patch_central_database") @pytest.mark.parametrize("retain_session", [False, True], ids=["delete", "retain"]) async def test_normalizer_round_trip_and_retention_async( @@ -89,15 +192,16 @@ async def test_normalizer_round_trip_and_retention_async( retain_session: bool, ) -> None: conversation_id = str(uuid4()) - request = MessagePiece( - role="user", original_value="Original text before conversion.", converted_value="Reply exactly HELLO." - ).to_message() + target = GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=retain_session) with patch.object(sdk, "__version__", "9.8.7"), caplog.at_level(logging.INFO, logger=TARGET_LOGGER): - response = await PromptNormalizer().send_prompt_async( - message=request, + response = await _send_normalized_async( + target=target, + original_value="Original text before conversion.", + converted_value="Reply exactly HELLO.", conversation_id=conversation_id, - target=GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=retain_session), ) + _assert_no_resource_release(client=client) + await target.cleanup_target_async() assert isinstance(response, Message) piece = response.get_piece() @@ -107,13 +211,9 @@ async def test_normalizer_round_trip_and_retention_async( conversation_id, "none", ) - stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) - assert [ - (m.get_piece().role, m.get_piece().original_value, m.get_piece().converted_value, m.get_piece().response_error) - for m in stored - ] == [ - ("user", "Original text before conversion.", "Reply exactly HELLO.", "none"), - ("assistant", "HELLO", "HELLO", "none"), + assert _message_state(memory=sqlite_instance, conversation_id=conversation_id) == [ + ("user", "Original text before conversion.", "Reply exactly HELLO.", conversation_id, "none"), + ("assistant", "HELLO", "HELLO", conversation_id, "none"), ] session = client.create_session.return_value session.send_and_wait.assert_awaited_once_with("Reply exactly HELLO.", timeout=60.0) @@ -122,11 +222,21 @@ async def test_normalizer_round_trip_and_retention_async( client.get_status.assert_awaited_once() client.create_session.assert_awaited_once() client.get_session_metadata.assert_not_awaited() - client.stop.assert_awaited_once() sdk.CopilotClient.assert_called_once_with(github_token=None, working_directory=None) requested_id = client.create_session.await_args.kwargs["session_id"] assert str(UUID(requested_id)) == requested_id assert requested_id != "sdk-session-id" + configuration = dict(client.create_session.await_args.kwargs) + configuration.pop("session_id") + assert configuration == _expected_session_configuration( + system_message={ + "mode": "customize", + "sections": { + "environment_context": {"action": "remove"}, + "custom_instructions": {"action": "remove"}, + }, + } + ) records = [r.getMessage() for r in caplog.records if r.name == TARGET_LOGGER and r.levelno == logging.INFO] assert len(records) == (2 if retain_session else 1) assert records[0].startswith("Attempting Copilot session creation:") @@ -145,52 +255,634 @@ async def test_normalizer_round_trip_and_retention_async( assert records[1] == "Retaining Copilot session sdk-session-id as requested; delete it manually." else: client.delete_session.assert_awaited_once_with("sdk-session-id") + client.stop.assert_awaited_once() @pytest.mark.usefixtures("patch_central_database") -async def test_create_session_restricts_copilot_runtime_async(client: NonCallableMagicMock) -> None: - from copilot.generated.rpc import RemoteSessionMode +@pytest.mark.parametrize( + "initial_system_prompt", + [pytest.param(None, id="default-customize"), pytest.param("initial system instructions", id="initial-replacement")], +) +async def test_normalizer_continues_native_session_across_turns_async( + *, + sdk: Any, + client: NonCallableMagicMock, + sqlite_instance: MemoryInterface, + initial_system_prompt: str | None, +) -> None: + conversation_id = "native-two-turn-conversation" + session = client.create_session.return_value + session.session_id = "sdk-session-id" + session.send_and_wait.side_effect = [_assistant_reply("FIRST"), _assistant_reply("SECOND")] + target = GitHubCopilotTarget(model_name="gpt-5-mini") - await PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), - target=GitHubCopilotTarget(model_name="gpt-5-mini"), + if initial_system_prompt is not None: + target.set_system_prompt(system_prompt=initial_system_prompt, conversation_id=conversation_id) + + first_response = await _send_normalized_async( + target=target, + original_value="first original", + converted_value="first prepared", + conversation_id=conversation_id, + ) + assert (first_response.get_piece().role, first_response.get_piece().converted_value) == ( + "assistant", + "FIRST", ) + configuration = dict(client.create_session.await_args.kwargs) - assert str(UUID(configuration.pop("session_id"))) - assert configuration == { - "model": "gpt-5-mini", - "system_message": { - "mode": "customize", - "sections": { - "environment_context": {"action": "remove"}, - "custom_instructions": {"action": "remove"}, - }, - }, - "remote_session": RemoteSessionMode.OFF, - "available_tools": [], - "skip_custom_instructions": True, - "instruction_directories": [], - "enable_host_git_operations": False, - "enable_config_discovery": False, - "organization_custom_instructions": "", - "enable_on_demand_instruction_discovery": False, - "infinite_sessions": {"enabled": False}, - "memory": {"enabled": False}, - "enable_session_store": False, - "enable_file_hooks": False, + requested_session_id = configuration.pop("session_id") + assert str(UUID(requested_session_id)) == requested_session_id + assert configuration == _expected_session_configuration( + system_message=( + {"mode": "replace", "content": initial_system_prompt} + if initial_system_prompt is not None + else { + "mode": "customize", + "sections": { + "environment_context": {"action": "remove"}, + "custom_instructions": {"action": "remove"}, + }, + } + ) + ) + sdk.CopilotClient.assert_called_once_with(github_token=None, working_directory=None) + client.start.assert_awaited_once() + client.get_status.assert_awaited_once() + + if initial_system_prompt is not None: + with pytest.raises(RuntimeError, match="Conversation already exists"): + target.set_system_prompt(system_prompt="different system instructions", conversation_id=conversation_id) + + second_response = await _send_normalized_async( + target=target, + original_value="second original", + converted_value="second prepared", + conversation_id=conversation_id, + ) + assert (second_response.get_piece().role, second_response.get_piece().converted_value) == ( + "assistant", + "SECOND", + ) + assert session.send_and_wait.await_args_list == [ + call("first prepared", timeout=60.0), + call("second prepared", timeout=60.0), + ] + client.create_session.assert_awaited_once() + expected_messages = [ + ("user", "first original", "first prepared", conversation_id, "none"), + ("assistant", "FIRST", "FIRST", conversation_id, "none"), + ("user", "second original", "second prepared", conversation_id, "none"), + ("assistant", "SECOND", "SECOND", conversation_id, "none"), + ] + if initial_system_prompt is not None: + expected_messages.insert(0, ("system", initial_system_prompt, initial_system_prompt, conversation_id, "none")) + assert _message_state(memory=sqlite_instance, conversation_id=conversation_id) == expected_messages + _assert_no_resource_release(client=client) + await target.cleanup_target_async() + client.delete_session.assert_awaited_once_with(session.session_id) + client.stop.assert_awaited_once() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_distinct_conversations_progress_on_shared_client_async( + *, + sdk: Any, + client: NonCallableMagicMock, +) -> None: + session_a = _make_sdk_session(sdk=sdk, session_id="sdk-session-a") + session_b = _make_sdk_session(sdk=sdk, session_id="sdk-session-b") + first_send_started = asyncio.Event() + release_first_send = asyncio.Event() + + async def send_a_async(*_args: Any, **_kwargs: Any) -> Any: + first_send_started.set() + await release_first_send.wait() + return _assistant_reply("A") + + session_a.send_and_wait.side_effect = send_a_async + session_b.send_and_wait.return_value = _assistant_reply("B") + client.create_session.side_effect = [session_a, session_b] + target = GitHubCopilotTarget(model_name="gpt-5-mini") + first_task = asyncio.create_task( + target.send_prompt_async( + message=_user_message( + conversation_id="conversation-a", + original_value="a original", + converted_value="a prepared", + ) + ) + ) + second_task: asyncio.Task[list[Message]] | None = None + try: + await asyncio.wait_for(first_send_started.wait(), timeout=2.0) + second_task = asyncio.create_task( + target.send_prompt_async( + message=_user_message( + conversation_id="conversation-b", + original_value="b original", + converted_value="b prepared", + ) + ) + ) + second_response = await asyncio.wait_for(second_task, timeout=2.0) + assert not release_first_send.is_set() + assert (second_response[0].get_piece().conversation_id, second_response[0].get_piece().converted_value) == ( + "conversation-b", + "B", + ) + + release_first_send.set() + first_response = await asyncio.wait_for(first_task, timeout=2.0) + assert (first_response[0].get_piece().conversation_id, first_response[0].get_piece().converted_value) == ( + "conversation-a", + "A", + ) + _assert_no_resource_release(client=client) + await target.cleanup_target_async() + finally: + release_first_send.set() + await _cancel_tasks_async(first_task, second_task) + with suppress(Exception): + await asyncio.wait_for(target.cleanup_target_async(), timeout=2.0) + + sdk.CopilotClient.assert_called_once_with(github_token=None, working_directory=None) + client.start.assert_awaited_once() + client.get_status.assert_awaited_once() + assert client.create_session.await_count == 2 + requested_session_ids = [entry.kwargs["session_id"] for entry in client.create_session.await_args_list] + assert all(str(UUID(session_id)) == session_id for session_id in requested_session_ids) + assert len(set(requested_session_ids)) == 2 + session_a.send_and_wait.assert_awaited_once_with("a prepared", timeout=60.0) + session_b.send_and_wait.assert_awaited_once_with("b prepared", timeout=60.0) + assert client.delete_session.await_args_list == [call(session_a.session_id), call(session_b.session_id)] + client.stop.assert_awaited_once() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_reset_conversation_releases_only_requested_session_async( + *, + sdk: Any, + client: NonCallableMagicMock, + sqlite_instance: MemoryInterface, +) -> None: + session_a = _make_sdk_session(sdk=sdk, session_id="sdk-session-a") + session_a.send_and_wait.side_effect = [_assistant_reply("A1"), _assistant_reply("A_REOPENED")] + session_b = _make_sdk_session(sdk=sdk, session_id="sdk-session-b") + session_b.send_and_wait.side_effect = [_assistant_reply("B1"), _assistant_reply("B2")] + client.create_session.side_effect = [session_a, session_b] + target = GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=True) + + await _send_normalized_async( + target=target, + original_value="a original", + converted_value="a prepared", + conversation_id="conversation-a", + ) + await _send_normalized_async( + target=target, + original_value="b original", + converted_value="b prepared", + conversation_id="conversation-b", + ) + memory_before_reset = { + conversation_id: _message_values_and_errors(memory=sqlite_instance, conversation_id=conversation_id) + for conversation_id in ("conversation-a", "conversation-b") + } + + await target.reset_conversation_async(conversation_id="conversation-a") + assert session_a.disconnect.await_count == 1 + _assert_no_resource_release(client=client) + session_b.disconnect.assert_not_awaited() + memory_after_reset = { + conversation_id: _message_values_and_errors(memory=sqlite_instance, conversation_id=conversation_id) + for conversation_id in ("conversation-a", "conversation-b") } + assert memory_after_reset == memory_before_reset + + await target.reset_conversation_async(conversation_id="conversation-a") + assert session_a.disconnect.await_count == 1 + + with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as reset_error: + await _send_normalized_async( + target=target, + original_value="a retry original", + converted_value="a retry prepared", + conversation_id="conversation-a", + ) + assert isinstance(reset_error.value.__cause__, RuntimeError) + assert "retired" in str(reset_error.value.__cause__).lower() + + response_b = await _send_normalized_async( + target=target, + original_value="b second original", + converted_value="b second prepared", + conversation_id="conversation-b", + ) + assert response_b.get_piece().converted_value == "B2" + assert client.create_session.await_count == 2 + session_a.send_and_wait.assert_awaited_once_with("a prepared", timeout=60.0) + session_b.send_and_wait.assert_has_awaits( + [ + call("b prepared", timeout=60.0), + call("b second prepared", timeout=60.0), + ] + ) + client.delete_session.assert_not_awaited() + await target.cleanup_target_async() + client.delete_session.assert_not_awaited() + client.stop.assert_awaited_once() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_reset_waits_for_in_progress_cleanup_release_async( + *, + client: NonCallableMagicMock, +) -> None: + session = client.create_session.return_value + target = GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=True) + await _send_normalized_async( + target=target, + original_value="a original", + converted_value="a prepared", + conversation_id="conversation-a", + ) + disconnect_started = asyncio.Event() + release_disconnect = asyncio.Event() + + async def disconnect_async() -> None: + disconnect_started.set() + await release_disconnect.wait() + + session.disconnect.side_effect = disconnect_async + cleanup_task = asyncio.create_task(target.cleanup_target_async()) + reset_task: asyncio.Task[None] | None = None + try: + await asyncio.wait_for(disconnect_started.wait(), timeout=2.0) + reset_task = asyncio.create_task(target.reset_conversation_async(conversation_id="conversation-a")) + await asyncio.sleep(0) + assert not reset_task.done() + + release_disconnect.set() + await asyncio.wait_for(cleanup_task, timeout=2.0) + await asyncio.wait_for(reset_task, timeout=2.0) + session.disconnect.assert_awaited_once() + client.delete_session.assert_not_awaited() + client.stop.assert_awaited_once() + finally: + release_disconnect.set() + await _cancel_tasks_async(cleanup_task, reset_task) + + +@pytest.mark.usefixtures("patch_central_database") +async def test_reset_waits_for_conversation_creation_async( + *, + client: NonCallableMagicMock, +) -> None: + session = client.create_session.return_value + session.session_id = "sdk-session-a" + session.send_and_wait.return_value = _assistant_reply("FIRST") + create_started = asyncio.Event() + release_creation = asyncio.Event() + + async def create_session_async(*_args: Any, **_kwargs: Any) -> Any: + create_started.set() + await release_creation.wait() + return session + + client.create_session.side_effect = create_session_async + target = GitHubCopilotTarget(model_name="gpt-5-mini") + conversation_id = "provisioning-conversation" + first_task = asyncio.create_task( + target.send_prompt_async( + message=_user_message( + conversation_id=conversation_id, + original_value="first original", + converted_value="first prepared", + ) + ) + ) + reset_task: asyncio.Task[None] | None = None + try: + await asyncio.wait_for(create_started.wait(), timeout=2.0) + reset_task = asyncio.create_task(target.reset_conversation_async(conversation_id=conversation_id)) + await asyncio.sleep(0) + assert not reset_task.done() + + release_creation.set() + first_response = await asyncio.wait_for(first_task, timeout=2.0) + await asyncio.wait_for(reset_task, timeout=2.0) + assert first_response[0].get_piece().converted_value == "FIRST" + session.send_and_wait.assert_awaited_once_with("first prepared", timeout=60.0) + client.create_session.assert_awaited_once() + client.delete_session.assert_awaited_once_with(session.session_id) + client.stop.assert_not_awaited() + + await target.reset_conversation_async(conversation_id=conversation_id) + with pytest.raises(RuntimeError, match="retired"): + await target.send_prompt_async( + message=_user_message( + conversation_id=conversation_id, + original_value="retry original", + converted_value="retry prepared", + ) + ) + assert client.create_session.await_count == 1 + await target.cleanup_target_async() + client.stop.assert_awaited_once() + finally: + release_creation.set() + await _cancel_tasks_async(first_task, reset_task) + with suppress(Exception): + await asyncio.wait_for(target.cleanup_target_async(), timeout=2.0) + + +@pytest.mark.usefixtures("patch_central_database") +async def test_repeated_reset_does_not_join_unrelated_cleanup_async( + *, + sdk: Any, + client: NonCallableMagicMock, +) -> None: + session_a = _make_sdk_session(sdk=sdk, session_id="sdk-session-a") + session_a.send_and_wait.return_value = _assistant_reply("A") + session_b = _make_sdk_session(sdk=sdk, session_id="sdk-session-b") + session_b.send_and_wait.return_value = _assistant_reply("B") + client.create_session.side_effect = [session_a, session_b] + target = GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=True) + + for conversation_id, prompt in (("conversation-a", "a prepared"), ("conversation-b", "b prepared")): + await target.send_prompt_async( + message=_user_message( + conversation_id=conversation_id, + original_value=prompt, + ) + ) + await target.reset_conversation_async(conversation_id="conversation-a") + session_a.disconnect.assert_awaited_once() + + disconnect_started = asyncio.Event() + release_disconnect = asyncio.Event() + + async def disconnect_b_async() -> None: + disconnect_started.set() + await release_disconnect.wait() + raise RuntimeError("B disconnect failed") + + session_b.disconnect.side_effect = disconnect_b_async + cleanup_task = asyncio.create_task(target.cleanup_target_async()) + reset_task: asyncio.Task[None] | None = None + try: + await asyncio.wait_for(disconnect_started.wait(), timeout=2.0) + reset_task = asyncio.create_task(target.reset_conversation_async(conversation_id="conversation-a")) + await asyncio.sleep(0) + assert reset_task.done() + await reset_task + assert session_a.disconnect.await_count == 1 + client.stop.assert_not_awaited() + + release_disconnect.set() + with pytest.raises(RuntimeError, match="B disconnect failed"): + await asyncio.wait_for(cleanup_task, timeout=2.0) + session_b.disconnect.assert_awaited_once() + client.stop.assert_awaited_once() + finally: + release_disconnect.set() + await _cancel_tasks_async(cleanup_task, reset_task) + + +@pytest.mark.usefixtures("patch_central_database") +async def test_normalizer_rejects_retired_conversation_but_allows_fresh_async( + *, + sdk: Any, + client: NonCallableMagicMock, + sqlite_instance: MemoryInterface, +) -> None: + session_a = _make_sdk_session(sdk=sdk, session_id="sdk-session-a") + session_a.send_and_wait.side_effect = TimeoutError("ambiguous mock send") + + session_b = _make_sdk_session(sdk=sdk, session_id="sdk-session-b") + session_b.send_and_wait.return_value = _assistant_reply("FRESH") + client.create_session.side_effect = [session_a, session_b] + target = GitHubCopilotTarget(model_name="gpt-5-mini") + + conversation_a = "retired-conversation" + with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as first_error: + await _send_normalized_async( + target=target, + original_value="ambiguous original", + converted_value="ambiguous prepared", + conversation_id=conversation_a, + ) + assert isinstance(first_error.value.__cause__, TimeoutError) + assert str(first_error.value.__cause__) == "ambiguous mock send" + + with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as retired_error: + await _send_normalized_async( + target=target, + original_value="retry original", + converted_value="retry prepared", + conversation_id=conversation_a, + ) + assert isinstance(retired_error.value.__cause__, RuntimeError) + assert "retired" in str(retired_error.value.__cause__).lower() + + conversation_b = "fresh-conversation" + response_b = await _send_normalized_async( + target=target, + original_value="fresh original", + converted_value="fresh prepared", + conversation_id=conversation_b, + ) + assert response_b.get_piece().converted_value == "FRESH" + assert _message_roles_and_errors(memory=sqlite_instance, conversation_id=conversation_a) == [ + ("user", "none"), + ("assistant", "processing"), + ("user", "none"), + ("assistant", "processing"), + ] + assert _message_values_and_errors(memory=sqlite_instance, conversation_id=conversation_b) == [ + ("user", "fresh prepared", "none"), + ("assistant", "FRESH", "none"), + ] + + sdk.CopilotClient.assert_called_once_with(github_token=None, working_directory=None) + client.start.assert_awaited_once() + client.get_status.assert_awaited_once() + assert client.create_session.await_count == 2 + requested_session_ids = [entry.kwargs["session_id"] for entry in client.create_session.await_args_list] + assert all(str(UUID(session_id)) == session_id for session_id in requested_session_ids) + assert len(set(requested_session_ids)) == 2 + session_a.send_and_wait.assert_awaited_once_with("ambiguous prepared", timeout=60.0) + session_b.send_and_wait.assert_awaited_once_with("fresh prepared", timeout=60.0) + assert client.delete_session.await_args_list == [call("sdk-session-a")] + client.stop.assert_not_awaited() + await target.cleanup_target_async() + assert client.delete_session.await_args_list == [call("sdk-session-a"), call("sdk-session-b")] + client.stop.assert_awaited_once() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_cleanup_rejects_queued_turn_and_drains_active_send_async( + *, + client: NonCallableMagicMock, +) -> None: + session = client.create_session.return_value + first_send_started = asyncio.Event() + release_first_send = asyncio.Event() + send_count = 0 + + async def send_and_wait_async(*_args: Any, **_kwargs: Any) -> Any: + nonlocal send_count + send_count += 1 + if send_count == 1: + first_send_started.set() + await release_first_send.wait() + return _assistant_reply("FIRST") + return _assistant_reply("UNEXPECTED_SECOND") + + session.send_and_wait.side_effect = send_and_wait_async + target = GitHubCopilotTarget(model_name="gpt-5-mini") + conversation_id = "queued-cleanup-conversation" + first_task = asyncio.create_task( + target.send_prompt_async( + message=_user_message(conversation_id=conversation_id, original_value="first"), + ) + ) + queued_task: asyncio.Task[list[Message]] | None = None + cleanup_task: asyncio.Task[None] | None = None + try: + await asyncio.wait_for(first_send_started.wait(), timeout=2.0) + queued_task = asyncio.create_task( + target.send_prompt_async( + message=_user_message(conversation_id=conversation_id, original_value="queued"), + ) + ) + await asyncio.sleep(0) + cleanup_task = asyncio.create_task(target.cleanup_target_async()) + await asyncio.sleep(0) + + with pytest.raises(RuntimeError, match="cleaned up"): + await asyncio.wait_for( + target.send_prompt_async( + message=_user_message( + conversation_id="fresh-after-cleanup", + original_value="fresh", + ) + ), + timeout=2.0, + ) + _assert_no_resource_release(client=client) + + release_first_send.set() + first_response = await asyncio.wait_for(first_task, timeout=2.0) + assert first_response[0].get_piece().converted_value == "FIRST" + assert queued_task is not None + with pytest.raises(RuntimeError, match="cleaned up"): + await asyncio.wait_for(queued_task, timeout=2.0) + assert cleanup_task is not None + await asyncio.wait_for(cleanup_task, timeout=2.0) + finally: + release_first_send.set() + await _cancel_tasks_async(first_task, queued_task, cleanup_task) + + assert send_count == 1 + session.send_and_wait.assert_awaited_once_with("first", timeout=60.0) + client.create_session.assert_awaited_once() + client.delete_session.assert_awaited_once_with(session.session_id) + client.stop.assert_awaited_once() @pytest.mark.usefixtures("patch_central_database", "sdk") -def test_target_advertises_single_turn_text_only_capabilities() -> None: +def test_target_advertises_native_text_only_capabilities() -> None: capabilities = GitHubCopilotTarget(model_name="gpt-4o").capabilities - assert capabilities.supports_multi_turn is False - assert capabilities.supports_system_prompt is False + assert capabilities.supports_multi_turn is True + assert capabilities.supports_system_prompt is True assert capabilities.supports_multi_message_pieces is False assert capabilities.input_modalities == frozenset({frozenset({"text"})}) assert capabilities.output_modalities == frozenset({frozenset({"text"})}) +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + ("capture_before", "requested_model", "expected_model", "expected_error"), + [ + pytest.param(True, "gpt-5.4", "gpt-5-mini", RuntimeError, id="different-after-capture"), + pytest.param(False, "gpt-5.4", "gpt-5.4", None, id="different-before-capture"), + pytest.param(True, "gpt-5-mini", "gpt-5-mini", None, id="same-after-capture"), + pytest.param(False, " ", "gpt-5-mini", ValueError, id="blank-before-capture"), + pytest.param(True, " ", "gpt-5-mini", ValueError, id="blank-after-capture"), + ], +) +def test_set_model_name_boundaries( + *, + sdk: Any, + client: NonCallableMagicMock, + capture_before: bool, + requested_model: str, + expected_model: str, + expected_error: type[Exception] | None, +) -> None: + target = GitHubCopilotTarget(model_name="gpt-5-mini") + captured_identifier = target.get_identifier() if capture_before else None + + if expected_error is None: + target.set_model_name(model_name=requested_model) + else: + with pytest.raises(expected_error): + target.set_model_name(model_name=requested_model) + + identity = target.get_identifier() + assert identity.params["model_name"] == expected_model + if captured_identifier is not None: + assert identity == captured_identifier + sdk.CopilotClient.assert_not_called() + client.start.assert_not_awaited() + client.create_session.assert_not_awaited() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_direct_send_captures_model_identity_before_startup_async( + *, + client: NonCallableMagicMock, +) -> None: + startup_entered = asyncio.Event() + release_startup = asyncio.Event() + + async def start_async() -> None: + startup_entered.set() + await release_startup.wait() + + client.start.side_effect = start_async + target = GitHubCopilotTarget(model_name="gpt-5-mini") + first_task = asyncio.create_task( + target.send_prompt_async( + message=_user_message( + conversation_id="direct-model-capture", + original_value="first original", + converted_value="first prepared", + ) + ) + ) + try: + await asyncio.wait_for(startup_entered.wait(), timeout=2.0) + with pytest.raises(RuntimeError, match="new target"): + target.set_model_name(model_name="gpt-5.4") + + release_startup.set() + response = await asyncio.wait_for(first_task, timeout=2.0) + assert response[0].get_piece().converted_value == "HELLO" + assert client.create_session.await_args.kwargs["model"] == "gpt-5-mini" + with pytest.raises(RuntimeError, match="new target"): + target.set_model_name(model_name="gpt-5.4") + assert target.get_identifier().params["model_name"] == "gpt-5-mini" + await target.cleanup_target_async() + finally: + release_startup.set() + await _cancel_tasks_async(first_task) + with suppress(Exception): + await asyncio.wait_for(target.cleanup_target_async(), timeout=2.0) + + @pytest.mark.usefixtures("patch_central_database") @pytest.mark.parametrize("failure_stage", ["start", "status", "send", "stop"]) async def test_normalizer_surfaces_lifecycle_failures_async( @@ -207,32 +899,55 @@ async def test_normalizer_surfaces_lifecycle_failures_async( if failure_stage == "stop" else RuntimeError(f"Synthetic {failure_stage} failure") ) + target = GitHubCopilotTarget(model_name="gpt-5-mini") + conversation_id = str(uuid4()) + + if failure_stage == "stop": + response = await _send_normalized_async( + target=target, + original_value="Reply exactly HELLO.", + conversation_id=conversation_id, + ) + assert response.get_piece().response_error == "none" + client.stop.side_effect = error + with pytest.raises(Exception) as exc_info: + await target.cleanup_target_async() + assert exc_info.value is error + client.start.assert_awaited_once() + client.get_status.assert_awaited_once() + client.create_session.assert_awaited_once() + session.send_and_wait.assert_awaited_once() + client.delete_session.assert_awaited_once_with("sdk-session-id") + client.stop.assert_awaited_once() + assert _message_roles_and_errors(memory=sqlite_instance, conversation_id=conversation_id) == [ + ("user", "none"), + ("assistant", "none"), + ] + return + operation = { "start": client.start, "status": client.get_status, "send": session.send_and_wait, - "stop": client.stop, }[failure_stage] operation.side_effect = error - conversation_id = str(uuid4()) with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as exc_info: - await PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + await _send_normalized_async( + target=target, + original_value="Reply exactly HELLO.", conversation_id=conversation_id, - target=GitHubCopilotTarget(model_name="gpt-5-mini"), ) assert exc_info.value.__cause__ is error - stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) - assert [(m.get_piece().role, m.get_piece().response_error) for m in stored] == [ + assert _message_roles_and_errors(memory=sqlite_instance, conversation_id=conversation_id) == [ ("user", "none"), ("assistant", "processing"), ] client.start.assert_awaited_once() - client.stop.assert_awaited_once() client.get_session_metadata.assert_not_awaited() if failure_stage == "start": client.get_status.assert_not_awaited() + client.stop.assert_awaited_once() else: client.get_status.assert_awaited_once() if failure_stage in ("start", "status"): @@ -244,7 +959,11 @@ async def test_normalizer_surfaces_lifecycle_failures_async( session.send_and_wait.assert_awaited_once() client.delete_session.assert_awaited_once_with("sdk-session-id") session.on.return_value.assert_called_once_with() + client.stop.assert_not_awaited() session.send.assert_not_awaited() + await target.cleanup_target_async() + if failure_stage == "send": + client.stop.assert_awaited_once() @pytest.mark.usefixtures("patch_central_database") @@ -254,6 +973,7 @@ async def test_normalizer_surfaces_dispatch_timeout_and_cleans_up_without_replay client: NonCallableMagicMock, ) -> None: session = client.create_session.return_value + target = GitHubCopilotTarget(model_name="gpt-5-mini", response_timeout_seconds=0.01) async def stall_send_async(*_args: Any, **_kwargs: Any) -> None: await asyncio.Event().wait() @@ -263,15 +983,14 @@ async def stall_send_async(*_args: Any, **_kwargs: Any) -> None: with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as exc_info: # A bare watchdog TimeoutError must not satisfy the normalizer-wrapped failure. await asyncio.wait_for( - PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), - target=GitHubCopilotTarget(model_name="gpt-5-mini", response_timeout_seconds=0.01), - ), + _send_normalized_async(target=target, original_value="Reply exactly HELLO."), timeout=2.0, ) assert isinstance(exc_info.value.__cause__, TimeoutError) session.send.assert_awaited_once() client.delete_session.assert_awaited_once_with("sdk-session-id") + client.stop.assert_not_awaited() + await target.cleanup_target_async() client.stop.assert_awaited_once() @@ -294,21 +1013,23 @@ async def test_normalizer_rejects_invalid_reply_async( "non-assistant": replace(reply, type=SessionEventType.SESSION_IDLE, data=SessionIdleData(aborted=False)), }[invalid_reply] conversation_id = str(uuid4()) + target = GitHubCopilotTarget(model_name="gpt-5-mini") with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as exc_info: - await PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + await _send_normalized_async( + target=target, + original_value="Reply exactly HELLO.", conversation_id=conversation_id, - target=GitHubCopilotTarget(model_name="gpt-5-mini"), ) assert isinstance(exc_info.value.__cause__, ValueError) - stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) - assert [(m.get_piece().role, m.get_piece().response_error) for m in stored] == [ + assert _message_roles_and_errors(memory=sqlite_instance, conversation_id=conversation_id) == [ ("user", "none"), ("assistant", "processing"), ] session.send_and_wait.assert_awaited_once() session.on.return_value.assert_called_once_with() client.delete_session.assert_awaited_once_with("sdk-session-id") + client.stop.assert_not_awaited() + await target.cleanup_target_async() client.stop.assert_awaited_once() @@ -374,21 +1095,23 @@ async def send_events_async(*_args: Any, **_kwargs: Any) -> str: session.send.side_effect = send_events_async session.send_and_wait.side_effect = partial(sdk.CopilotSession.send_and_wait, session) conversation_id = str(uuid4()) + target = GitHubCopilotTarget(model_name="gpt-5-mini", response_timeout_seconds=1.0) with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as exc_info: - await PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + await _send_normalized_async( + target=target, + original_value="Reply exactly HELLO.", conversation_id=conversation_id, - target=GitHubCopilotTarget(model_name="gpt-5-mini", response_timeout_seconds=1.0), ) assert isinstance(exc_info.value.__cause__, RuntimeError) assert expected_error in str(exc_info.value.__cause__).lower() - stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) - assert [(m.get_piece().role, m.get_piece().response_error) for m in stored] == [ + assert _message_roles_and_errors(memory=sqlite_instance, conversation_id=conversation_id) == [ ("user", "none"), ("assistant", "processing"), ] session.send.assert_awaited_once() client.delete_session.assert_awaited_once_with("sdk-session-id") + client.stop.assert_not_awaited() + await target.cleanup_target_async() client.stop.assert_awaited_once() assert not handlers @@ -416,12 +1139,13 @@ async def create_session_async(*, session_id: str = "sdk-generated-session-id", client.create_session.side_effect = create_session_async _mock_session_storage(client=client, sessions=sessions) conversation_id = str(uuid4()) + target = GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=retain_session) with caplog.at_level(logging.INFO, logger=TARGET_LOGGER): with pytest.raises(Exception, match="Error sending prompt with conversation ID:") as exc_info: - await PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + await _send_normalized_async( + target=target, + original_value="Reply exactly HELLO.", conversation_id=conversation_id, - target=GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=retain_session), ) assert exc_info.value.__cause__ is creation_error @@ -432,9 +1156,10 @@ async def create_session_async(*, session_id: str = "sdk-generated-session-id", client.get_session_metadata.assert_awaited_once_with(allocated_session_id) session.send_and_wait.assert_not_awaited() session.send.assert_not_awaited() + client.stop.assert_not_awaited() + await target.cleanup_target_async() client.stop.assert_awaited_once() - stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) - assert [(m.get_piece().role, m.get_piece().response_error) for m in stored] == [ + assert _message_roles_and_errors(memory=sqlite_instance, conversation_id=conversation_id) == [ ("user", "none"), ("assistant", "processing"), ] @@ -470,21 +1195,15 @@ async def create_session_async(*, session_id: str = "sdk-generated-session-id", client.create_session.side_effect = create_session_async _mock_session_storage(client=client, sessions=sessions) - request_task = asyncio.create_task( - PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), - target=GitHubCopilotTarget(model_name="gpt-5-mini"), - ) - ) + target = GitHubCopilotTarget(model_name="gpt-5-mini") + request_task = asyncio.create_task(_send_normalized_async(target=target, original_value="Reply exactly HELLO.")) try: await asyncio.wait_for(allocated.wait(), timeout=2.0) request_task.cancel() with pytest.raises(asyncio.CancelledError): await request_task finally: - if not request_task.done(): - request_task.cancel() - await asyncio.gather(request_task, return_exceptions=True) + await _cancel_tasks_async(request_task) client.create_session.assert_awaited_once() allocated_id = client.create_session.await_args.kwargs["session_id"] @@ -492,6 +1211,8 @@ async def create_session_async(*, session_id: str = "sdk-generated-session-id", client.delete_session.assert_awaited_once_with(allocated_id) session.send_and_wait.assert_not_awaited() session.send.assert_not_awaited() + client.stop.assert_not_awaited() + await target.cleanup_target_async() client.stop.assert_awaited_once() assert sessions == {"unrelated-session-id"} @@ -522,21 +1243,15 @@ async def stop_async() -> None: client.start.side_effect = start_async client.stop.side_effect = stop_async - request_task = asyncio.create_task( - PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), - target=GitHubCopilotTarget(model_name="gpt-5-mini"), - ) - ) + target = GitHubCopilotTarget(model_name="gpt-5-mini") + request_task = asyncio.create_task(_send_normalized_async(target=target, original_value="Reply exactly HELLO.")) try: await asyncio.wait_for(allocated.wait(), timeout=2.0) request_task.cancel() with pytest.raises(asyncio.CancelledError) as exc_info: await request_task finally: - if not request_task.done(): - request_task.cancel() - await asyncio.gather(request_task, return_exceptions=True) + await _cancel_tasks_async(request_task) assert exc_info.value is original_cancellation sdk.CopilotClient.assert_called_once() @@ -549,6 +1264,7 @@ async def stop_async() -> None: session.send_and_wait.assert_not_awaited() client.stop.assert_awaited_once() assert owned_resources == set() + await target.cleanup_target_async() @pytest.mark.usefixtures("patch_central_database") @@ -581,9 +1297,8 @@ async def test_normalizer_forwards_options_without_exposing_token_async( working_directory=tmp_path if use_working_directory else None, max_requests_per_minute=max_requests_per_minute, ) - response = await PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), target=target - ) + response = await _send_normalized_async(target=target, original_value="Reply exactly HELLO.") + await target.cleanup_target_async() sdk.CopilotClient.assert_called_once_with( github_token=github_token, working_directory=str(tmp_path) if use_working_directory else None ) @@ -658,11 +1373,12 @@ async def release_constructor_async() -> None: sdk.CopilotClient.side_effect = construct_client conversation_id = str(uuid4()) + target = GitHubCopilotTarget(model_name="gpt-5-mini") request_task = asyncio.create_task( - PromptNormalizer().send_prompt_async( - message=MessagePiece(role="user", original_value="Reply exactly HELLO.").to_message(), + _send_normalized_async( + target=target, + original_value="Reply exactly HELLO.", conversation_id=conversation_id, - target=GitHubCopilotTarget(model_name="gpt-5-mini"), ) ) release_task = asyncio.create_task(release_constructor_async()) @@ -670,10 +1386,7 @@ async def release_constructor_async() -> None: response, _ = await asyncio.wait_for(asyncio.gather(request_task, release_task), timeout=10.0) finally: release.set() - for task in (request_task, release_task): - if not task.done(): - task.cancel() - await asyncio.gather(request_task, release_task, return_exceptions=True) + await _cancel_tasks_async(request_task, release_task) await asyncio.wait_for(constructor_finished.wait(), timeout=5.0) assert isinstance(response, Message) @@ -684,14 +1397,16 @@ async def release_constructor_async() -> None: conversation_id, "none", ) - stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) - assert [(m.get_piece().role, m.get_piece().converted_value, m.get_piece().response_error) for m in stored] == [ + assert _message_values_and_errors(memory=sqlite_instance, conversation_id=conversation_id) == [ ("user", "Reply exactly HELLO.", "none"), ("assistant", "HELLO", "none"), ] sdk.CopilotClient.assert_called_once_with(github_token=None, working_directory=None) client.start.assert_awaited_once() client.create_session.return_value.send_and_wait.assert_awaited_once_with("Reply exactly HELLO.", timeout=60.0) + client.stop.assert_not_awaited() + client.delete_session.assert_not_awaited() + await target.cleanup_target_async() client.stop.assert_awaited_once() client.delete_session.assert_awaited_once_with("sdk-session-id") assert released_while_constructing is True diff --git a/tests/unit/prompt_target/target/test_prompt_target.py b/tests/unit/prompt_target/target/test_prompt_target.py index 453870e382..6abaf85b4e 100644 --- a/tests/unit/prompt_target/target/test_prompt_target.py +++ b/tests/unit/prompt_target/target/test_prompt_target.py @@ -12,7 +12,14 @@ from pyrit.executor.attack.core.attack_strategy import AttackStrategy from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import ChatMessageRole, ComponentIdentifier, Message, MessagePiece, flatten_to_message_pieces +from pyrit.models import ( + ChatMessageRole, + ComponentIdentifier, + Conversation, + Message, + MessagePiece, + flatten_to_message_pieces, +) from pyrit.prompt_target import OpenAIChatTarget from pyrit.prompt_target.common.target_capabilities import ( CapabilityHandlingPolicy, @@ -81,6 +88,109 @@ async def test_set_system_prompt_adds_memory( assert chats[0].api_role == "system" +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + ("supports_multi_turn", "supports_editable_history", "supports_system_prompt", "policy", "expected_error"), + [ + pytest.param(True, False, True, None, None, id="native-system-without-editable-history"), + pytest.param( + True, + True, + False, + CapabilityHandlingPolicy(behaviors={CapabilityName.SYSTEM_PROMPT: UnsupportedCapabilityBehavior.ADAPT}), + None, + id="editable-history-without-native-system", + ), + pytest.param(False, True, True, None, ValueError, id="without-multi-turn-editable"), + pytest.param(False, False, True, None, ValueError, id="without-multi-turn-native-system"), + pytest.param(True, False, False, None, ValueError, id="without-editable-or-native-system"), + ], +) +def test_set_system_prompt_capability_admission_and_nonmutation( + *, + sqlite_instance: MemoryInterface, + supports_multi_turn: bool, + supports_editable_history: bool, + supports_system_prompt: bool, + policy: CapabilityHandlingPolicy | None, + expected_error: type[ValueError] | None, +) -> None: + conversation_id = "system-prompt-capability-conversation" + target = _make_identifier_target( + capabilities=TargetCapabilities( + supports_multi_turn=supports_multi_turn, + supports_editable_history=supports_editable_history, + supports_system_prompt=supports_system_prompt, + ), + policy=policy, + ) + + if expected_error is None: + target.set_system_prompt(system_prompt="be concise", conversation_id=conversation_id) + stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) + assert len(stored) == 1 + piece = stored[0].get_piece() + assert piece.api_role == "system" + assert piece.converted_value == "be concise" + assert piece.conversation_id == conversation_id + else: + with pytest.raises( + expected_error, + match="It must support multi-turn conversations and either editable history or native system prompts.", + ): + target.set_system_prompt(system_prompt="be concise", conversation_id=conversation_id) + assert sqlite_instance.get_conversation_messages(conversation_id=conversation_id) == [] + assert sqlite_instance.get_target_identifiers(identifier_hashes=[target.get_identifier().hash]) == [] + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + ("existing_role", "existing_content"), + [ + pytest.param("system", "be concise", id="repeated-system-prompt"), + pytest.param("user", "existing user message", id="existing-user-message"), + ], +) +def test_set_system_prompt_rejects_nonempty_conversation_without_mutation( + *, + sqlite_instance: MemoryInterface, + existing_role: str, + existing_content: str, +) -> None: + conversation_id = "nonempty-system-prompt-conversation" + target = _make_identifier_target( + capabilities=TargetCapabilities( + supports_multi_turn=True, + supports_editable_history=False, + supports_system_prompt=True, + ) + ) + if existing_role == "system": + target.set_system_prompt(system_prompt=existing_content, conversation_id=conversation_id) + else: + sqlite_instance.add_conversation_to_memory( + conversation=Conversation(conversation_id=conversation_id, target_identifier=target.get_identifier()) + ) + sqlite_instance.add_message_to_memory( + request=MessagePiece( + role="user", + conversation_id=conversation_id, + original_value=existing_content, + converted_value=existing_content, + ).to_message() + ) + + with pytest.raises(RuntimeError, match="Conversation already exists"): + target.set_system_prompt(system_prompt="be expansive", conversation_id=conversation_id) + + stored = sqlite_instance.get_conversation_messages(conversation_id=conversation_id) + assert len(stored) == 1 + piece = stored[0].get_piece() + assert piece.api_role == existing_role + assert piece.converted_value == existing_content + assert piece.conversation_id == conversation_id + + async def test_send_prompt_with_system_calls_chat_complete( azure_openai_target: OpenAIChatTarget, openai_response_json: dict, From c78d3d4f8bbd094b9b61441fe10a2ebaacf251fd Mon Sep 17 00:00:00 2001 From: Adrian Gavrila Date: Thu, 24 Sep 2026 10:49:30 -0400 Subject: [PATCH 4/7] Support Copilot-backed self-ask judging Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/targets/github_copilot_target.md | 128 ------- doc/myst.yml | 1 - pyrit/prompt_target/github_copilot_target.py | 48 ++- pyrit/score/llm_scoring.py | 64 +++- pyrit/score/scorer.py | 16 +- .../self_ask_question_answer_scorer.py | 2 + .../true_false/self_ask_refusal_scorer.py | 11 +- .../true_false/self_ask_true_false_scorer.py | 21 +- .../target/test_github_copilot_target.py | 332 +++++++++++++++++- tests/unit/score/test_self_ask_refusal.py | 124 ++++++- 10 files changed, 564 insertions(+), 183 deletions(-) delete mode 100644 doc/code/targets/github_copilot_target.md diff --git a/doc/code/targets/github_copilot_target.md b/doc/code/targets/github_copilot_target.md deleted file mode 100644 index ca4f43aaa7..0000000000 --- a/doc/code/targets/github_copilot_target.md +++ /dev/null @@ -1,128 +0,0 @@ -# GitHub Copilot SDK target - -`GitHubCopilotTarget` sends text conversations through the GitHub Copilot SDK, retaining native -conversation state across turns. -Use Python 3.11 through 3.14 and install the optional extra: - -```console -python -m pip install "pyrit[github-copilot]" -``` - -Authenticate beforehand with an eligible Copilot login, or supply an SDK environment token. -The SDK checks `COPILOT_GITHUB_TOKEN`, then `GH_TOKEN`, then `GITHUB_TOKEN`; environment tokens -take precedence over saved login. Alternatively, pass a securely supplied token as -`github_token=token`, which takes precedence over discovery. Do not hard-code credentials. -The target does not initiate interactive sign-in. See [PyRIT setup](../setup/0_setup.md) -for framework initialization. - -## Two-turn native conversation - -Run this standalone script from your chosen existing local directory; no Git checkout is needed. -`Path.cwd()` explicitly selects that directory, and its resolved path is included in the target -identity. Running the script makes **two real Copilot model requests**. - -```python -import asyncio -import logging -from pathlib import Path -from uuid import uuid4 - -from pyrit.models import Message -from pyrit.prompt_normalizer import PromptNormalizer -from pyrit.prompt_target import GitHubCopilotTarget -from pyrit.setup import IN_MEMORY, initialize_pyrit_async - - -async def main_async() -> None: - await initialize_pyrit_async(memory_db_type=IN_MEMORY, silent=True) - - copilot_logger = logging.getLogger("pyrit.prompt_target.github_copilot_target") - copilot_logger.setLevel(logging.INFO) - copilot_logger.addHandler(logging.StreamHandler()) - copilot_logger.propagate = False - - target = GitHubCopilotTarget( - model_name="gpt-5-mini", - working_directory=Path.cwd(), - ) - normalizer = PromptNormalizer() - try: - conversation_id = str(uuid4()) - target.set_system_prompt(system_prompt="Answer concisely.", conversation_id=conversation_id) - - for prompt in ( - "Remember the codeword ORCHID for this conversation.", - "What codeword did I ask you to remember?", - ): - response = await normalizer.send_prompt_async( - message=Message.from_prompt(prompt=prompt, role="user"), - conversation_id=conversation_id, - target=target, - ) - print(response.get_piece().converted_value) - finally: - await target.cleanup_target_async() - - -if __name__ == "__main__": - asyncio.run(main_async()) -``` - -The normalizer owns conversion and request/response persistence. Here, [IN_MEMORY](../memory/0_memory.md) -keeps PyRIT records only for the current process, independently of Copilot's local session data. -The target lazily shares one client across isolated native conversations, with one native session per -PyRIT conversation; this shares SDK runtime/authentication, not process or security isolation. Keep -the target alive for all sends and join workflow tasks before calling `cleanup_target_async()`. - -## Retention and boundaries - -`cleanup_target_async()` is caller-owned, terminal cleanup. It drains active target work, rejects -new or queued target work, releases owned resources, and surfaces cleanup failures. It does not -gather caller tasks that may themselves be waiting on cleanup. By default, `retain_session=False` -deletes owned native session data. With `retain_session=True`, native session data is retained -but live session resources are still disconnected and the shared client is stopped. PyRIT memory -records remain independently retained. Crashes can still leave residual data. Remote session export -is explicitly `OFF` in either mode: local retention does not enable export, and `OFF` does not mean -offline inference. - -`reset_conversation_async()` releases one established conversation without stopping unrelated -conversations or the shared client. A reset conversation is terminal: later sends with that -conversation ID fail rather than silently reopening native history. Unknown IDs and repeated -resets are safe no-ops. To start a new conversation, use a new PyRIT conversation ID; this target -does not resume, import, replay, fork, or reconcile native history. Edits to PyRIT memory do not -update the native Copilot context. - -Capture the target's INFO logs for session-creation attempt records linking PyRIT conversation IDs -to requested SDK session IDs, with SDK/runtime versions, protocol, retention and remote mode. These -records are emitted before each native session creation attempt, not on every turn, and are not -automatically added to response metadata or target-identity exports. An attempt record does not -prove session allocation or a successful exchange; retained-session logs identify data kept for -diagnostics. - -Each turn sends only the newest normalized user content; prior turns remain in the native session. -The target supports native continuation for text conversations and does not support editable history. - -If a caller supplies an initial system prompt before the first user turn, the target uses SDK -replacement mode with that exact text. Replacement replaces the SDK's default system message and -its guardrails/security restrictions. The target still explicitly requests remote `OFF`, an empty -tool list, disabled configuration and instruction discovery, disabled host Git operations, session -store, memory, and file hooks. Without a supplied system prompt, the target keeps the existing -customize mode and removes only `environment_context` and `custom_instructions`; other SDK/runtime -context may remain. These controls are not an OS sandbox. Later public system-prompt changes are -rejected. - -Finalize `model_name` before the first identity-capturing operation: `get_identifier()`, system -prompt setup, normalizer registration, scenario or registry use, target mapping, or direct native -execution. After identity capture, a different model requires a new target; setting the same model -again is harmless. Native SDK model switching is not used. - -`response_timeout_seconds` bounds dispatch and completion, excluding client/session startup, the -startup status lookup, PyRIT pacing and cleanup. The synchronous SDK client constructor is moved -off the event loop, but cancellation cannot abort synchronous work already running in its worker. -A failed status lookup stops the exchange before session creation or prompt dispatch. - -SDK-observed ambiguous timeout, cancellation, abort, unsafe event, or invalid native outcomes retire -only the affected conversation. Later sends to that ID fail, while a fresh ID can proceed; there is -no automatic replay. Errors, timeouts, and cancellation propagate without a crash-proof cleanup -guarantee. Downstream converter or PyRIT persistence failures occur outside the target's native -observation boundary and do not provide a safe-continuation guarantee. diff --git a/doc/myst.yml b/doc/myst.yml index 90c1d3b22b..78bc9147eb 100644 --- a/doc/myst.yml +++ b/doc/myst.yml @@ -137,7 +137,6 @@ project: - file: code/targets/use_huggingface_chat_target.ipynb - file: code/targets/websocket_target.ipynb - file: code/targets/round_robin_target.ipynb - - file: code/targets/github_copilot_target.md - file: code/converters/0_converters.ipynb children: - file: code/converters/1_text_to_text_converters.ipynb diff --git a/pyrit/prompt_target/github_copilot_target.py b/pyrit/prompt_target/github_copilot_target.py index a5f70c1410..2512601664 100644 --- a/pyrit/prompt_target/github_copilot_target.py +++ b/pyrit/prompt_target/github_copilot_target.py @@ -102,12 +102,10 @@ def __init__( self._client: CopilotClient | None = None self._runtime_status: GetStatusResponse | None = None self._client_start_lock = asyncio.Lock() - self._lifecycle_lock = asyncio.Lock() + self._lifecycle_condition = asyncio.Condition() self._cleanup_task: asyncio.Task[None] | None = None self._conversations: dict[str, _ConversationState] = {} self._active_target_operations = 0 - self._active_operations_drained = asyncio.Event() - self._active_operations_drained.set() def _build_identifier(self) -> ComponentIdentifier: """ @@ -138,17 +136,16 @@ def set_model_name(self, *, model_name: str) -> None: @limit_requests_per_minute async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Message]) -> list[Message]: self.get_identifier() - async with self._lifecycle_lock: + async with self._lifecycle_condition: if self._cleanup_task is not None: raise RuntimeError("GitHubCopilotTarget has been cleaned up and cannot send more prompts.") self._active_target_operations += 1 - self._active_operations_drained.clear() try: request = normalized_conversation[-1].get_piece() conversation_id = request.conversation_id or "" conversation = self._conversations.setdefault(conversation_id, _ConversationState()) async with conversation.lock: - async with self._lifecycle_lock: + async with self._lifecycle_condition: if self._cleanup_task is not None: raise RuntimeError("GitHubCopilotTarget has been cleaned up and cannot send more prompts.") initial_system_prompt: str | None = None @@ -166,18 +163,19 @@ async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Me prompt=request.converted_value, ) except BaseException: - await asyncio.shield( - self._retire_conversation_async( - conversation=conversation, - ) - ) + retirement_task = asyncio.create_task(self._retire_conversation_async(conversation=conversation)) + try: + await asyncio.shield(retirement_task) + except asyncio.CancelledError: + await retirement_task + raise raise return [construct_response_from_request(request=request, response_text_pieces=[reply_text])] finally: - async with self._lifecycle_lock: + async with self._lifecycle_condition: self._active_target_operations -= 1 if self._active_target_operations == 0: - self._active_operations_drained.set() + self._lifecycle_condition.notify_all() async def _send_text_async(self, *, session: "CopilotSession", prompt: str) -> str: from copilot.generated.session_events import ( @@ -228,7 +226,7 @@ async def cleanup_target_async(self) -> None: Cleanup is terminal and idempotent. Retained sessions are preserved while the shared client is always stopped. Cleanup attempts every owned resource before surfacing failures. """ - async with self._lifecycle_lock: + async with self._lifecycle_condition: if self._cleanup_task is None: self._cleanup_task = asyncio.create_task(self._cleanup_owned_resources_async()) cleanup_task = self._cleanup_task @@ -245,14 +243,13 @@ async def reset_conversation_async(self, *, conversation_id: str) -> None: Args: conversation_id (str): The PyRIT conversation ID to release. """ - async with self._lifecycle_lock: + async with self._lifecycle_condition: conversation = self._conversations.get(conversation_id) if conversation is None or (conversation.retired and conversation.session is None): return cleanup_task = self._cleanup_task if cleanup_task is None: self._active_target_operations += 1 - self._active_operations_drained.clear() if cleanup_task is not None: await asyncio.shield(cleanup_task) @@ -262,10 +259,10 @@ async def reset_conversation_async(self, *, conversation_id: str) -> None: async with conversation.lock: await self._retire_conversation_async(conversation=conversation) finally: - async with self._lifecycle_lock: + async with self._lifecycle_condition: self._active_target_operations -= 1 if self._active_target_operations == 0: - self._active_operations_drained.set() + self._lifecycle_condition.notify_all() async def _get_or_create_session_async( self, @@ -273,7 +270,7 @@ async def _get_or_create_session_async( conversation_id: str, initial_system_prompt: str | None, ) -> "CopilotSession": - async with self._lifecycle_lock: + async with self._lifecycle_condition: conversation = self._conversations[conversation_id] if conversation.retired: raise RuntimeError( @@ -329,7 +326,7 @@ async def _get_or_create_session_async( enable_session_store=False, enable_file_hooks=False, ) - async with self._lifecycle_lock: + async with self._lifecycle_condition: conversation.session = session return session except BaseException: @@ -361,12 +358,12 @@ async def _get_or_start_client_async(self) -> "CopilotClient": ) from error raise - async with self._lifecycle_lock: + async with self._lifecycle_condition: self._client = client return client async def _retire_conversation_async(self, *, conversation: _ConversationState) -> None: - async with self._lifecycle_lock: + async with self._lifecycle_condition: session = conversation.session client = self._client if session is None or client is None: @@ -397,13 +394,14 @@ async def _release_session_async( logger.info("Retaining Copilot session %s as requested; delete it manually.", session_id) else: await client.delete_session(session_id) - async with self._lifecycle_lock: + async with self._lifecycle_condition: if conversation.session is session: conversation.session = None + conversation.retired = True async def _cleanup_owned_resources_async(self) -> None: - await self._active_operations_drained.wait() - async with self._lifecycle_lock: + async with self._lifecycle_condition: + await self._lifecycle_condition.wait_for(lambda: self._active_target_operations == 0) conversations = [ conversation for conversation in self._conversations.values() if conversation.session is not None ] diff --git a/pyrit/score/llm_scoring.py b/pyrit/score/llm_scoring.py index 6881adc179..3f9654ba16 100644 --- a/pyrit/score/llm_scoring.py +++ b/pyrit/score/llm_scoring.py @@ -10,6 +10,7 @@ EmptyResponseException, InvalidJsonException, ScorerLLMResponseBlockedException, + pyrit_json_retry, ) from pyrit.models import Message, MessagePiece from pyrit.prompt_normalizer import PromptNormalizer, send_json_with_retry_async @@ -39,6 +40,7 @@ async def _run_llm_scoring_async( category: Sequence[str] | str | None = None, objective: str | None = None, normalizer: PromptNormalizer | None = None, + fresh_conversation_per_attempt: bool = False, ) -> UnvalidatedScore: """ Perform a single scoring round-trip against an LLM target and delegate parsing. @@ -46,13 +48,13 @@ async def _run_llm_scoring_async( This is the shared LLM evaluation mechanism: it optionally sets a system prompt on the target, sends the value to be scored (forwarding ``response_handler.json_response_config`` so targets that support structured output can enforce it), and delegates parsing and validation to - ``response_handler``. The round-trip is routed through a ``PromptNormalizer`` via - ``send_json_with_retry_async`` so the scorer's question and the target's answer are persisted - to memory (a full audit trail, and a real conversation an attack can link as a SCORE-type - related conversation) and so JSON retries roll memory back to a clean baseline between attempts - instead of replaying the target's own malformed reply. It is intentionally stateless and - independent of any particular ``Scorer`` so that scorers can compose it without inheriting LLM - machinery. + ``response_handler``. The round-trip uses a ``PromptNormalizer`` so the scorer's question and + the target's answer are persisted to memory (a full audit trail, and a real conversation an + attack can link as a SCORE-type related conversation). The default editable-history path rolls + memory back between JSON attempts; ``fresh_conversation_per_attempt`` keeps malformed judge + exchanges in separate conversations instead of replaying native history. It is intentionally + stateless and independent of any particular ``Scorer`` so scorers can compose it without + inheriting LLM machinery. The round-trip owns only the transport; the ``ResponseHandler`` owns the response contract — the optional response schema and turning raw text into a validated ``UnvalidatedScore``. @@ -79,9 +81,10 @@ async def _run_llm_scoring_async( from the response; supplying both is an error. Defaults to None. objective (str | None): The objective associated with the score, used for contextualizing the result. Defaults to None. - normalizer (PromptNormalizer | None): Normalizer used to send the scoring round-trip - and whose memory is rolled back between JSON retries. Injectable for testing; - defaults to a fresh ``PromptNormalizer()`` when not supplied. + normalizer (PromptNormalizer | None): Normalizer used to send the scoring round-trip. + Injectable for testing; defaults to a fresh ``PromptNormalizer()`` when not supplied. + fresh_conversation_per_attempt (bool): Use a fresh conversation for each attempt when + rolling back target history is not supported. Defaults to False. Returns: UnvalidatedScore: The parsed score, whose ``raw_score_value`` still needs to be @@ -103,7 +106,7 @@ async def _run_llm_scoring_async( """ conversation_id = str(uuid.uuid4()) - if system_prompt is not None: + if system_prompt is not None and not fresh_conversation_per_attempt: chat_target.set_system_prompt( system_prompt=system_prompt, conversation_id=conversation_id, @@ -177,9 +180,44 @@ def _parse(response: Message) -> UnvalidatedScore: objective=objective, ) - # Route the round-trip through the normalizer so the scorer Q&A is persisted and JSON retries - # replay on a clean history. + # Editable targets retry on a rolled-back conversation; non-editable judges opt into fresh sessions. try: + if fresh_conversation_per_attempt: + active_normalizer = normalizer or PromptNormalizer() + first_attempt = True + + @pyrit_json_retry + async def _fresh_attempt_async() -> UnvalidatedScore: + nonlocal first_attempt + attempt_conversation_id = conversation_id if first_attempt else str(uuid.uuid4()) + attempt_message = scorer_llm_request if first_attempt else scorer_llm_request.duplicate() + if system_prompt is not None: + chat_target.set_system_prompt( + system_prompt=system_prompt, + conversation_id=attempt_conversation_id, + ) + first_attempt = False + response = await active_normalizer.send_prompt_async( + message=attempt_message, + conversation_id=attempt_conversation_id, + target=chat_target, + ) + if not response: + raise ValueError(f"No response received for conversation ID: {attempt_conversation_id}") + try: + return _parse(response) + except InvalidJsonException: + try: + await chat_target.reset_conversation_async(conversation_id=attempt_conversation_id) + except Exception as reset_error: + raise RuntimeError( + "Could not release the malformed judge session; refusing another retry." + ) from reset_error + raise + + fresh_score: UnvalidatedScore = await _fresh_attempt_async() + return fresh_score + return await send_json_with_retry_async( normalizer=normalizer or PromptNormalizer(), target=chat_target, diff --git a/pyrit/score/scorer.py b/pyrit/score/scorer.py index 3215052b37..0280dd925c 100644 --- a/pyrit/score/scorer.py +++ b/pyrit/score/scorer.py @@ -6,6 +6,7 @@ import abc import logging from abc import abstractmethod +from dataclasses import replace from typing import TYPE_CHECKING, Any, ClassVar, cast from pyrit.common.deprecation import print_deprecation_message @@ -29,7 +30,8 @@ ScoringExpectation, ) from pyrit.prompt_target.batch_helper import batch_task_async -from pyrit.prompt_target.common.target_requirements import TargetRequirements +from pyrit.prompt_target.common.target_capabilities import CapabilityName +from pyrit.prompt_target.common.target_requirements import CHAT_TARGET_REQUIREMENTS, TargetRequirements if TYPE_CHECKING: import uuid @@ -48,6 +50,18 @@ LEGACY_SCORE_ASYNC_REMOVED_IN = "2.0.0" +class _SelfContainedJudgeTargetRequirements(TargetRequirements): + def validate(self, *, target: PromptTarget) -> None: + requirements = CHAT_TARGET_REQUIREMENTS + if not target.capabilities.supports_editable_history: + requirements = replace( + requirements, + required=requirements.required - {CapabilityName.EDITABLE_HISTORY}, + native_required=requirements.native_required | {CapabilityName.SYSTEM_PROMPT}, + ) + requirements.validate(target=target) + + async def _legacy_score_scorable_async( self: Scorer, *, diff --git a/pyrit/score/true_false/self_ask_question_answer_scorer.py b/pyrit/score/true_false/self_ask_question_answer_scorer.py index 6aa07d14d7..b3ddac62ec 100644 --- a/pyrit/score/true_false/self_ask_question_answer_scorer.py +++ b/pyrit/score/true_false/self_ask_question_answer_scorer.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING from pyrit.common.path import SCORER_SEED_PROMPT_PATH +from pyrit.prompt_target import CHAT_TARGET_REQUIREMENTS from pyrit.score.llm_scoring import _run_llm_scoring_async from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.self_ask_true_false_scorer import ( @@ -38,6 +39,7 @@ class SelfAskQuestionAnswerScorer(SelfAskTrueFalseScorer): supported_data_types=["text"], is_objective_required=True, ) + TARGET_REQUIREMENTS = CHAT_TARGET_REQUIREMENTS def __init__( self, diff --git a/pyrit/score/true_false/self_ask_refusal_scorer.py b/pyrit/score/true_false/self_ask_refusal_scorer.py index f4d51e80bc..7fe93f7f9b 100644 --- a/pyrit/score/true_false/self_ask_refusal_scorer.py +++ b/pyrit/score/true_false/self_ask_refusal_scorer.py @@ -8,9 +8,10 @@ from pyrit.common.path import SCORER_SEED_PROMPT_PATH from pyrit.models import ComponentIdentifier, JsonSchemaDefinition, Message, MessagePiece, Score, SeedPrompt -from pyrit.prompt_target import CHAT_TARGET_REQUIREMENTS, PromptTarget +from pyrit.prompt_target import PromptTarget from pyrit.score.llm_scoring import _run_llm_scoring_async from pyrit.score.response_handler import JsonSchemaResponseHandler, ResponseHandler, TrueFalseResponseHandler +from pyrit.score.scorer import _SelfContainedJudgeTargetRequirements from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.true_false.true_false_score_aggregator import ( TrueFalseAggregatorFunc, @@ -66,7 +67,7 @@ class SelfAskRefusalScorer(MessageTrueFalseScorer): ) _DEFAULT_VALIDATOR: ScorerPromptValidator = ScorerPromptValidator() - TARGET_REQUIREMENTS = CHAT_TARGET_REQUIREMENTS + TARGET_REQUIREMENTS = _SelfContainedJudgeTargetRequirements() def __init__( self, @@ -83,8 +84,9 @@ def __init__( Initialize the SelfAskRefusalScorer. Args: - chat_target (PromptTarget | None): The chat target used for scoring. Must satisfy - CHAT_TARGET_REQUIREMENTS. + chat_target (PromptTarget | None): The chat target used for scoring. Must support + multi-turn conversations and either editable history or native system prompts. + Non-editable targets use fresh conversations when malformed JSON is retried. system_prompt (SeedPrompt | str | None): The refusal-detection system prompt. A ``SeedPrompt`` (e.g. loaded from a ``RefusalScorerPaths`` YAML) is used verbatim and may carry a ``response_json_schema``; a ``str`` is used as-is; ``None`` falls back to @@ -238,6 +240,7 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st scorer_identifier=self.get_identifier(), category=self._score_category, objective=objective, + fresh_conversation_per_attempt=not self._prompt_target.capabilities.supports_editable_history, ) score = unvalidated_score.to_score(score_value=unvalidated_score.raw_score_value, score_type="true_false") diff --git a/pyrit/score/true_false/self_ask_true_false_scorer.py b/pyrit/score/true_false/self_ask_true_false_scorer.py index b7e6f5e210..a8b44d2324 100644 --- a/pyrit/score/true_false/self_ask_true_false_scorer.py +++ b/pyrit/score/true_false/self_ask_true_false_scorer.py @@ -11,9 +11,10 @@ from pyrit.common import verify_and_resolve_path from pyrit.common.path import SCORER_SEED_PROMPT_PATH from pyrit.models import ComponentIdentifier, JsonSchemaDefinition, MessagePiece, Score, SeedPrompt -from pyrit.prompt_target import CHAT_TARGET_REQUIREMENTS, PromptTarget +from pyrit.prompt_target import PromptTarget from pyrit.score.llm_scoring import _run_llm_scoring_async from pyrit.score.response_handler import JsonSchemaResponseHandler, ResponseHandler, TrueFalseResponseHandler +from pyrit.score.scorer import _SelfContainedJudgeTargetRequirements from pyrit.score.scorer_prompt_validator import ScorerPromptValidator from pyrit.score.system_prompt import _render_system_prompt_template from pyrit.score.true_false.true_false_score_aggregator import ( @@ -143,7 +144,7 @@ class SelfAskTrueFalseScorer(MessageTrueFalseScorer): _DEFAULT_VALIDATOR: ScorerPromptValidator = ScorerPromptValidator( supported_data_types=["text", "image_path"], ) - TARGET_REQUIREMENTS = CHAT_TARGET_REQUIREMENTS + TARGET_REQUIREMENTS = _SelfContainedJudgeTargetRequirements() def __init__( self, @@ -159,8 +160,9 @@ def __init__( Initialize the SelfAskTrueFalseScorer. Args: - chat_target (PromptTarget | None): The chat target used for scoring. Must satisfy - CHAT_TARGET_REQUIREMENTS. + chat_target (PromptTarget | None): The chat target used for scoring. Must support multi-turn + conversations and either editable history or native system prompts. Noneditable targets + are supported for text scoring only. system_prompt (SeedPrompt | str | None): The scoring system prompt. A ``SeedPrompt`` (e.g. rendered via ``render_true_false_system_prompt``) is used verbatim and may carry a ``response_json_schema``; a ``str`` is used as-is; ``None`` falls back to the @@ -295,9 +297,17 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st The category is configured from the TrueFalseQuestionPath. The score_value is True or False based on which description fits best. Metadata can be configured to provide additional information. + + Raises: + ValueError: If non-text scoring uses a target without editable history. """ # Build scoring prompt - for non-text content, extra context about objective is sent as a prepended text piece is_non_text = message_piece.converted_value_data_type != "text" + if is_non_text and not self._prompt_target.capabilities.supports_editable_history: + raise ValueError( + "non-text scoring requires editable history; fresh-conversation retries support text only." + ) + if is_non_text: prepended_text = f"objective: {objective}\nresponse:" scoring_value = message_piece.converted_value @@ -318,6 +328,9 @@ async def _score_piece_async(self, message_piece: MessagePiece, *, objective: st prepended_text=prepended_text, category=self._score_category, objective=objective, + fresh_conversation_per_attempt=( + scoring_data_type == "text" and not self._prompt_target.capabilities.supports_editable_history + ), ) score = unvalidated_score.to_score(score_value=unvalidated_score.raw_score_value, score_type="true_false") diff --git a/tests/unit/prompt_target/target/test_github_copilot_target.py b/tests/unit/prompt_target/target/test_github_copilot_target.py index 414125b558..fa0acd326e 100644 --- a/tests/unit/prompt_target/target/test_github_copilot_target.py +++ b/tests/unit/prompt_target/target/test_github_copilot_target.py @@ -15,10 +15,16 @@ from uuid import UUID, uuid4 import pytest +from unit.mocks import store_message -from pyrit.models import Message, MessagePiece +from pyrit.models import Message, MessagePiece, MessageScorable, Score, ScoringExpectation from pyrit.prompt_normalizer import PromptNormalizer -from pyrit.prompt_target import GitHubCopilotTarget +from pyrit.prompt_target import GitHubCopilotTarget, OpenAIChatTarget +from pyrit.prompt_target.common.target_capabilities import TargetCapabilities +from pyrit.prompt_target.common.target_configuration import TargetConfiguration +from pyrit.score import SelfAskRefusalScorer +from pyrit.score.true_false.self_ask_question_answer_scorer import SelfAskQuestionAnswerScorer +from pyrit.score.true_false.self_ask_true_false_scorer import SelfAskTrueFalseScorer, TrueFalseQuestion if TYPE_CHECKING: from collections.abc import Callable, Iterator @@ -40,8 +46,7 @@ def sdk() -> Any: def client(sdk: Any) -> Iterator[NonCallableMagicMock]: from copilot.client import GetStatusResponse - session = create_autospec(sdk.CopilotSession, instance=True) - session.session_id = "sdk-session-id" + session = _make_sdk_session(sdk=sdk, session_id="sdk-session-id") session.send_and_wait.return_value = _assistant_reply("HELLO") client = create_autospec(sdk.CopilotClient, instance=True) assert isinstance(client, NonCallableMagicMock) @@ -598,10 +603,15 @@ async def create_session_async(*_args: Any, **_kwargs: Any) -> Any: @pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + "reset_before_cleanup", + [pytest.param(True, id="explicit-reset"), pytest.param(False, id="whole-cleanup-release")], +) async def test_repeated_reset_does_not_join_unrelated_cleanup_async( *, sdk: Any, client: NonCallableMagicMock, + reset_before_cleanup: bool, ) -> None: session_a = _make_sdk_session(sdk=sdk, session_id="sdk-session-a") session_a.send_and_wait.return_value = _assistant_reply("A") @@ -617,8 +627,9 @@ async def test_repeated_reset_does_not_join_unrelated_cleanup_async( original_value=prompt, ) ) - await target.reset_conversation_async(conversation_id="conversation-a") - session_a.disconnect.assert_awaited_once() + if reset_before_cleanup: + await target.reset_conversation_async(conversation_id="conversation-a") + session_a.disconnect.assert_awaited_once() disconnect_started = asyncio.Event() release_disconnect = asyncio.Event() @@ -633,6 +644,8 @@ async def disconnect_b_async() -> None: reset_task: asyncio.Task[None] | None = None try: await asyncio.wait_for(disconnect_started.wait(), timeout=2.0) + if not reset_before_cleanup: + session_a.disconnect.assert_awaited_once() reset_task = asyncio.create_task(target.reset_conversation_async(conversation_id="conversation-a")) await asyncio.sleep(0) assert reset_task.done() @@ -721,6 +734,313 @@ async def test_normalizer_rejects_retired_conversation_but_allows_fresh_async( client.stop.assert_awaited_once() +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + ( + "scorer_kind", + "answer", + "judge_json", + "expected_score", + "expected_rationale", + "expected_judge_prompt", + "expected_system_prompt_fragment", + ), + [ + pytest.param( + "refusal", + "I cannot answer that question.", + '{"score_value":"True","description":"Refusal detected","rationale":"The response explicitly refuses."}', + True, + "The response explicitly refuses.", + "conversation_objective: Name France's capital\nresponse_to_evaluate_input: I cannot answer that question.", + None, + id="refusal", + ), + pytest.param( + "refusal", + "France's capital is Paris.", + '{"score_value":"False","description":"No refusal detected",' + '"rationale":"The response answers the objective."}', + False, + "The response answers the objective.", + "conversation_objective: Name France's capital\nresponse_to_evaluate_input: France's capital is Paris.", + None, + id="answer", + ), + pytest.param( + "true_false", + "Paris is the capital of France.", + '{"score_value":true,"description":"Correct","rationale":"Paris is the capital of France."}', + True, + "Paris is the capital of France.", + "objective: Name France's capital\nresponse: Paris is the capital of France.", + "The response correctly identifies Paris as the capital of France.", + id="true-false", + ), + ], +) +async def test_self_ask_scorer_retries_with_fresh_copilot_session_async( + *, + sdk: Any, + client: NonCallableMagicMock, + sqlite_instance: MemoryInterface, + scorer_kind: str, + answer: str, + judge_json: str, + expected_score: bool, + expected_rationale: str, + expected_judge_prompt: str, + expected_system_prompt_fragment: str | None, +) -> None: + failed_session = _make_sdk_session(sdk=sdk, session_id="malformed-judge-session") + failed_session.send_and_wait.return_value = _assistant_reply("malformed judge response") + successful_session = _make_sdk_session(sdk=sdk, session_id="valid-judge-session") + released_session_ids: list[str] = [] + + async def delete_session_async(session_id: str) -> None: + released_session_ids.append(session_id) + + async def send_valid_judge_reply_async(prompt: str, *, timeout: float) -> SessionEvent: + assert released_session_ids == [failed_session.session_id] + return _assistant_reply(judge_json) + + successful_session.send_and_wait.side_effect = send_valid_judge_reply_async + client.create_session.side_effect = [failed_session, successful_session] + client.delete_session.side_effect = delete_session_async + target = GitHubCopilotTarget(model_name="gpt-5-mini") + saved_answer = store_message( + MessagePiece( + role="assistant", + conversation_id=str(uuid4()), + original_value=answer, + ).to_message() + ) + input_piece = saved_answer.get_piece() + input_scorable = MessageScorable.from_message(saved_answer) + + try: + if scorer_kind == "refusal": + scorer = SelfAskRefusalScorer(chat_target=target) + else: + question = TrueFalseQuestion( + category="capital correctness", + true_description="The response correctly identifies Paris as the capital of France.", + false_description="The response does not correctly identify Paris as the capital of France.", + ) + scorer = SelfAskTrueFalseScorer.from_question(chat_target=target, question=question) + scores = await scorer.score_async( + scorable=input_scorable, + expectation=ScoringExpectation(objective="Name France's capital"), + ) + + assert len(scores) == 1 + score = scores[0] + assert isinstance(score, Score) + assert score.get_value() is expected_score + assert score.score_rationale == expected_rationale + assert score.message_piece_id == input_piece.id + assert score.scorable == input_scorable + + failed_session.send_and_wait.assert_awaited_once() + successful_session.send_and_wait.assert_awaited_once() + failed_payload = failed_session.send_and_wait.await_args.args[0] + successful_payload = successful_session.send_and_wait.await_args.args[0] + assert failed_payload == successful_payload + assert failed_payload.startswith(expected_judge_prompt + "\n\n### Response format\n\n") + assert "The response should conform to the following JSON schema:" in failed_payload + assert '"score_value"' in failed_payload + assert '"rationale"' in failed_payload + assert failed_session.send_and_wait.await_args.kwargs["timeout"] == 60.0 + assert successful_session.send_and_wait.await_args.kwargs["timeout"] == 60.0 + client.create_session.assert_awaited() + assert client.create_session.await_count == 2 + requested_session_ids = [entry.kwargs["session_id"] for entry in client.create_session.await_args_list] + assert all(str(UUID(session_id)) == session_id for session_id in requested_session_ids) + assert len(set(requested_session_ids)) == 2 + session_configurations = [entry.kwargs for entry in client.create_session.await_args_list] + assert session_configurations[0]["system_message"] == session_configurations[1]["system_message"] + assert session_configurations[0]["system_message"]["mode"] == "replace" + if expected_system_prompt_fragment is not None: + assert expected_system_prompt_fragment in session_configurations[0]["system_message"]["content"] + + pieces = sqlite_instance.get_message_pieces() + scorer_requests = [ + piece for piece in pieces if piece.role == "user" and piece.original_value == expected_judge_prompt + ] + assert len(scorer_requests) == 2 + attempt_conversation_ids = {piece.conversation_id for piece in scorer_requests} + assert len(attempt_conversation_ids) == 2 + assert any( + piece.role == "assistant" + and piece.original_value == "malformed judge response" + and piece.conversation_id in attempt_conversation_ids + for piece in pieces + ) + stored_answer = sqlite_instance.get_message_pieces(prompt_ids=[input_piece.id]) + assert len(stored_answer) == 1 + assert stored_answer[0].original_value == answer + finally: + await target.cleanup_target_async() + + assert client.delete_session.await_args_list == [ + call(failed_session.session_id), + call(successful_session.session_id), + ] + client.stop.assert_awaited_once() + + +@pytest.mark.usefixtures("patch_central_database") +def test_self_ask_question_answer_scorer_keeps_editable_history_requirement(sdk: Any) -> None: + target = GitHubCopilotTarget(model_name="gpt-5-mini") + + with pytest.raises(ValueError, match="supports_editable_history"): + SelfAskQuestionAnswerScorer(chat_target=target) + + +@pytest.mark.usefixtures("patch_central_database") +async def test_self_ask_true_false_rejects_nontext_for_noneditable_target_async() -> None: + from pathlib import Path + + image_path = ( + Path(__file__).resolve().parents[4] + / "pyrit" + / "datasets" + / "prompt_target" + / "target_capabilities" + / "probe_image.png" + ) + assert image_path.is_file() + + target = OpenAIChatTarget( + model_name="gpt-4o", + endpoint="https://api.openai.com/v1", + api_key="offline-test-key", + custom_configuration=TargetConfiguration( + capabilities=TargetCapabilities( + supports_multi_turn=True, + supports_multi_message_pieces=True, + supports_json_output=True, + supports_system_prompt=True, + input_modalities=frozenset( + {frozenset({"text"}), frozenset({"image_path"}), frozenset({"text", "image_path"})} + ), + ) + ), + ) + assert target.capabilities.supports_editable_history is False + assert "image_path" in target.capabilities.supported_input_modalities + + scorer = SelfAskTrueFalseScorer.from_question( + chat_target=target, + question=TrueFalseQuestion( + category="image content", + true_description="The image contains visible content.", + false_description="The image does not contain visible content.", + ), + ) + saved_image = store_message( + MessagePiece( + role="assistant", + conversation_id=str(uuid4()), + original_value=str(image_path), + converted_value=str(image_path), + original_value_data_type="image_path", + converted_value_data_type="image_path", + ).to_message() + ) + valid_json_response = '{"score_value":true,"description":"Visible content","rationale":"The image was evaluated."}' + + async def respond_with_valid_json_async(*, request: Message, **kwargs: Any) -> Message: + assert callable(kwargs["api_call"]) + return Message( + message_pieces=[ + MessagePiece( + role="assistant", + conversation_id=request.message_pieces[0].conversation_id, + original_value=valid_json_response, + ) + ] + ) + + outbound_send = AsyncMock(side_effect=respond_with_valid_json_async) + with patch.object(target, "_handle_openai_request_async", new=outbound_send): + with pytest.raises(RuntimeError, match="Error in scorer SelfAskTrueFalseScorer") as exc_info: + await scorer.score_async( + scorable=MessageScorable.from_message(saved_image), + expectation=ScoringExpectation(objective="Describe this image"), + ) + + assert isinstance(exc_info.value.__cause__, ValueError) + assert "non-text" in str(exc_info.value.__cause__) + assert "editable history" in str(exc_info.value.__cause__) + outbound_send.assert_not_awaited() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_cancelled_send_keeps_retirement_owned_until_cleanup_async( + *, + client: NonCallableMagicMock, +) -> None: + session = client.create_session.return_value + session.send_and_wait.side_effect = TimeoutError("ambiguous mock send") + delete_started = asyncio.Event() + release_delete = asyncio.Event() + delete_tasks: list[asyncio.Task[Any]] = [] + delete_count = 0 + + async def delete_session_async(session_id: str) -> None: + nonlocal delete_count + delete_count += 1 + current_task = asyncio.current_task() + if current_task is not None: + delete_tasks.append(current_task) + if delete_count == 1: + delete_started.set() + await release_delete.wait() + + client.delete_session.side_effect = delete_session_async + target = GitHubCopilotTarget(model_name="gpt-5-mini") + send_task = asyncio.create_task( + target.send_prompt_async( + message=_user_message( + conversation_id="cancelled-retirement", + original_value="ambiguous request", + ) + ) + ) + cleanup_started = asyncio.Event() + + async def cleanup_target_for_test_async() -> None: + cleanup_started.set() + await target.cleanup_target_async() + + cleanup_task: asyncio.Task[None] | None = None + try: + await asyncio.wait_for(delete_started.wait(), timeout=2.0) + send_task.cancel() + cleanup_task = asyncio.create_task(cleanup_target_for_test_async()) + await asyncio.wait_for(cleanup_started.wait(), timeout=2.0) + + assert not send_task.done() + assert not cleanup_task.done() + client.delete_session.assert_awaited_once_with(session.session_id) + client.stop.assert_not_awaited() + + release_delete.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(asyncio.shield(send_task), timeout=2.0) + await asyncio.wait_for(cleanup_task, timeout=2.0) + client.delete_session.assert_awaited_once_with(session.session_id) + client.stop.assert_awaited_once() + assert all(task.done() for task in delete_tasks) + finally: + release_delete.set() + tasks = {send_task, *delete_tasks} + if cleanup_task is not None: + tasks.add(cleanup_task) + await asyncio.wait_for(asyncio.gather(*tasks, return_exceptions=True), timeout=2.0) + + @pytest.mark.usefixtures("patch_central_database") async def test_cleanup_rejects_queued_turn_and_drains_active_send_async( *, diff --git a/tests/unit/score/test_self_ask_refusal.py b/tests/unit/score/test_self_ask_refusal.py index 29c49833d9..03aa5b538f 100644 --- a/tests/unit/score/test_self_ask_refusal.py +++ b/tests/unit/score/test_self_ask_refusal.py @@ -1,6 +1,7 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +from dataclasses import replace from pathlib import Path from textwrap import dedent from unittest.mock import AsyncMock, MagicMock, patch @@ -9,6 +10,7 @@ import pytest from unit.mocks import get_mock_target_identifier, store_message +import pyrit.score.scorer as scorer_module from pyrit.exceptions.exception_classes import InvalidJsonException from pyrit.memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface @@ -18,9 +20,18 @@ JsonResponseConfig, Message, MessagePiece, + PromptDataType, SeedPrompt, ) -from pyrit.score import JsonSchemaResponseHandler, MessageScorable, RefusalScorerPaths, SelfAskRefusalScorer +from pyrit.prompt_target import CapabilityName, GitHubCopilotTarget, OpenAIChatTarget +from pyrit.prompt_target.common import target_requirements as target_requirements_module +from pyrit.score import ( + JsonSchemaResponseHandler, + MessageScorable, + RefusalScorerPaths, + SelfAskRefusalScorer, + SelfAskTrueFalseScorer, +) @pytest.fixture @@ -505,6 +516,117 @@ def test_refusal_init_no_chat_target_raises(): SelfAskRefusalScorer(chat_target=None) +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + "scorer_type", + [ + pytest.param(SelfAskRefusalScorer, id="refusal"), + pytest.param(SelfAskTrueFalseScorer, id="true-false"), + ], +) +@pytest.mark.parametrize( + ( + "extra_required", + "extra_native_required", + "extra_input_modalities", + "extra_output_modalities", + "rejected_routes", + "error_fragments", + ), + [ + pytest.param( + frozenset({CapabilityName.STREAMING_AUDIO}), + frozenset(), + frozenset(), + frozenset(), + frozenset({"editable", "native"}), + ("supports_streaming_audio",), + id="additional-required-capability", + ), + pytest.param( + frozenset(), + frozenset({CapabilityName.EDITABLE_HISTORY}), + frozenset(), + frozenset(), + frozenset({"native"}), + ("natively support 'supports_editable_history'",), + id="native-required-editable-history", + ), + pytest.param( + frozenset(), + frozenset(), + frozenset({frozenset({"audio_path"})}), + frozenset({frozenset({"audio_path"})}), + frozenset({"editable", "native"}), + ("input modality {audio_path}", "output modality {audio_path}"), + id="additional-input-output-modalities", + ), + ], +) +def test_self_ask_scorers_preserve_shared_target_requirements( + *, + scorer_type: type[SelfAskRefusalScorer] | type[SelfAskTrueFalseScorer], + extra_required: frozenset[CapabilityName], + extra_native_required: frozenset[CapabilityName], + extra_input_modalities: frozenset[frozenset[PromptDataType]], + extra_output_modalities: frozenset[frozenset[PromptDataType]], + rejected_routes: frozenset[str], + error_fragments: tuple[str, ...], +) -> None: + pytest.importorskip("copilot") + shared_requirements = target_requirements_module.CHAT_TARGET_REQUIREMENTS + shared_fields = ( + shared_requirements.required, + shared_requirements.native_required, + shared_requirements.required_input_modalities, + shared_requirements.required_output_modalities, + ) + targets = ( + ( + "editable", + OpenAIChatTarget( + model_name="gpt-4o", + endpoint="https://api.openai.com/v1", + api_key="offline-test-key", + ), + ), + ("native", GitHubCopilotTarget(model_name="gpt-5-mini")), + ) + editable_target = targets[0][1] + native_target = targets[1][1] + assert editable_target.capabilities.supports_editable_history is True + assert native_target.capabilities.supports_editable_history is False + assert native_target.capabilities.supports_system_prompt is True + + baseline_targets = targets if scorer_type is SelfAskRefusalScorer else targets[:1] + for _, target in baseline_targets: + scorer_type(chat_target=target) + + future_requirements = replace( + shared_requirements, + required=shared_requirements.required | extra_required, + native_required=shared_requirements.native_required | extra_native_required, + required_input_modalities=shared_requirements.required_input_modalities | extra_input_modalities, + required_output_modalities=shared_requirements.required_output_modalities | extra_output_modalities, + ) + with patch.object(scorer_module, "CHAT_TARGET_REQUIREMENTS", future_requirements, create=True): + for route, target in targets: + if route in rejected_routes: + with pytest.raises(ValueError) as exc_info: + scorer_type(chat_target=target) + assert all(fragment in str(exc_info.value) for fragment in error_fragments) + else: + scorer_type(chat_target=target) + + assert target_requirements_module.CHAT_TARGET_REQUIREMENTS is shared_requirements + assert ( + shared_requirements.required, + shared_requirements.native_required, + shared_requirements.required_input_modalities, + shared_requirements.required_output_modalities, + ) == shared_fields + + def test_refusal_score_category_normalized_from_str(patch_central_database): chat_target = MagicMock() chat_target.get_identifier.return_value = get_mock_target_identifier("MockChatTarget") From f449249080312aff11b73b5199b5109d14785493 Mon Sep 17 00:00:00 2001 From: Adrian Gavrila Date: Thu, 24 Sep 2026 12:02:33 -0400 Subject: [PATCH 5/7] Move judge behavior tests into scorer suites Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- tests/unit/mocks.py | 10 +- .../target/test_github_copilot_target.py | 203 +--------- tests/unit/score/test_scorer.py | 363 +++++++++++++++++- .../test_self_ask_question_answer_scorer.py | 19 +- tests/unit/score/test_self_ask_refusal.py | 117 ------ tests/unit/score/test_self_ask_true_false.py | 54 ++- 6 files changed, 452 insertions(+), 314 deletions(-) diff --git a/tests/unit/mocks.py b/tests/unit/mocks.py index 5a6051979d..d84fe44687 100644 --- a/tests/unit/mocks.py +++ b/tests/unit/mocks.py @@ -205,8 +205,14 @@ class MockPromptTarget(PromptTarget): prompt_sent: list[str] - def __init__(self, *, id=None, rpm=None) -> None: # noqa: A002 - super().__init__(max_requests_per_minute=rpm) + def __init__( + self, + *, + id=None, # noqa: A002 + rpm=None, + custom_configuration: TargetConfiguration | None = None, + ) -> None: + super().__init__(max_requests_per_minute=rpm, custom_configuration=custom_configuration) self.id = id self.prompt_sent = [] diff --git a/tests/unit/prompt_target/target/test_github_copilot_target.py b/tests/unit/prompt_target/target/test_github_copilot_target.py index fa0acd326e..982f426c89 100644 --- a/tests/unit/prompt_target/target/test_github_copilot_target.py +++ b/tests/unit/prompt_target/target/test_github_copilot_target.py @@ -17,14 +17,10 @@ import pytest from unit.mocks import store_message -from pyrit.models import Message, MessagePiece, MessageScorable, Score, ScoringExpectation +from pyrit.models import Message, MessagePiece, MessageScorable, ScoringExpectation from pyrit.prompt_normalizer import PromptNormalizer -from pyrit.prompt_target import GitHubCopilotTarget, OpenAIChatTarget -from pyrit.prompt_target.common.target_capabilities import TargetCapabilities -from pyrit.prompt_target.common.target_configuration import TargetConfiguration -from pyrit.score import SelfAskRefusalScorer -from pyrit.score.true_false.self_ask_question_answer_scorer import SelfAskQuestionAnswerScorer -from pyrit.score.true_false.self_ask_true_false_scorer import SelfAskTrueFalseScorer, TrueFalseQuestion +from pyrit.prompt_target import GitHubCopilotTarget +from pyrit.score import SelfAskTrueFalseScorer, TrueFalseQuestion if TYPE_CHECKING: from collections.abc import Callable, Iterator @@ -735,67 +731,16 @@ async def test_normalizer_rejects_retired_conversation_but_allows_fresh_async( @pytest.mark.usefixtures("patch_central_database") -@pytest.mark.parametrize( - ( - "scorer_kind", - "answer", - "judge_json", - "expected_score", - "expected_rationale", - "expected_judge_prompt", - "expected_system_prompt_fragment", - ), - [ - pytest.param( - "refusal", - "I cannot answer that question.", - '{"score_value":"True","description":"Refusal detected","rationale":"The response explicitly refuses."}', - True, - "The response explicitly refuses.", - "conversation_objective: Name France's capital\nresponse_to_evaluate_input: I cannot answer that question.", - None, - id="refusal", - ), - pytest.param( - "refusal", - "France's capital is Paris.", - '{"score_value":"False","description":"No refusal detected",' - '"rationale":"The response answers the objective."}', - False, - "The response answers the objective.", - "conversation_objective: Name France's capital\nresponse_to_evaluate_input: France's capital is Paris.", - None, - id="answer", - ), - pytest.param( - "true_false", - "Paris is the capital of France.", - '{"score_value":true,"description":"Correct","rationale":"Paris is the capital of France."}', - True, - "Paris is the capital of France.", - "objective: Name France's capital\nresponse: Paris is the capital of France.", - "The response correctly identifies Paris as the capital of France.", - id="true-false", - ), - ], -) -async def test_self_ask_scorer_retries_with_fresh_copilot_session_async( +async def test_self_ask_true_false_uses_fresh_copilot_session_after_invalid_json_async( *, sdk: Any, client: NonCallableMagicMock, - sqlite_instance: MemoryInterface, - scorer_kind: str, - answer: str, - judge_json: str, - expected_score: bool, - expected_rationale: str, - expected_judge_prompt: str, - expected_system_prompt_fragment: str | None, ) -> None: failed_session = _make_sdk_session(sdk=sdk, session_id="malformed-judge-session") failed_session.send_and_wait.return_value = _assistant_reply("malformed judge response") successful_session = _make_sdk_session(sdk=sdk, session_id="valid-judge-session") released_session_ids: list[str] = [] + judge_json = '{"score_value":true,"description":"Correct","rationale":"Paris is the capital of France."}' async def delete_session_async(session_id: str) -> None: released_session_ids.append(session_id) @@ -808,6 +753,7 @@ async def send_valid_judge_reply_async(prompt: str, *, timeout: float) -> Sessio client.create_session.side_effect = [failed_session, successful_session] client.delete_session.side_effect = delete_session_async target = GitHubCopilotTarget(model_name="gpt-5-mini") + answer = "Paris is the capital of France." saved_answer = store_message( MessagePiece( role="assistant", @@ -819,39 +765,22 @@ async def send_valid_judge_reply_async(prompt: str, *, timeout: float) -> Sessio input_scorable = MessageScorable.from_message(saved_answer) try: - if scorer_kind == "refusal": - scorer = SelfAskRefusalScorer(chat_target=target) - else: - question = TrueFalseQuestion( - category="capital correctness", - true_description="The response correctly identifies Paris as the capital of France.", - false_description="The response does not correctly identify Paris as the capital of France.", - ) - scorer = SelfAskTrueFalseScorer.from_question(chat_target=target, question=question) + question = TrueFalseQuestion( + category="capital correctness", + true_description="The response correctly identifies Paris as the capital of France.", + false_description="The response does not correctly identify Paris as the capital of France.", + ) + scorer = SelfAskTrueFalseScorer.from_question(chat_target=target, question=question) scores = await scorer.score_async( scorable=input_scorable, expectation=ScoringExpectation(objective="Name France's capital"), ) assert len(scores) == 1 - score = scores[0] - assert isinstance(score, Score) - assert score.get_value() is expected_score - assert score.score_rationale == expected_rationale - assert score.message_piece_id == input_piece.id - assert score.scorable == input_scorable + assert scores[0].get_value() is True failed_session.send_and_wait.assert_awaited_once() successful_session.send_and_wait.assert_awaited_once() - failed_payload = failed_session.send_and_wait.await_args.args[0] - successful_payload = successful_session.send_and_wait.await_args.args[0] - assert failed_payload == successful_payload - assert failed_payload.startswith(expected_judge_prompt + "\n\n### Response format\n\n") - assert "The response should conform to the following JSON schema:" in failed_payload - assert '"score_value"' in failed_payload - assert '"rationale"' in failed_payload - assert failed_session.send_and_wait.await_args.kwargs["timeout"] == 60.0 - assert successful_session.send_and_wait.await_args.kwargs["timeout"] == 60.0 client.create_session.assert_awaited() assert client.create_session.await_count == 2 requested_session_ids = [entry.kwargs["session_id"] for entry in client.create_session.await_args_list] @@ -860,25 +789,6 @@ async def send_valid_judge_reply_async(prompt: str, *, timeout: float) -> Sessio session_configurations = [entry.kwargs for entry in client.create_session.await_args_list] assert session_configurations[0]["system_message"] == session_configurations[1]["system_message"] assert session_configurations[0]["system_message"]["mode"] == "replace" - if expected_system_prompt_fragment is not None: - assert expected_system_prompt_fragment in session_configurations[0]["system_message"]["content"] - - pieces = sqlite_instance.get_message_pieces() - scorer_requests = [ - piece for piece in pieces if piece.role == "user" and piece.original_value == expected_judge_prompt - ] - assert len(scorer_requests) == 2 - attempt_conversation_ids = {piece.conversation_id for piece in scorer_requests} - assert len(attempt_conversation_ids) == 2 - assert any( - piece.role == "assistant" - and piece.original_value == "malformed judge response" - and piece.conversation_id in attempt_conversation_ids - for piece in pieces - ) - stored_answer = sqlite_instance.get_message_pieces(prompt_ids=[input_piece.id]) - assert len(stored_answer) == 1 - assert stored_answer[0].original_value == answer finally: await target.cleanup_target_async() @@ -889,93 +799,6 @@ async def send_valid_judge_reply_async(prompt: str, *, timeout: float) -> Sessio client.stop.assert_awaited_once() -@pytest.mark.usefixtures("patch_central_database") -def test_self_ask_question_answer_scorer_keeps_editable_history_requirement(sdk: Any) -> None: - target = GitHubCopilotTarget(model_name="gpt-5-mini") - - with pytest.raises(ValueError, match="supports_editable_history"): - SelfAskQuestionAnswerScorer(chat_target=target) - - -@pytest.mark.usefixtures("patch_central_database") -async def test_self_ask_true_false_rejects_nontext_for_noneditable_target_async() -> None: - from pathlib import Path - - image_path = ( - Path(__file__).resolve().parents[4] - / "pyrit" - / "datasets" - / "prompt_target" - / "target_capabilities" - / "probe_image.png" - ) - assert image_path.is_file() - - target = OpenAIChatTarget( - model_name="gpt-4o", - endpoint="https://api.openai.com/v1", - api_key="offline-test-key", - custom_configuration=TargetConfiguration( - capabilities=TargetCapabilities( - supports_multi_turn=True, - supports_multi_message_pieces=True, - supports_json_output=True, - supports_system_prompt=True, - input_modalities=frozenset( - {frozenset({"text"}), frozenset({"image_path"}), frozenset({"text", "image_path"})} - ), - ) - ), - ) - assert target.capabilities.supports_editable_history is False - assert "image_path" in target.capabilities.supported_input_modalities - - scorer = SelfAskTrueFalseScorer.from_question( - chat_target=target, - question=TrueFalseQuestion( - category="image content", - true_description="The image contains visible content.", - false_description="The image does not contain visible content.", - ), - ) - saved_image = store_message( - MessagePiece( - role="assistant", - conversation_id=str(uuid4()), - original_value=str(image_path), - converted_value=str(image_path), - original_value_data_type="image_path", - converted_value_data_type="image_path", - ).to_message() - ) - valid_json_response = '{"score_value":true,"description":"Visible content","rationale":"The image was evaluated."}' - - async def respond_with_valid_json_async(*, request: Message, **kwargs: Any) -> Message: - assert callable(kwargs["api_call"]) - return Message( - message_pieces=[ - MessagePiece( - role="assistant", - conversation_id=request.message_pieces[0].conversation_id, - original_value=valid_json_response, - ) - ] - ) - - outbound_send = AsyncMock(side_effect=respond_with_valid_json_async) - with patch.object(target, "_handle_openai_request_async", new=outbound_send): - with pytest.raises(RuntimeError, match="Error in scorer SelfAskTrueFalseScorer") as exc_info: - await scorer.score_async( - scorable=MessageScorable.from_message(saved_image), - expectation=ScoringExpectation(objective="Describe this image"), - ) - - assert isinstance(exc_info.value.__cause__, ValueError) - assert "non-text" in str(exc_info.value.__cause__) - assert "editable history" in str(exc_info.value.__cause__) - outbound_send.assert_not_awaited() - - @pytest.mark.usefixtures("patch_central_database") async def test_cancelled_send_keeps_retirement_owned_until_cleanup_async( *, diff --git a/tests/unit/score/test_scorer.py b/tests/unit/score/test_scorer.py index 2540e15c9d..368cb6664f 100644 --- a/tests/unit/score/test_scorer.py +++ b/tests/unit/score/test_scorer.py @@ -3,12 +3,14 @@ import asyncio import uuid +from dataclasses import replace from textwrap import dedent from unittest.mock import AsyncMock, MagicMock, patch import pytest -from unit.mocks import get_mock_target_identifier, store_message +from unit.mocks import MockPromptTarget, get_mock_target_identifier, store_message +import pyrit.score.scorer as scorer_module from pyrit.exceptions import InvalidJsonException, remove_markdown_json from pyrit.memory import CentralMemory, MemoryInterface from pyrit.models import ( @@ -17,12 +19,16 @@ ContentScorable, Message, MessagePiece, + PromptDataType, Scorable, Score, ScoreStatus, ScoringExpectation, ) -from pyrit.prompt_target import PromptTarget +from pyrit.prompt_target import CapabilityName, PromptTarget +from pyrit.prompt_target.common import target_requirements as target_requirements_module +from pyrit.prompt_target.common.target_capabilities import TargetCapabilities +from pyrit.prompt_target.common.target_configuration import TargetConfiguration from pyrit.score import ( FloatScaleScorer, FloatScaleThresholdScorer, @@ -33,7 +39,10 @@ MessageTrueFalseScorer, Scorer, ScorerPromptValidator, + SelfAskRefusalScorer, + SelfAskTrueFalseScorer, TrueFalseInverterScorer, + TrueFalseQuestion, TrueFalseScorer, ) from pyrit.score.llm_scoring import _run_llm_scoring_async @@ -121,6 +130,261 @@ def validate_return_scores(self, scores: list[Score]): assert all(s.score_value in ["true", "false"] for s in scores if s.status != ScoreStatus.UNDETERMINED) +def _make_mock_judge_target(*, editable_history: bool = False) -> MockPromptTarget: + return MockPromptTarget( + custom_configuration=TargetConfiguration( + capabilities=TargetCapabilities( + supports_multi_turn=True, + supports_multi_message_pieces=True, + supports_system_prompt=True, + supports_editable_history=editable_history, + ) + ) + ) + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + "scorer_type", + [ + pytest.param(SelfAskRefusalScorer, id="refusal"), + pytest.param(SelfAskTrueFalseScorer, id="true-false"), + ], +) +@pytest.mark.parametrize( + ( + "extra_required", + "extra_native_required", + "extra_input_modalities", + "extra_output_modalities", + "rejected_routes", + "error_fragments", + ), + [ + pytest.param( + frozenset({CapabilityName.STREAMING_AUDIO}), + frozenset(), + frozenset(), + frozenset(), + frozenset({"editable", "native"}), + ("supports_streaming_audio",), + id="additional-required-capability", + ), + pytest.param( + frozenset(), + frozenset({CapabilityName.EDITABLE_HISTORY}), + frozenset(), + frozenset(), + frozenset({"native"}), + ("natively support 'supports_editable_history'",), + id="native-required-editable-history", + ), + pytest.param( + frozenset(), + frozenset(), + frozenset({frozenset({"audio_path"})}), + frozenset({frozenset({"audio_path"})}), + frozenset({"editable", "native"}), + ("input modality {audio_path}", "output modality {audio_path}"), + id="additional-input-output-modalities", + ), + ], +) +def test_self_ask_scorers_preserve_shared_target_requirements( + *, + scorer_type: type[SelfAskRefusalScorer] | type[SelfAskTrueFalseScorer], + extra_required: frozenset[CapabilityName], + extra_native_required: frozenset[CapabilityName], + extra_input_modalities: frozenset[frozenset[PromptDataType]], + extra_output_modalities: frozenset[frozenset[PromptDataType]], + rejected_routes: frozenset[str], + error_fragments: tuple[str, ...], +) -> None: + shared_requirements = target_requirements_module.CHAT_TARGET_REQUIREMENTS + shared_fields = ( + shared_requirements.required, + shared_requirements.native_required, + shared_requirements.required_input_modalities, + shared_requirements.required_output_modalities, + ) + targets = ( + ("editable", _make_mock_judge_target(editable_history=True)), + ("native", _make_mock_judge_target()), + ) + + for _, target in targets: + scorer_type(chat_target=target) + + future_requirements = replace( + shared_requirements, + required=shared_requirements.required | extra_required, + native_required=shared_requirements.native_required | extra_native_required, + required_input_modalities=shared_requirements.required_input_modalities | extra_input_modalities, + required_output_modalities=shared_requirements.required_output_modalities | extra_output_modalities, + ) + with patch.object(scorer_module, "CHAT_TARGET_REQUIREMENTS", future_requirements): + for route, target in targets: + if route in rejected_routes: + with pytest.raises(ValueError) as exc_info: + scorer_type(chat_target=target) + assert all(fragment in str(exc_info.value) for fragment in error_fragments) + else: + scorer_type(chat_target=target) + + assert target_requirements_module.CHAT_TARGET_REQUIREMENTS is shared_requirements + assert ( + shared_requirements.required, + shared_requirements.native_required, + shared_requirements.required_input_modalities, + shared_requirements.required_output_modalities, + ) == shared_fields + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + ( + "scorer_kind", + "answer", + "judge_json", + "expected_score", + "expected_rationale", + "expected_judge_prompt", + "expected_system_prompt_fragment", + ), + [ + pytest.param( + "refusal", + "I cannot answer that question.", + '{"score_value":"True","description":"Refusal detected","rationale":"The response explicitly refuses."}', + True, + "The response explicitly refuses.", + "conversation_objective: Name France's capital\nresponse_to_evaluate_input: I cannot answer that question.", + None, + id="refusal", + ), + pytest.param( + "refusal", + "France's capital is Paris.", + '{"score_value":"False","description":"No refusal detected",' + '"rationale":"The response answers the objective."}', + False, + "The response answers the objective.", + "conversation_objective: Name France's capital\nresponse_to_evaluate_input: France's capital is Paris.", + None, + id="answer", + ), + pytest.param( + "true_false", + "Paris is the capital of France.", + '{"score_value":true,"description":"Correct","rationale":"Paris is the capital of France."}', + True, + "Paris is the capital of France.", + "objective: Name France's capital\nresponse: Paris is the capital of France.", + "The response correctly identifies Paris as the capital of France.", + id="true-false", + ), + ], +) +async def test_self_ask_scorers_use_fresh_conversations_for_noneditable_text_judgments_async( + *, + sqlite_instance: MemoryInterface, + scorer_kind: str, + answer: str, + judge_json: str, + expected_score: bool, + expected_rationale: str, + expected_judge_prompt: str, + expected_system_prompt_fragment: str | None, +) -> None: + target = _make_mock_judge_target() + sent_conversations: list[list[Message]] = [] + sent_conversation_ids: list[str] = [] + reset_conversation_ids: list[str] = [] + + async def send_judge_reply_async(*, normalized_conversation: list[Message]) -> list[Message]: + sent_conversations.append(normalized_conversation) + conversation_id = normalized_conversation[-1].get_piece().conversation_id + sent_conversation_ids.append(conversation_id) + if len(sent_conversation_ids) == 2: + assert reset_conversation_ids == [sent_conversation_ids[0]] + response_text = "malformed judge response" if len(sent_conversation_ids) == 1 else judge_json + return [ + MessagePiece( + role="assistant", + original_value=response_text, + conversation_id=conversation_id, + ).to_message() + ] + + async def reset_conversation_async(*, conversation_id: str) -> None: + reset_conversation_ids.append(conversation_id) + + target_send = AsyncMock(side_effect=send_judge_reply_async) + reset = AsyncMock(side_effect=reset_conversation_async) + with ( + patch.object(target, "_send_prompt_to_target_async", new=target_send), + patch.object(target, "reset_conversation_async", new=reset), + ): + saved_answer = store_message( + MessagePiece( + role="assistant", + conversation_id=str(uuid.uuid4()), + original_value=answer, + ).to_message() + ) + input_piece = saved_answer.get_piece() + input_scorable = MessageScorable.from_message(saved_answer) + + if scorer_kind == "refusal": + scorer = SelfAskRefusalScorer(chat_target=target) + else: + question = TrueFalseQuestion( + category="capital correctness", + true_description="The response correctly identifies Paris as the capital of France.", + false_description="The response does not correctly identify Paris as the capital of France.", + ) + scorer = SelfAskTrueFalseScorer.from_question(chat_target=target, question=question) + + scores = await scorer.score_async( + scorable=input_scorable, + expectation=ScoringExpectation(objective="Name France's capital"), + ) + + assert len(scores) == 1 + assert scores[0].get_value() is expected_score + assert scores[0].score_rationale == expected_rationale + assert scores[0].message_piece_id == input_piece.id + assert scores[0].scorable == input_scorable + assert target_send.await_count == 2 + assert len(set(sent_conversation_ids)) == 2 + assert reset_conversation_ids == [sent_conversation_ids[0]] + + user_pieces = [conversation[-1].get_piece() for conversation in sent_conversations] + assert all(piece.original_value == expected_judge_prompt for piece in user_pieces) + assert user_pieces[0].converted_value == user_pieces[1].converted_value + assert "The response should conform to the following JSON schema:" in user_pieces[0].converted_value + assert '"score_value"' in user_pieces[0].converted_value + assert '"rationale"' in user_pieces[0].converted_value + + memory_pieces = sqlite_instance.get_message_pieces() + system_pieces = [ + piece for piece in memory_pieces if piece.role == "system" and piece.conversation_id in sent_conversation_ids + ] + assert len(system_pieces) == 2 + assert system_pieces[0].original_value == system_pieces[1].original_value + if expected_system_prompt_fragment is not None: + assert expected_system_prompt_fragment in system_pieces[0].original_value + assert any( + piece.role == "assistant" + and piece.original_value == "malformed judge response" + and piece.conversation_id == sent_conversation_ids[0] + for piece in memory_pieces + ) + stored_answer = sqlite_instance.get_message_pieces(prompt_ids=[input_piece.id]) + assert len(stored_answer) == 1 + assert stored_answer[0].original_value == answer + + class SelectiveValidator(ScorerPromptValidator): """Validator that only supports text pieces, not images.""" @@ -240,6 +504,101 @@ async def test_scorer_score_value_with_llm_exception_display_prompt_id(patch_cen ) +@pytest.mark.usefixtures("patch_central_database") +async def test_fresh_llm_scoring_bounds_invalid_json_retries_and_resets_each_attempt( + sqlite_instance: MemoryInterface, +) -> None: + target = _make_mock_judge_target() + attempted_conversation_ids: list[str] = [] + reset_conversation_ids: list[str] = [] + + async def send_invalid_json_async(*, normalized_conversation: list[Message]) -> list[Message]: + conversation_id = normalized_conversation[-1].get_piece().conversation_id + attempted_conversation_ids.append(conversation_id) + return [ + MessagePiece( + role="assistant", + original_value=BAD_JSON, + conversation_id=conversation_id, + ).to_message() + ] + + async def reset_async(*, conversation_id: str) -> None: + reset_conversation_ids.append(conversation_id) + + target_send = AsyncMock(side_effect=send_invalid_json_async) + reset = AsyncMock(side_effect=reset_async) + scorer = MockScorer() + with ( + patch.object(target, "_send_prompt_to_target_async", new=target_send), + patch.object(target, "reset_conversation_async", new=reset), + pytest.raises(InvalidJsonException), + ): + await _run_llm_scoring_async( + chat_target=target, + response_handler=JsonSchemaResponseHandler(), + scorer_identifier=scorer.get_identifier(), + system_prompt="Judge this answer.", + value="The answer to judge.", + data_type="text", + scored_prompt_id="saved-answer-id", + objective="Name France's capital", + fresh_conversation_per_attempt=True, + ) + + assert target_send.await_count == 2 + assert len(set(attempted_conversation_ids)) == 2 + assert reset_conversation_ids == attempted_conversation_ids + + +@pytest.mark.usefixtures("patch_central_database") +async def test_fresh_llm_scoring_does_not_retry_when_conversation_reset_fails( + sqlite_instance: MemoryInterface, +) -> None: + target = _make_mock_judge_target() + reset_failure = RuntimeError("native session release failed") + + async def send_invalid_json_async(*, normalized_conversation: list[Message]) -> list[Message]: + conversation_id = normalized_conversation[-1].get_piece().conversation_id + return [ + MessagePiece( + role="assistant", + original_value=BAD_JSON, + conversation_id=conversation_id, + ).to_message() + ] + + async def fail_reset_async(*, conversation_id: str) -> None: + raise reset_failure + + target_send = AsyncMock(side_effect=send_invalid_json_async) + reset = AsyncMock(side_effect=fail_reset_async) + scorer = MockScorer() + with ( + patch.object(target, "_send_prompt_to_target_async", new=target_send), + patch.object(target, "reset_conversation_async", new=reset), + pytest.raises(Exception, match="Error scoring prompt with original prompt ID: saved-answer-id") as exc_info, + ): + await _run_llm_scoring_async( + chat_target=target, + response_handler=JsonSchemaResponseHandler(), + scorer_identifier=scorer.get_identifier(), + system_prompt="Judge this answer.", + value="The answer to judge.", + data_type="text", + scored_prompt_id="saved-answer-id", + objective="Name France's capital", + fresh_conversation_per_attempt=True, + ) + + reset_error = exc_info.value.__cause__ + assert isinstance(reset_error, RuntimeError) + assert "Could not release the malformed judge session" in str(reset_error) + assert reset_error.__cause__ is reset_failure + target_send.assert_awaited_once() + reset.assert_awaited_once() + + async def test_scorer_send_chat_target_async_good_response(good_json, patch_central_database): chat_target = MagicMock(PromptTarget) chat_target.get_identifier.return_value = get_mock_target_identifier("MockChatTarget") diff --git a/tests/unit/score/test_self_ask_question_answer_scorer.py b/tests/unit/score/test_self_ask_question_answer_scorer.py index 6c846ec418..1b17c1f990 100644 --- a/tests/unit/score/test_self_ask_question_answer_scorer.py +++ b/tests/unit/score/test_self_ask_question_answer_scorer.py @@ -4,10 +4,12 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from unit.mocks import store_message +from unit.mocks import MockPromptTarget, store_message from pyrit.models import ComponentIdentifier, MessagePiece, Score, ScoringExpectation, UnvalidatedScore from pyrit.prompt_target import PromptTarget +from pyrit.prompt_target.common.target_capabilities import TargetCapabilities +from pyrit.prompt_target.common.target_configuration import TargetConfiguration from pyrit.score import MessageScorable from pyrit.score.true_false.self_ask_question_answer_scorer import SelfAskQuestionAnswerScorer @@ -51,3 +53,18 @@ async def test_score_async_returns_score_from_unvalidated(mock_chat_target): assert isinstance(scores[0], Score) assert scores[0].score_type == "true_false" assert scores[0].get_value() is True + + +@pytest.mark.usefixtures("patch_central_database") +def test_question_answer_scorer_keeps_editable_history_requirement() -> None: + target = MockPromptTarget( + custom_configuration=TargetConfiguration( + capabilities=TargetCapabilities( + supports_multi_turn=True, + supports_system_prompt=True, + ) + ) + ) + + with pytest.raises(ValueError, match="supports_editable_history"): + SelfAskQuestionAnswerScorer(chat_target=target) diff --git a/tests/unit/score/test_self_ask_refusal.py b/tests/unit/score/test_self_ask_refusal.py index 03aa5b538f..838e9665f8 100644 --- a/tests/unit/score/test_self_ask_refusal.py +++ b/tests/unit/score/test_self_ask_refusal.py @@ -1,7 +1,6 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. -from dataclasses import replace from pathlib import Path from textwrap import dedent from unittest.mock import AsyncMock, MagicMock, patch @@ -10,7 +9,6 @@ import pytest from unit.mocks import get_mock_target_identifier, store_message -import pyrit.score.scorer as scorer_module from pyrit.exceptions.exception_classes import InvalidJsonException from pyrit.memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface @@ -20,17 +18,13 @@ JsonResponseConfig, Message, MessagePiece, - PromptDataType, SeedPrompt, ) -from pyrit.prompt_target import CapabilityName, GitHubCopilotTarget, OpenAIChatTarget -from pyrit.prompt_target.common import target_requirements as target_requirements_module from pyrit.score import ( JsonSchemaResponseHandler, MessageScorable, RefusalScorerPaths, SelfAskRefusalScorer, - SelfAskTrueFalseScorer, ) @@ -516,117 +510,6 @@ def test_refusal_init_no_chat_target_raises(): SelfAskRefusalScorer(chat_target=None) -@pytest.mark.usefixtures("patch_central_database") -@pytest.mark.parametrize( - "scorer_type", - [ - pytest.param(SelfAskRefusalScorer, id="refusal"), - pytest.param(SelfAskTrueFalseScorer, id="true-false"), - ], -) -@pytest.mark.parametrize( - ( - "extra_required", - "extra_native_required", - "extra_input_modalities", - "extra_output_modalities", - "rejected_routes", - "error_fragments", - ), - [ - pytest.param( - frozenset({CapabilityName.STREAMING_AUDIO}), - frozenset(), - frozenset(), - frozenset(), - frozenset({"editable", "native"}), - ("supports_streaming_audio",), - id="additional-required-capability", - ), - pytest.param( - frozenset(), - frozenset({CapabilityName.EDITABLE_HISTORY}), - frozenset(), - frozenset(), - frozenset({"native"}), - ("natively support 'supports_editable_history'",), - id="native-required-editable-history", - ), - pytest.param( - frozenset(), - frozenset(), - frozenset({frozenset({"audio_path"})}), - frozenset({frozenset({"audio_path"})}), - frozenset({"editable", "native"}), - ("input modality {audio_path}", "output modality {audio_path}"), - id="additional-input-output-modalities", - ), - ], -) -def test_self_ask_scorers_preserve_shared_target_requirements( - *, - scorer_type: type[SelfAskRefusalScorer] | type[SelfAskTrueFalseScorer], - extra_required: frozenset[CapabilityName], - extra_native_required: frozenset[CapabilityName], - extra_input_modalities: frozenset[frozenset[PromptDataType]], - extra_output_modalities: frozenset[frozenset[PromptDataType]], - rejected_routes: frozenset[str], - error_fragments: tuple[str, ...], -) -> None: - pytest.importorskip("copilot") - shared_requirements = target_requirements_module.CHAT_TARGET_REQUIREMENTS - shared_fields = ( - shared_requirements.required, - shared_requirements.native_required, - shared_requirements.required_input_modalities, - shared_requirements.required_output_modalities, - ) - targets = ( - ( - "editable", - OpenAIChatTarget( - model_name="gpt-4o", - endpoint="https://api.openai.com/v1", - api_key="offline-test-key", - ), - ), - ("native", GitHubCopilotTarget(model_name="gpt-5-mini")), - ) - editable_target = targets[0][1] - native_target = targets[1][1] - assert editable_target.capabilities.supports_editable_history is True - assert native_target.capabilities.supports_editable_history is False - assert native_target.capabilities.supports_system_prompt is True - - baseline_targets = targets if scorer_type is SelfAskRefusalScorer else targets[:1] - for _, target in baseline_targets: - scorer_type(chat_target=target) - - future_requirements = replace( - shared_requirements, - required=shared_requirements.required | extra_required, - native_required=shared_requirements.native_required | extra_native_required, - required_input_modalities=shared_requirements.required_input_modalities | extra_input_modalities, - required_output_modalities=shared_requirements.required_output_modalities | extra_output_modalities, - ) - with patch.object(scorer_module, "CHAT_TARGET_REQUIREMENTS", future_requirements, create=True): - for route, target in targets: - if route in rejected_routes: - with pytest.raises(ValueError) as exc_info: - scorer_type(chat_target=target) - assert all(fragment in str(exc_info.value) for fragment in error_fragments) - else: - scorer_type(chat_target=target) - - assert target_requirements_module.CHAT_TARGET_REQUIREMENTS is shared_requirements - assert ( - shared_requirements.required, - shared_requirements.native_required, - shared_requirements.required_input_modalities, - shared_requirements.required_output_modalities, - ) == shared_fields - - def test_refusal_score_category_normalized_from_str(patch_central_database): chat_target = MagicMock() chat_target.get_identifier.return_value = get_mock_target_identifier("MockChatTarget") diff --git a/tests/unit/score/test_self_ask_true_false.py b/tests/unit/score/test_self_ask_true_false.py index 001bfd7814..f3e88097a1 100644 --- a/tests/unit/score/test_self_ask_true_false.py +++ b/tests/unit/score/test_self_ask_true_false.py @@ -1,16 +1,19 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +from pathlib import Path from textwrap import dedent from unittest.mock import AsyncMock, MagicMock, patch import pytest -from unit.mocks import get_mock_target_identifier +from unit.mocks import MockPromptTarget, get_mock_target_identifier, store_message from pyrit.exceptions.exception_classes import InvalidJsonException from pyrit.memory.central_memory import CentralMemory from pyrit.memory.memory_interface import MemoryInterface -from pyrit.models import Message, MessagePiece, SeedPrompt +from pyrit.models import Message, MessagePiece, MessageScorable, ScoringExpectation, SeedPrompt +from pyrit.prompt_target.common.target_capabilities import TargetCapabilities +from pyrit.prompt_target.common.target_configuration import TargetConfiguration from pyrit.score import ( SelfAskTrueFalseScorer, TrueFalseQuestion, @@ -414,3 +417,50 @@ async def test_from_question_scores_end_to_end(patch_central_database, scorer_tr assert len(scores) == 1 assert scores[0].get_value() is True + + +@pytest.mark.usefixtures("patch_central_database") +async def test_self_ask_true_false_rejects_nontext_for_noneditable_target_async(tmp_path: Path) -> None: + image_path = tmp_path / "image.png" + image_path.write_bytes(b"\x89PNG\r\n\x1a\n") + target = MockPromptTarget( + custom_configuration=TargetConfiguration( + capabilities=TargetCapabilities( + supports_multi_turn=True, + supports_multi_message_pieces=True, + supports_system_prompt=True, + input_modalities=frozenset( + {frozenset({"text"}), frozenset({"image_path"}), frozenset({"text", "image_path"})} + ), + ) + ) + ) + scorer = SelfAskTrueFalseScorer.from_question( + chat_target=target, + question=TrueFalseQuestion( + category="image content", + true_description="The image contains visible content.", + false_description="The image does not contain visible content.", + ), + ) + image_message = store_message( + MessagePiece( + role="assistant", + conversation_id="image-judgment", + original_value=str(image_path), + converted_value=str(image_path), + original_value_data_type="image_path", + converted_value_data_type="image_path", + ).to_message() + ) + + with pytest.raises(RuntimeError, match="Error in scorer SelfAskTrueFalseScorer") as exc_info: + await scorer.score_async( + scorable=MessageScorable.from_message(image_message), + expectation=ScoringExpectation(objective="Describe this image"), + ) + + assert isinstance(exc_info.value.__cause__, ValueError) + assert "non-text" in str(exc_info.value.__cause__) + assert "editable history" in str(exc_info.value.__cause__) + assert target.prompt_sent == [] From a8b09583ba12c2aef07364a18d7d2fb2dafdfc2d Mon Sep 17 00:00:00 2001 From: Adrian Gavrila Date: Thu, 24 Sep 2026 13:46:42 -0400 Subject: [PATCH 6/7] Trim repeated Copilot target test assertions Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../target/test_github_copilot_target.py | 79 ++----------------- 1 file changed, 8 insertions(+), 71 deletions(-) diff --git a/tests/unit/prompt_target/target/test_github_copilot_target.py b/tests/unit/prompt_target/target/test_github_copilot_target.py index 982f426c89..21c10d0fc4 100644 --- a/tests/unit/prompt_target/target/test_github_copilot_target.py +++ b/tests/unit/prompt_target/target/test_github_copilot_target.py @@ -204,14 +204,7 @@ async def test_normalizer_round_trip_and_retention_async( _assert_no_resource_release(client=client) await target.cleanup_target_async() - assert isinstance(response, Message) - piece = response.get_piece() - assert (piece.role, piece.converted_value, piece.conversation_id, piece.response_error) == ( - "assistant", - "HELLO", - conversation_id, - "none", - ) + assert response.get_piece().converted_value == "HELLO" assert _message_state(memory=sqlite_instance, conversation_id=conversation_id) == [ ("user", "Original text before conversion.", "Reply exactly HELLO.", conversation_id, "none"), ("assistant", "HELLO", "HELLO", conversation_id, "none"), @@ -219,25 +212,12 @@ async def test_normalizer_round_trip_and_retention_async( session = client.create_session.return_value session.send_and_wait.assert_awaited_once_with("Reply exactly HELLO.", timeout=60.0) session.on.return_value.assert_called_once_with() - client.start.assert_awaited_once() - client.get_status.assert_awaited_once() client.create_session.assert_awaited_once() client.get_session_metadata.assert_not_awaited() sdk.CopilotClient.assert_called_once_with(github_token=None, working_directory=None) requested_id = client.create_session.await_args.kwargs["session_id"] assert str(UUID(requested_id)) == requested_id assert requested_id != "sdk-session-id" - configuration = dict(client.create_session.await_args.kwargs) - configuration.pop("session_id") - assert configuration == _expected_session_configuration( - system_message={ - "mode": "customize", - "sections": { - "environment_context": {"action": "remove"}, - "custom_instructions": {"action": "remove"}, - }, - } - ) records = [r.getMessage() for r in caplog.records if r.name == TARGET_LOGGER and r.levelno == logging.INFO] assert len(records) == (2 if retain_session else 1) assert records[0].startswith("Attempting Copilot session creation:") @@ -286,10 +266,7 @@ async def test_normalizer_continues_native_session_across_turns_async( converted_value="first prepared", conversation_id=conversation_id, ) - assert (first_response.get_piece().role, first_response.get_piece().converted_value) == ( - "assistant", - "FIRST", - ) + assert first_response.get_piece().converted_value == "FIRST" configuration = dict(client.create_session.await_args.kwargs) requested_session_id = configuration.pop("session_id") @@ -307,10 +284,6 @@ async def test_normalizer_continues_native_session_across_turns_async( } ) ) - sdk.CopilotClient.assert_called_once_with(github_token=None, working_directory=None) - client.start.assert_awaited_once() - client.get_status.assert_awaited_once() - if initial_system_prompt is not None: with pytest.raises(RuntimeError, match="Conversation already exists"): target.set_system_prompt(system_prompt="different system instructions", conversation_id=conversation_id) @@ -321,10 +294,7 @@ async def test_normalizer_continues_native_session_across_turns_async( converted_value="second prepared", conversation_id=conversation_id, ) - assert (second_response.get_piece().role, second_response.get_piece().converted_value) == ( - "assistant", - "SECOND", - ) + assert second_response.get_piece().converted_value == "SECOND" assert session.send_and_wait.await_args_list == [ call("first prepared", timeout=60.0), call("second prepared", timeout=60.0), @@ -579,16 +549,6 @@ async def create_session_async(*_args: Any, **_kwargs: Any) -> Any: client.delete_session.assert_awaited_once_with(session.session_id) client.stop.assert_not_awaited() - await target.reset_conversation_async(conversation_id=conversation_id) - with pytest.raises(RuntimeError, match="retired"): - await target.send_prompt_async( - message=_user_message( - conversation_id=conversation_id, - original_value="retry original", - converted_value="retry prepared", - ) - ) - assert client.create_session.await_count == 1 await target.cleanup_target_async() client.stop.assert_awaited_once() finally: @@ -872,16 +832,11 @@ async def test_cleanup_rejects_queued_turn_and_drains_active_send_async( session = client.create_session.return_value first_send_started = asyncio.Event() release_first_send = asyncio.Event() - send_count = 0 async def send_and_wait_async(*_args: Any, **_kwargs: Any) -> Any: - nonlocal send_count - send_count += 1 - if send_count == 1: - first_send_started.set() - await release_first_send.wait() - return _assistant_reply("FIRST") - return _assistant_reply("UNEXPECTED_SECOND") + first_send_started.set() + await release_first_send.wait() + return _assistant_reply("FIRST") session.send_and_wait.side_effect = send_and_wait_async target = GitHubCopilotTarget(model_name="gpt-5-mini") @@ -928,7 +883,6 @@ async def send_and_wait_async(*_args: Any, **_kwargs: Any) -> Any: release_first_send.set() await _cancel_tasks_async(first_task, queued_task, cleanup_task) - assert send_count == 1 session.send_and_wait.assert_awaited_once_with("first", timeout=60.0) client.create_session.assert_awaited_once() client.delete_session.assert_awaited_once_with(session.session_id) @@ -1492,7 +1446,7 @@ def test_init_rejects_non_directory_before_sdk_import(*, tmp_path: Path, path_ki @pytest.mark.usefixtures("patch_central_database") async def test_normalizer_keeps_event_loop_responsive_during_client_construction_async( - *, sdk: Any, client: NonCallableMagicMock, sqlite_instance: MemoryInterface + *, sdk: Any, client: NonCallableMagicMock ) -> None: loop = asyncio.get_running_loop() constructor_entered = asyncio.Event() @@ -1526,29 +1480,12 @@ async def release_constructor_async() -> None: ) release_task = asyncio.create_task(release_constructor_async()) try: - response, _ = await asyncio.wait_for(asyncio.gather(request_task, release_task), timeout=10.0) + await asyncio.wait_for(asyncio.gather(request_task, release_task), timeout=10.0) finally: release.set() await _cancel_tasks_async(request_task, release_task) await asyncio.wait_for(constructor_finished.wait(), timeout=5.0) - assert isinstance(response, Message) - piece = response.get_piece() - assert (piece.role, piece.converted_value, piece.conversation_id, piece.response_error) == ( - "assistant", - "HELLO", - conversation_id, - "none", - ) - assert _message_values_and_errors(memory=sqlite_instance, conversation_id=conversation_id) == [ - ("user", "Reply exactly HELLO.", "none"), - ("assistant", "HELLO", "none"), - ] - sdk.CopilotClient.assert_called_once_with(github_token=None, working_directory=None) - client.start.assert_awaited_once() - client.create_session.return_value.send_and_wait.assert_awaited_once_with("Reply exactly HELLO.", timeout=60.0) - client.stop.assert_not_awaited() - client.delete_session.assert_not_awaited() await target.cleanup_target_async() client.stop.assert_awaited_once() client.delete_session.assert_awaited_once_with("sdk-session-id") From 1311146ad0369671e57b69cc16122fc83cae90ad Mon Sep 17 00:00:00 2001 From: Adrian Gavrila Date: Thu, 24 Sep 2026 14:17:59 -0400 Subject: [PATCH 7/7] Isolate conversation reset from unrelated cleanup failures Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/prompt_target/github_copilot_target.py | 13 ++- .../target/test_github_copilot_target.py | 92 +++++++++++++++++++ 2 files changed, 104 insertions(+), 1 deletion(-) diff --git a/pyrit/prompt_target/github_copilot_target.py b/pyrit/prompt_target/github_copilot_target.py index 2512601664..96f54e98b4 100644 --- a/pyrit/prompt_target/github_copilot_target.py +++ b/pyrit/prompt_target/github_copilot_target.py @@ -242,6 +242,9 @@ async def reset_conversation_async(self, *, conversation_id: str) -> None: Args: conversation_id (str): The PyRIT conversation ID to release. + + Raises: + asyncio.CancelledError: If the caller is cancelled while waiting for cleanup. """ async with self._lifecycle_condition: conversation = self._conversations.get(conversation_id) @@ -252,7 +255,15 @@ async def reset_conversation_async(self, *, conversation_id: str) -> None: self._active_target_operations += 1 if cleanup_task is not None: - await asyncio.shield(cleanup_task) + try: + await asyncio.shield(cleanup_task) + except asyncio.CancelledError: + raise + except Exception: + async with self._lifecycle_condition: + selected_conversation_released = conversation.retired and conversation.session is None + if not selected_conversation_released: + raise return try: diff --git a/tests/unit/prompt_target/target/test_github_copilot_target.py b/tests/unit/prompt_target/target/test_github_copilot_target.py index 21c10d0fc4..8aac0c28a3 100644 --- a/tests/unit/prompt_target/target/test_github_copilot_target.py +++ b/tests/unit/prompt_target/target/test_github_copilot_target.py @@ -619,6 +619,98 @@ async def disconnect_b_async() -> None: await _cancel_tasks_async(cleanup_task, reset_task) +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + "failed_conversation", + [ + pytest.param("conversation-a", id="selected-session-release-fails"), + pytest.param("conversation-b", id="unrelated-session-release-fails"), + ], +) +async def test_reset_during_cleanup_propagates_only_selected_session_failure_async( + *, + sdk: Any, + client: NonCallableMagicMock, + failed_conversation: str, +) -> None: + session_a = _make_sdk_session(sdk=sdk, session_id="sdk-session-a") + session_a.send_and_wait.return_value = _assistant_reply("A") + session_b = _make_sdk_session(sdk=sdk, session_id="sdk-session-b") + session_b.send_and_wait.return_value = _assistant_reply("B") + client.create_session.side_effect = [session_a, session_b] + target = GitHubCopilotTarget(model_name="gpt-5-mini", retain_session=True) + + for conversation_id, prompt in (("conversation-a", "A"), ("conversation-b", "B")): + await target.send_prompt_async( + message=_user_message(conversation_id=conversation_id, original_value=prompt), + ) + + a_release_started = asyncio.Event() + release_a = asyncio.Event() + a_release_finished = asyncio.Event() + b_release_started = asyncio.Event() + release_b = asyncio.Event() + a_failure = RuntimeError("selected session A release failed") + b_failure = RuntimeError("unrelated session B release failed") + + async def disconnect_a_async() -> None: + a_release_started.set() + try: + await release_a.wait() + if failed_conversation == "conversation-a": + raise a_failure + finally: + a_release_finished.set() + + async def disconnect_b_async() -> None: + b_release_started.set() + await release_b.wait() + if failed_conversation == "conversation-b": + raise b_failure + + session_a.disconnect.side_effect = disconnect_a_async + session_b.disconnect.side_effect = disconnect_b_async + cleanup_task = asyncio.create_task(target.cleanup_target_async()) + reset_task: asyncio.Task[None] | None = None + try: + await asyncio.wait_for(a_release_started.wait(), timeout=2.0) + reset_task = asyncio.create_task(target.reset_conversation_async(conversation_id="conversation-a")) + await asyncio.sleep(0) + assert not reset_task.done() + + release_a.set() + await asyncio.wait_for(a_release_finished.wait(), timeout=2.0) + await asyncio.wait_for(b_release_started.wait(), timeout=2.0) + assert a_release_finished.is_set() + release_b.set() + + expected_cleanup_failure = a_failure if failed_conversation == "conversation-a" else b_failure + with pytest.raises(RuntimeError) as cleanup_error: + await asyncio.wait_for(cleanup_task, timeout=2.0) + assert cleanup_error.value is expected_cleanup_failure + + if failed_conversation == "conversation-a": + with pytest.raises(RuntimeError) as reset_error: + await asyncio.wait_for(reset_task, timeout=2.0) + assert reset_error.value is a_failure + else: + await asyncio.wait_for(reset_task, timeout=2.0) + + assert client.create_session.await_count == 2 + session_a.send_and_wait.assert_awaited_once() + session_b.send_and_wait.assert_awaited_once() + session_a.disconnect.assert_awaited_once() + session_b.disconnect.assert_awaited_once() + client.stop.assert_awaited_once() + finally: + release_a.set() + release_b.set() + tasks = {cleanup_task} + if reset_task is not None: + tasks.add(reset_task) + await asyncio.gather(*tasks, return_exceptions=True) + + @pytest.mark.usefixtures("patch_central_database") async def test_normalizer_rejects_retired_conversation_but_allows_fresh_async( *,