diff --git a/.pyrit_conf_example b/.pyrit_conf_example index 5f45135bfd..171a952531 100644 --- a/.pyrit_conf_example +++ b/.pyrit_conf_example @@ -94,6 +94,14 @@ operation: op_trash_panda # - https://my-vault.vault.azure.net/secrets/my-pyrit-env # env_akv_strict: true +# API-created Target Persistence +# ------------------------------ +# Non-secret target settings are stored in the configured memory database. +# Set this vault URL to persist API keys supplied through the backend UI/API. +# The backend authenticates to Key Vault with DefaultAzureCredential. +# Without it, targets created with an explicit API key remain memory-only. +# target_secret_key_vault_url: https://my-vault.vault.azure.net + # Auto-discovered ~/.pyrit/.env remains supported but emits a security warning. # Prefer env_akv_ref for shared or deployed secrets. # Use ~/.pyrit/.env.local for quick local plaintext patches or when Azure is unavailable. diff --git a/doc/getting_started/pyrit_conf.md b/doc/getting_started/pyrit_conf.md index 10032e1590..5acb7c7f37 100644 --- a/doc/getting_started/pyrit_conf.md +++ b/doc/getting_started/pyrit_conf.md @@ -181,6 +181,17 @@ Complete-value `kv:`, `akv:`, `azure_key_vault:`, and `env_akv_ref:` references Ordinary malformed dotenv lines retain python-dotenv's permissive behavior. `env_akv_strict` controls malformed Key Vault reference syntax in all sources: strict mode raises; non-strict mode warns and skips that assignment. Authentication, authorization, transport, missing-secret, and missing-value failures always raise. +### `target_secret_key_vault_url` + +Optional Azure Key Vault URL for API keys submitted when users create targets through the backend UI or REST API. +Non-secret target settings are stored in the configured memory database, while API keys are stored in this vault +using `DefaultAzureCredential`. Without this setting, a target created with an explicit API key is memory-only and +the UI warns that it will be cleared when the backend restarts. + +```yaml +target_secret_key_vault_url: https://my-vault.vault.azure.net +``` + Environment loading preserves the historical non-transactional dotenv behavior. The bootstrap document and each local file update `os.environ` as they load. If a later source or child-secret lookup fails, assignments made by earlier sources remain in the process environment. When `env_akv_ref` is not configured, an empty `env_files` list or missing default files leaves existing process environment variables unchanged and initialization continues. @@ -411,6 +422,9 @@ initializers: # - https://my-vault.vault.azure.net/secrets/my-pyrit-env # env_akv_strict: true +# Optional: persist API keys for targets created through the backend +# target_secret_key_vault_url: https://my-vault.vault.azure.net + # Optional plaintext local patch or non-Azure workflow # env_files: # - /path/to/.env.local diff --git a/frontend/src/components/Config/CreateTargetDialog.test.tsx b/frontend/src/components/Config/CreateTargetDialog.test.tsx index 319bed9128..312c6bcd27 100644 --- a/frontend/src/components/Config/CreateTargetDialog.test.tsx +++ b/frontend/src/components/Config/CreateTargetDialog.test.tsx @@ -18,6 +18,10 @@ jest.mock("@/services/api", () => ({ const mockedTargetsApi = targetsApi as jest.Mocked; const TARGET_CATALOG: TargetCatalogResponse = { + persistence: { + definitions_enabled: true, + api_keys_enabled: true, + }, items: [ { target_type: "AzureMLChatTarget", @@ -259,6 +263,28 @@ describe("CreateTargetDialog", () => { expect(screen.getByText("Cancel")).toBeInTheDocument(); }); + it("should warn when a submitted API key cannot be persisted", async () => { + const user = userEvent.setup(); + mockedTargetsApi.listTargetCatalog.mockResolvedValue({ + ...TARGET_CATALOG, + persistence: { + definitions_enabled: true, + api_keys_enabled: false, + }, + }); + + render( + + + + ); + await flushCatalogFetch(); + await selectTargetType("OpenAIChatTarget"); + await user.type(screen.getByPlaceholderText("API key"), "temporary-key"); + + expect(screen.getByText(/only be registered in memory/i)).toBeInTheDocument(); + }); + it("should show friendly names, catalog descriptions, implementation identifiers, and auth for all target types", async () => { render( @@ -425,7 +451,7 @@ describe("CreateTargetDialog", () => { screen.queryByRole("radio", { name: /Identity-based/ }) ).not.toBeInTheDocument(); expect( - screen.getByPlaceholderText("API key (stored in memory only)") + screen.getByPlaceholderText("API key") ).toBeInTheDocument(); // Selecting an identity-capable type should reveal the Authentication field. @@ -632,7 +658,7 @@ describe("CreateTargetDialog", () => { fireEvent.change(endpointInput, { target: { value: "https://api.openai.com" } }); // Fill API key — use fireEvent.change for the same reason as endpoint input. - fireEvent.change(screen.getByPlaceholderText("API key (stored in memory only)"), { + fireEvent.change(screen.getByPlaceholderText("API key"), { target: { value: "sk-test-key-123" }, }); @@ -903,7 +929,7 @@ describe("CreateTargetDialog", () => { // check that API Key field is visible by default. expect( - screen.getByPlaceholderText("API key (stored in memory only)") + screen.getByPlaceholderText("API key") ).toBeInTheDocument(); // Select identity option. @@ -915,7 +941,7 @@ describe("CreateTargetDialog", () => { // check that API Key field is hidden when identity mode is selected. expect( - screen.queryByPlaceholderText("API key (stored in memory only)") + screen.queryByPlaceholderText("API key") ).not.toBeInTheDocument(); await user.click(screen.getByText("Create Target")); @@ -957,7 +983,7 @@ describe("CreateTargetDialog", () => { // Type a key, then switch to identity option. fireEvent.change( - screen.getByPlaceholderText("API key (stored in memory only)"), + screen.getByPlaceholderText("API key"), { target: { value: "sk-typed-before-switch" } } ); diff --git a/frontend/src/components/Config/CreateTargetDialog.tsx b/frontend/src/components/Config/CreateTargetDialog.tsx index bcfb2a3e1c..11c459d0f1 100644 --- a/frontend/src/components/Config/CreateTargetDialog.tsx +++ b/frontend/src/components/Config/CreateTargetDialog.tsx @@ -27,7 +27,7 @@ import { import { DeleteRegular } from '@fluentui/react-icons' import { targetsApi } from '@/services/api' import { toApiError } from '@/services/errors' -import type { TargetInstance, TargetCatalogEntry } from '@/types' +import type { TargetInstance, TargetCatalogEntry, TargetPersistenceStatus } from '@/types' import { targetIdentifierHash, targetModelName, @@ -225,6 +225,7 @@ export default function CreateTargetDialog({ open, onClose, onCreated, existingT // Available target types + their auth facts, fetched from the backend registry. const [catalogEntries, setCatalogEntries] = useState([]) const [catalogStatus, setCatalogStatus] = useState('loading') + const [persistenceStatus, setPersistenceStatus] = useState(null) const catalogByType = useMemo( () => new Map(catalogEntries.map((entry) => [entry.target_type, entry])), [catalogEntries], @@ -240,6 +241,7 @@ export default function CreateTargetDialog({ open, onClose, onCreated, existingT if (open) { setCatalogEntries([]) setCatalogStatus('loading') + setPersistenceStatus(null) } } @@ -252,12 +254,14 @@ export default function CreateTargetDialog({ open, onClose, onCreated, existingT .then((res) => { if (!cancelled) { setCatalogEntries(res.items) + setPersistenceStatus(res.persistence) setCatalogStatus('loaded') } }) .catch(() => { if (!cancelled) { setCatalogEntries([]) + setPersistenceStatus(null) setCatalogStatus('error') } }) @@ -298,6 +302,16 @@ export default function CreateTargetDialog({ open, onClose, onCreated, existingT return null })() const showIdentityEndpointError = identityEndpointError !== null + const persistenceWarning = (() => { + if (!targetType || !persistenceStatus) return null + if (!persistenceStatus.definitions_enabled) { + return 'Persistent target storage is not configured. This target will only be registered in memory and will be cleared when the backend restarts.' + } + if (apiKey && !isIdentity && !persistenceStatus.api_keys_enabled) { + return 'Persistent API-key storage is not configured. This target will only be registered in memory and will be cleared when the backend restarts.' + } + return null + })() // Fetch the available targets when the dialog opens with RoundRobin selected. // If the parent already passed targets, derive availableTargets from them @@ -518,6 +532,12 @@ export default function CreateTargetDialog({ open, onClose, onCreated, existingT )} + {persistenceWarning && ( + + {persistenceWarning} + + )} + setApiKey(data.value)} /> diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index fc7346eb56..6955921bd6 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -259,6 +259,12 @@ export interface TargetCatalogEntry { export interface TargetCatalogResponse { items: TargetCatalogEntry[] + persistence: TargetPersistenceStatus +} + +export interface TargetPersistenceStatus { + definitions_enabled: boolean + api_keys_enabled: boolean } // --- Attacks --- diff --git a/pyrit/backend/main.py b/pyrit/backend/main.py index b195ae3b9c..723c0f5174 100644 --- a/pyrit/backend/main.py +++ b/pyrit/backend/main.py @@ -36,6 +36,8 @@ version, ) from pyrit.backend.services.initializer_service import get_initializer_service +from pyrit.backend.services.target_service import get_target_service +from pyrit.memory import CentralMemory from pyrit.setup.configuration_loader import ConfigurationLoader # Check for development mode from environment variable @@ -70,6 +72,13 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: for order_index, initializer in enumerate(config.initializer_configs) ] await get_initializer_service().run_additional_initializers_async() + target_service = get_target_service() + target_service.configure_persistence( + memory=CentralMemory.get_memory_instance(), + definitions_enabled=config.memory_db_type != "in_memory", + target_secret_key_vault_url=config.target_secret_key_vault_url, + ) + await target_service.restore_persisted_targets_async() # Expose config values to route handlers via app.state default_labels: dict[str, str] = {} diff --git a/pyrit/backend/models/targets.py b/pyrit/backend/models/targets.py index 87797d32d1..5e0c6b3498 100644 --- a/pyrit/backend/models/targets.py +++ b/pyrit/backend/models/targets.py @@ -19,6 +19,7 @@ __all__ = [ "CreateTargetRequest", "TargetCatalogEntry", + "TargetPersistenceStatus", "TargetCatalogResponse", "TargetListResponse", ] @@ -43,10 +44,18 @@ class TargetCatalogEntry(BaseModel): description: str | None = Field(None, description="Short description of the target from its docstring") +class TargetPersistenceStatus(BaseModel): + """Persistence capabilities for targets created through the backend.""" + + definitions_enabled: bool = Field(..., description="Whether non-secret target definitions survive restarts") + api_keys_enabled: bool = Field(..., description="Whether submitted API keys can be stored in Azure Key Vault") + + class TargetCatalogResponse(BaseModel): """Response for listing available target types from the registry.""" items: list[TargetCatalogEntry] = Field(..., description="List of available target types") + persistence: TargetPersistenceStatus = Field(..., description="Backend target persistence capabilities") class TargetListResponse(BaseModel): diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index 77d078a18c..5685fa5ff1 100644 --- a/pyrit/backend/services/target_service.py +++ b/pyrit/backend/services/target_service.py @@ -16,6 +16,7 @@ import logging from functools import lru_cache from typing import Any, Literal, cast +from uuid import NAMESPACE_URL, uuid5 from pyrit.backend.mappers.target_mappers import target_object_to_instance from pyrit.backend.models.common import PaginationInfo @@ -24,8 +25,11 @@ TargetCatalogEntry, TargetCatalogResponse, TargetListResponse, + TargetPersistenceStatus, ) +from pyrit.memory.memory_interface import MemoryInterface from pyrit.models.catalog.target import TargetInstance +from pyrit.models.persisted_target import PersistedTarget from pyrit.registry import TargetRegistry logger = logging.getLogger(__name__) @@ -43,6 +47,21 @@ class TargetService: def __init__(self) -> None: """Initialize the target service.""" self._registry = TargetRegistry.get_registry_singleton() + self._memory: MemoryInterface | None = None + self._definition_persistence_enabled = False + self._target_secret_key_vault_url: str | None = None + + def configure_persistence( + self, + *, + memory: MemoryInterface, + definitions_enabled: bool, + target_secret_key_vault_url: str | None, + ) -> None: + """Configure persistence after Central Memory has been initialized.""" + self._memory = memory + self._definition_persistence_enabled = definitions_enabled + self._target_secret_key_vault_url = target_secret_key_vault_url def _build_instance_from_object(self, *, target_registry_name: str, target_obj: Any) -> TargetInstance: """ @@ -148,7 +167,15 @@ async def list_target_catalog_async(self) -> TargetCatalogResponse: ) for metadata in metadata_items ] - return TargetCatalogResponse(items=items) + return TargetCatalogResponse( + items=items, + persistence=TargetPersistenceStatus( + definitions_enabled=self._definition_persistence_enabled, + api_keys_enabled=bool( + self._definition_persistence_enabled and self._target_secret_key_vault_url + ), + ), + ) async def create_target_async(self, *, request: CreateTargetRequest) -> TargetInstance: """ @@ -189,12 +216,106 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn params.pop("api_key", None) target_obj = self._registry.create_instance(request.type, **params) - - self._registry.instances.register(target_obj) - target_registry_name = target_obj.get_identifier().unique_name + await self._persist_target_async( + request=request, + params=params, + target_registry_name=target_registry_name, + ) + self._registry.instances.register(target_obj) return self._build_instance_from_object(target_registry_name=target_registry_name, target_obj=target_obj) + async def restore_persisted_targets_async(self) -> None: + """Recreate persisted API targets and register them in original creation order.""" + if not self._definition_persistence_enabled or self._memory is None: + return + + definitions = await asyncio.to_thread(self._memory.get_persisted_targets) + for definition in definitions: + try: + params: dict[str, Any] = dict(definition.parameters) + if definition.secret_name: + params["api_key"] = await self._get_api_key_async(secret_name=definition.secret_name) + if definition.auth_mode == "identity": + params.pop("api_key", None) + target_obj = self._registry.create_instance(definition.target_type, **params) + self._registry.instances.register(target_obj, name=definition.target_registry_name) + except Exception: + logger.exception("Failed to restore persisted target '%s'.", definition.target_registry_name) + + async def _persist_target_async( + self, + *, + request: CreateTargetRequest, + params: dict[str, Any], + target_registry_name: str, + ) -> None: + """Persist a target definition when durable storage is configured.""" + if not self._definition_persistence_enabled or self._memory is None: + return + + persisted_params = dict(params) + api_key = persisted_params.pop("api_key", None) + if api_key is not None and not self._target_secret_key_vault_url: + logger.warning( + "Target '%s' is memory-only because target_secret_key_vault_url is not configured.", + target_registry_name, + ) + return + + target_id = str(uuid5(NAMESPACE_URL, f"pyrit-target:{target_registry_name}")) + secret_name = f"pyrit-target-{target_id}" if api_key is not None else None + if secret_name: + await self._set_api_key_async(secret_name=secret_name, api_key=str(api_key)) + + definition = PersistedTarget( + id=target_id, + target_registry_name=target_registry_name, + target_type=request.type, + parameters=persisted_params, + auth_mode=request.auth_mode, + secret_name=secret_name, + ) + await asyncio.to_thread(self._memory.add_persisted_target, target=definition) + + async def _set_api_key_async(self, *, secret_name: str, api_key: str) -> None: + """Store one API key in the configured Azure Key Vault.""" + if not self._target_secret_key_vault_url: + raise RuntimeError("Target secret Key Vault is not configured.") + + from azure.identity.aio import DefaultAzureCredential + from azure.keyvault.secrets.aio import SecretClient + + async with DefaultAzureCredential() as credential: + async with SecretClient( + vault_url=self._target_secret_key_vault_url, + credential=credential, + ) as client: + await client.set_secret(secret_name, api_key) + + async def _get_api_key_async(self, *, secret_name: str) -> str: + """ + Load one API key from the configured Azure Key Vault. + + Returns: + str: The stored API key. + """ + if not self._target_secret_key_vault_url: + raise RuntimeError("Target secret Key Vault is not configured.") + + from azure.identity.aio import DefaultAzureCredential + from azure.keyvault.secrets.aio import SecretClient + + async with DefaultAzureCredential() as credential: + async with SecretClient( + vault_url=self._target_secret_key_vault_url, + credential=credential, + ) as client: + secret = await client.get_secret(secret_name) + if secret.value is None: + raise ValueError(f"Azure Key Vault secret '{secret_name}' has no value.") + return secret.value + @lru_cache(maxsize=1) def get_target_service() -> TargetService: diff --git a/pyrit/memory/alembic/versions/6d8f0a2c4e6b_add_persisted_targets_table.py b/pyrit/memory/alembic/versions/6d8f0a2c4e6b_add_persisted_targets_table.py new file mode 100644 index 0000000000..8370b6cb0e --- /dev/null +++ b/pyrit/memory/alembic/versions/6d8f0a2c4e6b_add_persisted_targets_table.py @@ -0,0 +1,42 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +""" +Add persisted API-created targets table. + +Revision ID: 6d8f0a2c4e6b +Revises: 4c9a6e1f2b7d +Create Date: 2026-08-28 12:30:00.000000 +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision: str = "6d8f0a2c4e6b" +down_revision: str | None = "4c9a6e1f2b7d" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + """Apply this schema upgrade.""" + op.create_table( + "PersistedTargets", + sa.Column("id", sa.String(length=36), nullable=False), + sa.Column("target_registry_name", sa.String(length=512), nullable=False), + sa.Column("target_type", sa.String(length=128), nullable=False), + sa.Column("parameters", sa.JSON(), nullable=False), + sa.Column("auth_mode", sa.String(length=16), nullable=False), + sa.Column("secret_name", sa.String(length=127), nullable=True), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint("target_registry_name"), + ) + + +def downgrade() -> None: + """Revert this schema upgrade.""" + op.drop_table("PersistedTargets") diff --git a/pyrit/memory/memory_interface.py b/pyrit/memory/memory_interface.py index d698593a5b..e57577fd41 100644 --- a/pyrit/memory/memory_interface.py +++ b/pyrit/memory/memory_interface.py @@ -35,6 +35,7 @@ ConversationEntry, ConverterIdentifierEntry, EmbeddingDataEntry, + PersistedTargetEntry, PromptConverterIdentifierEntry, PromptMemoryEntry, ScenarioIdentifierEntry, @@ -67,6 +68,7 @@ IdentifierType, Message, MessagePiece, + PersistedTarget, ScenarioIdentifier, ScenarioResult, ScenarioRunState, @@ -492,6 +494,23 @@ def delete_additional_initializer(self, *, initializer_id: str) -> None: logger.exception(f"Error deleting additional initializer '{initializer_id}': {e}") raise + def add_persisted_target(self, *, target: PersistedTarget) -> None: + """Insert or replace a persisted target definition, keyed by its ``id``.""" + self._update_entry(PersistedTargetEntry.from_domain_model(target)) + + def get_persisted_targets(self) -> Sequence[PersistedTarget]: + """ + Load persisted target definitions in creation order. + + Returns: + Sequence[PersistedTarget]: Persisted target definitions. + """ + entries = self._query_entries( + PersistedTargetEntry, + order_by=PersistedTargetEntry.created_at.asc(), + ) + return [entry.to_domain_model() for entry in entries] + @abc.abstractmethod def _init_storage_io(self) -> None: """ diff --git a/pyrit/memory/memory_models.py b/pyrit/memory/memory_models.py index 8fc1247bfd..a7637374cb 100644 --- a/pyrit/memory/memory_models.py +++ b/pyrit/memory/memory_models.py @@ -54,6 +54,7 @@ ConverterIdentifier, EvaluationIdentifier, MessagePiece, + PersistedTarget, PromptDataType, ScenarioEvaluationIdentifier, ScenarioIdentifier, @@ -464,6 +465,59 @@ def to_domain_model(self) -> AdditionalInitializer: ) +class PersistedTargetEntry(DomainBackedEntry[PersistedTarget]): + """Persistence row for a reconstructable API-created target.""" + + __tablename__ = "PersistedTargets" + __table_args__ = {"extend_existing": True} + + id: Mapped[str] = mapped_column(String(36), primary_key=True) + target_registry_name: Mapped[str] = mapped_column(String(512), nullable=False, unique=True) + target_type: Mapped[str] = mapped_column(String(128), nullable=False) + parameters: Mapped[dict[str, Any]] = mapped_column(JSON, nullable=False) + auth_mode: Mapped[str] = mapped_column(String(16), nullable=False) + secret_name: Mapped[str | None] = mapped_column(String(127), nullable=True) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + + @classmethod + def from_domain_model(cls, domain_model: PersistedTarget) -> Self: + """ + Build an unsaved persisted-target row from its domain model. + + Returns: + Self: A new, unsaved row. + """ + return cls( + id=domain_model.id, + target_registry_name=domain_model.target_registry_name, + target_type=domain_model.target_type, + parameters=domain_model.parameters, + auth_mode=domain_model.auth_mode, + secret_name=domain_model.secret_name, + created_at=domain_model.created_at, + ) + + def to_domain_model(self) -> PersistedTarget: + """ + Convert this row back into its domain model. + + Returns: + PersistedTarget: The reconstructed persisted target. + """ + created_at = self.created_at + if created_at.tzinfo is None: + created_at = created_at.replace(tzinfo=timezone.utc) + return PersistedTarget( + id=self.id, + target_registry_name=self.target_registry_name, + target_type=self.target_type, + parameters=self.parameters, + auth_mode=self.auth_mode, + secret_name=self.secret_name, + created_at=created_at, + ) + + T = TypeVar("T", bound=ComponentIdentifier) diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index 4a6dfc36ca..ff039577e3 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -90,6 +90,7 @@ RegistryReference, display_choices, ) + from pyrit.models.persisted_target import PersistedTarget from pyrit.models.question_answering import QuestionAnsweringDataset, QuestionAnsweringEntry, QuestionChoice from pyrit.models.results.attack_result import AttackOutcome, AttackResult, AttackResultT from pyrit.models.results.scenario_result import ScenarioResult, ScenarioRunState @@ -202,6 +203,7 @@ "NextMessageSystemPromptPaths": "pyrit.models.seeds", "ObjectiveTargetEvaluationIdentifier": "pyrit.models.identifiers", "Parameter": "pyrit.models.parameter", + "PersistedTarget": "pyrit.models.persisted_target", "ParameterDestination": "pyrit.models.parameter", "PromptDataType": "pyrit.models.literals", "PromptResponseError": "pyrit.models.literals", diff --git a/pyrit/models/persisted_target.py b/pyrit/models/persisted_target.py new file mode 100644 index 0000000000..617e210b7c --- /dev/null +++ b/pyrit/models/persisted_target.py @@ -0,0 +1,30 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Persisted definitions for targets created through the backend API.""" + +from datetime import datetime, timezone +from typing import Literal +from uuid import uuid4 + +from pydantic import BaseModel, Field + +from pyrit.models.identifiers import JSONValue + + +class PersistedTarget(BaseModel): + """A reconstructable API-created target definition without inline secrets.""" + + id: str = Field(default_factory=lambda: str(uuid4()), description="Stable unique row id.") + target_registry_name: str = Field(..., description="Registry name assigned when the target was created.") + target_type: str = Field(..., description="Target class name used by TargetRegistry.") + parameters: dict[str, JSONValue] = Field( + default_factory=dict, + description="JSON-serializable constructor parameters with secrets removed.", + ) + auth_mode: Literal["api_key", "identity"] = Field(default="api_key", description="Target authentication mode.") + secret_name: str | None = Field(default=None, description="Azure Key Vault secret name containing the API key.") + created_at: datetime = Field( + default_factory=lambda: datetime.now(timezone.utc), + description="Creation time used to restore dependent targets in order.", + ) diff --git a/pyrit/setup/configuration_loader.py b/pyrit/setup/configuration_loader.py index f3d601528d..13fc03fb34 100644 --- a/pyrit/setup/configuration_loader.py +++ b/pyrit/setup/configuration_loader.py @@ -105,6 +105,8 @@ class ConfigurationLoader(YamlLoadable): seed: Optional root seed for deterministic converter operations. operator: Name for the current operator, e.g. a team or username. operation: Name for the current operation. + target_secret_key_vault_url: Optional Azure Key Vault URL used to store + API keys for targets created through the backend. Example YAML configuration: memory_db_type: sqlite @@ -148,6 +150,7 @@ class ConfigurationLoader(YamlLoadable): operation: str | None = None max_concurrent_scenario_runs: int = 3 allow_custom_initializers: bool = False + target_secret_key_vault_url: str | None = None server: dict[str, Any] | None = None extensions: dict[str, Any] = field(default_factory=dict) @@ -157,6 +160,7 @@ def __post_init__(self) -> None: self._normalize_memory_db_type() self._normalize_initializers() self._validate_env_akv_ref() + self._normalize_target_secret_key_vault_url() self._normalize_server() def _validate_env_akv_ref(self) -> None: @@ -203,6 +207,19 @@ def _normalize_memory_db_type(self) -> None: # Store normalized snake_case value self.memory_db_type = normalized + def _normalize_target_secret_key_vault_url(self) -> None: + """ + Normalize the optional Key Vault URL used for API-created target secrets. + + Raises: + ValueError: If the configured value is not a non-empty string. + """ + if self.target_secret_key_vault_url is None: + return + if not isinstance(self.target_secret_key_vault_url, str) or not self.target_secret_key_vault_url.strip(): + raise ValueError("target_secret_key_vault_url must be a non-empty Azure Key Vault URL.") + self.target_secret_key_vault_url = self.target_secret_key_vault_url.strip().rstrip("/") + def _normalize_initializers(self) -> None: """ Normalize initializer entries to InitializerConfig objects. diff --git a/tests/unit/backend/test_api_routes.py b/tests/unit/backend/test_api_routes.py index c9f8ffba5f..d859356309 100644 --- a/tests/unit/backend/test_api_routes.py +++ b/tests/unit/backend/test_api_routes.py @@ -803,6 +803,10 @@ def test_list_target_catalog(self, client: TestClient) -> None: mock_service = MagicMock() mock_service.list_target_catalog_async = AsyncMock( return_value=TargetCatalogResponse( + persistence={ + "definitions_enabled": True, + "api_keys_enabled": False, + }, items=[ { "target_type": "OpenAIChatTarget", diff --git a/tests/unit/backend/test_main.py b/tests/unit/backend/test_main.py index 19e1471e2a..ad3931bdba 100644 --- a/tests/unit/backend/test_main.py +++ b/tests/unit/backend/test_main.py @@ -34,6 +34,14 @@ async def test_lifespan_yields(self) -> None: "pyrit.backend.main.get_initializer_service", return_value=MagicMock(run_additional_initializers_async=AsyncMock()), ), + patch( + "pyrit.backend.main.get_target_service", + return_value=MagicMock( + configure_persistence=MagicMock(), + restore_persisted_targets_async=AsyncMock(), + ), + ), + patch("pyrit.backend.main.CentralMemory.get_memory_instance", return_value=MagicMock()), patch("pyrit.backend.main.setup_frontend"), ): async with lifespan(app): @@ -54,6 +62,14 @@ async def test_lifespan_warns_when_custom_initializers_allowed(self) -> None: "pyrit.backend.main.get_initializer_service", return_value=MagicMock(run_additional_initializers_async=AsyncMock()), ), + patch( + "pyrit.backend.main.get_target_service", + return_value=MagicMock( + configure_persistence=MagicMock(), + restore_persisted_targets_async=AsyncMock(), + ), + ), + patch("pyrit.backend.main.CentralMemory.get_memory_instance", return_value=MagicMock()), patch("pyrit.backend.main.setup_frontend"), patch.object(logging.getLogger("pyrit.backend.main"), "warning") as mock_warning, ): @@ -72,6 +88,14 @@ async def test_lifespan_populates_default_labels_from_operator_and_operation(sel "pyrit.backend.main.get_initializer_service", return_value=MagicMock(run_additional_initializers_async=AsyncMock()), ), + patch( + "pyrit.backend.main.get_target_service", + return_value=MagicMock( + configure_persistence=MagicMock(), + restore_persisted_targets_async=AsyncMock(), + ), + ), + patch("pyrit.backend.main.CentralMemory.get_memory_instance", return_value=MagicMock()), patch("pyrit.backend.main.setup_frontend"), ): async with lifespan(app): @@ -90,6 +114,14 @@ async def test_lifespan_reads_config_file_env_var(self) -> None: "pyrit.backend.main.get_initializer_service", return_value=MagicMock(run_additional_initializers_async=AsyncMock()), ), + patch( + "pyrit.backend.main.get_target_service", + return_value=MagicMock( + configure_persistence=MagicMock(), + restore_persisted_targets_async=AsyncMock(), + ), + ), + patch("pyrit.backend.main.CentralMemory.get_memory_instance", return_value=MagicMock()), patch("pyrit.backend.main.setup_frontend"), ): async with lifespan(app): diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index e226d33416..a1defc3f22 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -6,13 +6,14 @@ """ import os -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from pyrit.backend.models.targets import CreateTargetRequest from pyrit.backend.services.target_service import TargetService, get_target_service -from pyrit.models import ComponentIdentifier +from pyrit.memory.memory_interface import MemoryInterface +from pyrit.models import ComponentIdentifier, PersistedTarget from pyrit.prompt_target import PromptTarget, TargetCapabilities from pyrit.registry import TargetRegistry from unit.mocks import MockPromptTarget @@ -261,6 +262,19 @@ async def test_catalog_cold_and_warm_results_are_equal(self) -> None: assert cold == warm + async def test_catalog_reports_persistence_capabilities(self) -> None: + service = TargetService() + service.configure_persistence( + memory=MagicMock(spec=MemoryInterface), + definitions_enabled=True, + target_secret_key_vault_url="https://test.vault.azure.net", + ) + + result = await service.list_target_catalog_async() + + assert result.persistence.definitions_enabled is True + assert result.persistence.api_keys_enabled is True + async def test_catalog_refreshes_after_runtime_class_registration(self) -> None: service = TargetService() initial = await service.list_target_catalog_async() @@ -421,6 +435,110 @@ async def test_create_target_registers_in_registry(self, sqlite_instance) -> Non target_obj = service.get_target_object(target_registry_name=result.target_registry_name) assert target_obj is not None + async def test_create_target_with_api_key_is_memory_only_without_key_vault(self, sqlite_instance) -> None: + memory = MagicMock(spec=MemoryInterface) + service = TargetService() + service.configure_persistence( + memory=memory, + definitions_enabled=True, + target_secret_key_vault_url=None, + ) + + await service.create_target_async( + request=CreateTargetRequest( + type="OpenAIChatTarget", + params={ + "endpoint": "https://test.openai.azure.com/", + "model_name": "gpt-4o", + "api_key": "temporary-key", + }, + ) + ) + + memory.add_persisted_target.assert_not_called() + + async def test_create_target_stores_api_key_outside_definition(self, sqlite_instance) -> None: + memory = MagicMock(spec=MemoryInterface) + service = TargetService() + service.configure_persistence( + memory=memory, + definitions_enabled=True, + target_secret_key_vault_url="https://test.vault.azure.net", + ) + + with patch.object(service, "_set_api_key_async", new=AsyncMock()) as set_secret: + await service.create_target_async( + request=CreateTargetRequest( + type="OpenAIChatTarget", + params={ + "endpoint": "https://test.openai.azure.com/", + "model_name": "gpt-4o", + "api_key": "durable-key", + }, + ) + ) + + set_secret.assert_awaited_once() + persisted = memory.add_persisted_target.call_args.kwargs["target"] + assert isinstance(persisted, PersistedTarget) + assert persisted.parameters == { + "endpoint": "https://test.openai.azure.com/", + "model_name": "gpt-4o", + } + assert persisted.secret_name is not None + + +class TestRestorePersistedTargets: + async def test_restore_recreates_target_with_stored_registry_name(self) -> None: + memory = MagicMock(spec=MemoryInterface) + memory.get_persisted_targets.return_value = [ + PersistedTarget( + target_registry_name="saved-text-target", + target_type="TextTarget", + ) + ] + service = TargetService() + service.configure_persistence( + memory=memory, + definitions_enabled=True, + target_secret_key_vault_url=None, + ) + + await service.restore_persisted_targets_async() + + assert service.get_target_object(target_registry_name="saved-text-target") is not None + + async def test_restore_resolves_api_key_from_key_vault(self) -> None: + memory = MagicMock(spec=MemoryInterface) + memory.get_persisted_targets.return_value = [ + PersistedTarget( + target_registry_name="saved-target", + target_type="MockPromptTarget", + parameters={"endpoint": "https://example.test"}, + secret_name="saved-secret", + ) + ] + service = TargetService() + service.configure_persistence( + memory=memory, + definitions_enabled=True, + target_secret_key_vault_url="https://test.vault.azure.net", + ) + target = _mock_prompt_target() + + with ( + patch.object(service, "_get_api_key_async", new=AsyncMock(return_value="restored-key")), + patch.object(service._registry, "create_instance", return_value=target) as create, + ): + await service.restore_persisted_targets_async() + + assert create.call_args.args == ("MockPromptTarget",) + assert create.call_args.kwargs == { + "endpoint": "https://example.test", + "api_key": "restored-key", + } + assert service.get_target_object(target_registry_name="saved-target") is target + async def test_create_target_model_name_not_overridden_by_env_var(self, sqlite_instance) -> None: """Test that explicit model_name is not overridden by underlying_model env var.""" with patch.dict(os.environ, {"OPENAI_CHAT_UNDERLYING_MODEL": "gpt-4o"}): diff --git a/tests/unit/memory/test_persisted_target_memory.py b/tests/unit/memory/test_persisted_target_memory.py new file mode 100644 index 0000000000..3f72097f49 --- /dev/null +++ b/tests/unit/memory/test_persisted_target_memory.py @@ -0,0 +1,39 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from pyrit.memory.memory_models import PersistedTargetEntry +from pyrit.models import PersistedTarget + + +def test_persisted_target_round_trips_without_api_key(sqlite_instance) -> None: + target = PersistedTarget( + target_registry_name="OpenAIChatTarget-abc123", + target_type="OpenAIChatTarget", + parameters={ + "endpoint": "https://example.openai.azure.com", + "model_name": "gpt-4o", + }, + secret_name="pyrit-target-secret", + ) + + sqlite_instance.add_persisted_target(target=target) + + entries = sqlite_instance._query_entries(PersistedTargetEntry) + assert len(entries) == 1 + assert entries[0].to_domain_model() == target + assert sqlite_instance.get_persisted_targets() == [target] + assert "api_key" not in entries[0].parameters + + +def test_add_persisted_target_upserts_by_id(sqlite_instance) -> None: + target = PersistedTarget( + id="target-id", + target_registry_name="target-name", + target_type="TextTarget", + ) + sqlite_instance.add_persisted_target(target=target) + + updated = target.model_copy(update={"parameters": {"value": "updated"}}) + sqlite_instance.add_persisted_target(target=updated) + + assert sqlite_instance.get_persisted_targets() == [updated] diff --git a/tests/unit/setup/test_configuration_loader.py b/tests/unit/setup/test_configuration_loader.py index 51c626b4c4..5727a3aa6e 100644 --- a/tests/unit/setup/test_configuration_loader.py +++ b/tests/unit/setup/test_configuration_loader.py @@ -688,6 +688,7 @@ def test_load_with_overrides_preserves_all_explicit_config_fields(self, mock_def """ max_concurrent_scenario_runs: 7 allow_custom_initializers: true +target_secret_key_vault_url: https://targets.vault.azure.net/ server: url: http://localhost:8765/ startup_timeout: 45 @@ -701,6 +702,7 @@ def test_load_with_overrides_preserves_all_explicit_config_fields(self, mock_def assert config.max_concurrent_scenario_runs == 7 assert config.allow_custom_initializers is True + assert config.target_secret_key_vault_url == "https://targets.vault.azure.net" assert config.server == {"url": "http://localhost:8765/", "startup_timeout": 45} assert config.server_config == ServerConfig(url="http://localhost:8765", startup_timeout=45.0) assert config.extensions == {"feature_flag": "enabled"}