Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions .pyrit_conf_example
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,12 @@ operation: op_trash_panda
# env_akv_ref:
# - https://my-vault.vault.azure.net/secrets/my-pyrit-env

# Target API Key Vault
# --------------------
# Vault used by the backend to store API keys for database-backed targets.
# This is separate from env_akv_ref: it stores one secret per persisted target.
# target_api_key_vault_url: https://my-vault.vault.azure.net

# Max Concurrent Scenario Runs
# ----------------------------
# Maximum number of scenario runs that can execute concurrently in the backend.
Expand Down
80 changes: 80 additions & 0 deletions pyrit/auth/key_vault_secret_store.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

"""Azure Key Vault storage for persisted target API keys."""

from urllib.parse import urlparse


class KeyVaultSecretStore:
"""Store and retrieve target API keys using Azure Key Vault."""

def __init__(self, *, vault_url: str) -> None:
"""
Initialize the secret store.

Args:
vault_url (str): Base URL of the Azure Key Vault.
"""
self._vault_url = vault_url.rstrip("/")

async def set_secret_async(self, *, name: str, value: str) -> str:
"""
Store a secret and return its versioned Key Vault URI.

Args:
name (str): Key Vault secret name.
value (str): Secret value to store.

Returns:
str: The versioned URI returned by Azure Key Vault.

Raises:
RuntimeError: If Azure Key Vault does not return a secret URI.
"""
from azure.identity.aio import DefaultAzureCredential
from azure.keyvault.secrets.aio import SecretClient

credential = DefaultAzureCredential()
try:
async with SecretClient(vault_url=self._vault_url, credential=credential) as client:
secret = await client.set_secret(name, value)
finally:
await credential.close()
if not secret.id:
raise RuntimeError(f"Azure Key Vault did not return a URI for secret '{name}'.")
return secret.id

@staticmethod
async def get_secret_async(*, secret_uri: str) -> str:
"""
Retrieve a secret value from its versioned Key Vault URI.

Args:
secret_uri (str): A versioned or unversioned Azure Key Vault secret URI.

Returns:
str: The stored secret value.

Raises:
ValueError: If the URI is invalid or the secret has no value.
"""
from azure.identity.aio import DefaultAzureCredential
from azure.keyvault.secrets.aio import SecretClient

parsed = urlparse(secret_uri)
path_parts = [part for part in parsed.path.split("/") if part]
if parsed.scheme != "https" or not parsed.netloc or len(path_parts) not in (2, 3) or path_parts[0] != "secrets":
raise ValueError(f"Invalid Azure Key Vault secret URI: '{secret_uri}'.")

credential = DefaultAzureCredential()
try:
async with SecretClient(
vault_url=f"{parsed.scheme}://{parsed.netloc}", credential=credential
) as client:
secret = await client.get_secret(path_parts[1], version=path_parts[2] if len(path_parts) == 3 else None)
finally:
await credential.close()
if secret.value is None:
raise ValueError(f"Azure Key Vault secret '{secret_uri}' has no value.")
return secret.value
2 changes: 2 additions & 0 deletions pyrit/backend/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@
targets,
version,
)
from pyrit.backend.services.target_service import configure_target_service
from pyrit.setup.configuration_loader import ConfigurationLoader

# Check for development mode from environment variable
Expand All @@ -57,6 +58,7 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:

config = ConfigurationLoader.load_with_overrides(config_file=config_file)
await config.initialize_pyrit_async()
configure_target_service(api_key_vault_url=config.target_api_key_vault_url)

# Expose config values to route handlers via app.state
default_labels: dict[str, str] = {}
Expand Down
51 changes: 46 additions & 5 deletions pyrit/backend/services/target_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,10 +12,13 @@
- Retrieved from registry (pre-registered at startup or created earlier)
"""

import asyncio
import logging
import re
from functools import lru_cache
from typing import Any, Literal, cast

from pyrit.auth.key_vault_secret_store import KeyVaultSecretStore
from pyrit.backend.mappers.target_mappers import target_object_to_instance
from pyrit.backend.models.common import PaginationInfo
from pyrit.backend.models.targets import (
Expand All @@ -24,11 +27,15 @@
TargetCatalogResponse,
TargetListResponse,
)
from pyrit.memory import CentralMemory
from pyrit.models import OpenAITargetConfig
from pyrit.models.catalog.target import TargetInstance
from pyrit.registry import TargetRegistry

logger = logging.getLogger(__name__)

_target_api_key_vault_url: str | None = None


class TargetService:
"""
Expand All @@ -40,9 +47,10 @@ class TargetService:
service only orchestrates the request → registry hand-off.
"""

def __init__(self) -> None:
def __init__(self, *, api_key_vault_url: str | None = None) -> None:
"""Initialize the target service."""
self._registry = TargetRegistry.get_registry_singleton()
self._secret_store = KeyVaultSecretStore(vault_url=api_key_vault_url) if api_key_vault_url else None

def _build_instance_from_object(self, *, target_registry_name: str, target_obj: Any) -> TargetInstance:
"""
Expand Down Expand Up @@ -182,7 +190,6 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn
raise ValueError(
f"Target type '{request.type}' not found. Available types: {self._registry.get_class_names()}"
)

target_cls = self._registry.get_class(request.type)
params: dict[str, Any] = dict(request.params)

Expand All @@ -199,11 +206,38 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn
else:
target_obj = target_cls(**params)

self._registry.instances.register(target_obj)

target_registry_name = target_obj.get_identifier().unique_name
if request.type != "OpenAIChatTarget":
self._registry.instances.register(target_obj, name=target_registry_name)
return self._build_instance_from_object(target_registry_name=target_registry_name, target_obj=target_obj)

secret_uri = await self._persist_api_key_async(
target_registry_name=target_registry_name,
api_key=cast("str | None", params.pop("api_key", None)),
)
persisted_target = OpenAITargetConfig(
target_registry_name=target_registry_name,
endpoint=cast("str", params["endpoint"]),
model_name=cast("str", params["model_name"]),
auth_mode=request.auth_mode,
api_key_secret_uri=secret_uri,
)
memory = CentralMemory.get_memory_instance()
await asyncio.to_thread(memory.add_openai_target_config, target=persisted_target)
self._registry.instances.register(target_obj, name=target_registry_name)

return self._build_instance_from_object(target_registry_name=target_registry_name, target_obj=target_obj)

async def _persist_api_key_async(self, *, target_registry_name: str, api_key: str | None) -> str | None:
if not api_key:
return None
if self._secret_store is None:
raise ValueError(
"Target API key persistence requires 'target_api_key_vault_url' in the PyRIT configuration."
)
secret_name = re.sub(r"[^0-9A-Za-z-]", "-", f"pyrit-target-{target_registry_name}")[:127]
return await self._secret_store.set_secret_async(name=secret_name, value=api_key)

def _has_reference_params(self, *, target_type: str) -> bool:
"""
Return True if the target type's build contract references other registry
Expand All @@ -229,4 +263,11 @@ def get_target_service() -> TargetService:
Returns:
The singleton TargetService instance.
"""
return TargetService()
return TargetService(api_key_vault_url=_target_api_key_vault_url)


def configure_target_service(*, api_key_vault_url: str | None) -> None:
"""Configure the cached backend target service."""
global _target_api_key_vault_url
_target_api_key_vault_url = api_key_vault_url
get_target_service.cache_clear()
39 changes: 39 additions & 0 deletions pyrit/memory/alembic/versions/4a7c9e1b3d5f_add_targets_table.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

"""
Add persisted target configurations.

Revision ID: 4a7c9e1b3d5f
Revises: 3f6e8a0c2d4b
Create Date: 2026-07-17 00:00:00.000000
"""

from collections.abc import Sequence

import sqlalchemy as sa
from alembic import op

# revision identifiers, used by Alembic.
revision: str = "4a7c9e1b3d5f"
down_revision: str | None = "3f6e8a0c2d4b"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None


def upgrade() -> None:
"""Create the Targets table."""
op.create_table(
"Targets",
sa.Column("target_registry_name", sa.String(), nullable=False),
sa.Column("endpoint", sa.String(), nullable=False),
sa.Column("model_name", sa.String(), nullable=False),
sa.Column("auth_mode", sa.String(), nullable=False),
sa.Column("api_key_secret_uri", sa.String(), nullable=True),
sa.PrimaryKeyConstraint("target_registry_name"),
)


def downgrade() -> None:
"""Drop the Targets table."""
op.drop_table("Targets")
11 changes: 11 additions & 0 deletions pyrit/memory/memory_interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
ConversationEntry,
ConverterIdentifierEntry,
EmbeddingDataEntry,
OpenAITargetConfigEntry,
PromptConverterIdentifierEntry,
PromptMemoryEntry,
ScenarioIdentifierEntry,
Expand Down Expand Up @@ -63,6 +64,7 @@
IdentifierType,
Message,
MessagePiece,
OpenAITargetConfig,
ScenarioIdentifier,
ScenarioResult,
Score,
Expand Down Expand Up @@ -149,6 +151,15 @@ def enable_embedding(self, embedding_model: Any | None = None) -> None:

self.memory_embedding = default_memory_embedding_factory(embedding_model=embedding_model)

def add_openai_target_config(self, *, target: OpenAITargetConfig) -> None:
"""Persist a sanitized, reconstructable target configuration."""
self._insert_entry(OpenAITargetConfigEntry.from_domain_model(target))

def get_openai_target_configs(self) -> Sequence[OpenAITargetConfig]:
"""Return all persisted target configurations."""
entries = self._query_entries(OpenAITargetConfigEntry, order_by=OpenAITargetConfigEntry.target_registry_name)
return [entry.to_domain_model() for entry in entries]

def disable_embedding(self) -> None:
"""
Disable embedding functionality for the memory interface.
Expand Down
42 changes: 42 additions & 0 deletions pyrit/memory/memory_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@
ConverterIdentifier,
EvaluationIdentifier,
MessagePiece,
OpenAITargetConfig,
PromptDataType,
ScenarioEvaluationIdentifier,
ScenarioIdentifier,
Expand Down Expand Up @@ -419,6 +420,47 @@ def __init_subclass__(cls, **kwargs: Any) -> None:
)


class OpenAITargetConfigEntry(DomainBackedEntry[OpenAITargetConfig]):
"""Persistence projection for a reconstructable target configuration."""

__tablename__ = "Targets"
__table_args__ = {"extend_existing": True}

target_registry_name: Mapped[str] = mapped_column(String, primary_key=True)
endpoint: Mapped[str] = mapped_column(String, nullable=False)
model_name: Mapped[str] = mapped_column(String, nullable=False)
auth_mode: Mapped[Literal["api_key", "identity"]] = mapped_column(String, nullable=False)
api_key_secret_uri: Mapped[str | None] = mapped_column(String, nullable=True)

@classmethod
def from_domain_model(cls, domain_model: OpenAITargetConfig) -> Self:
"""
Build an unsaved target row from its domain model.

Args:
domain_model (OpenAITargetConfig): The target configuration to persist.

Returns:
Self: An unsaved target configuration row.
"""
return cls(**domain_model.model_dump())

def to_domain_model(self) -> OpenAITargetConfig:
"""
Build the canonical domain model represented by this row.

Returns:
OpenAITargetConfig: The reconstructed target configuration.
"""
return OpenAITargetConfig(
target_registry_name=self.target_registry_name,
endpoint=self.endpoint,
model_name=self.model_name,
auth_mode=self.auth_mode,
api_key_secret_uri=self.api_key_secret_uri,
)


T = TypeVar("T", bound=ComponentIdentifier)


Expand Down
2 changes: 2 additions & 0 deletions pyrit/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,7 @@
RegistryReference,
display_choices,
)
from pyrit.models.openai_target_config import OpenAITargetConfig
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
Expand Down Expand Up @@ -181,6 +182,7 @@
"ObjectiveTargetEvaluationIdentifier",
"Parameter",
"ParameterDestination",
"OpenAITargetConfig",
"PromptDataType",
"PromptResponseError",
"QuestionAnsweringDataset",
Expand Down
26 changes: 26 additions & 0 deletions pyrit/models/openai_target_config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

"""Domain model for OpenAI target configurations persisted in memory."""

from typing import Literal

from pydantic import BaseModel, model_validator


class OpenAITargetConfig(BaseModel):
"""A reconstructable OpenAI target configuration with a secret reference."""

target_registry_name: str
endpoint: str
model_name: str
auth_mode: Literal["api_key", "identity"] = "api_key"
api_key_secret_uri: str | None = None

@model_validator(mode="after")
def _validate_secret_storage(self) -> "OpenAITargetConfig":
if self.auth_mode == "api_key" and not self.api_key_secret_uri:
raise ValueError("API key authentication requires an api_key_secret_uri.")
if self.auth_mode == "identity" and self.api_key_secret_uri:
raise ValueError("Identity authentication must not reference an API key secret.")
return self
Loading
Loading