diff --git a/doc/code/setup/1_configuration.ipynb b/doc/code/setup/1_configuration.ipynb index fc18179efd..1c2b8fc018 100644 --- a/doc/code/setup/1_configuration.ipynb +++ b/doc/code/setup/1_configuration.ipynb @@ -197,7 +197,24 @@ "\n", " ```bash\n", " az login\n", - " ```" + " ```\n", + "\n", + "### Choosing between key and Entra auth\n", + "\n", + "Authentication is explicit. `auth_mode=\"api_key\"` (the default) resolves a credential in this order: a token provider callable passed as `api_key`, an explicit `api_key` string, then the target's API key environment variable. If none of those yield a key it raises a `ValueError` rather than guessing.\n", + "\n", + "Pass `auth_mode=\"identity\"` when you want Entra auth. It skips the key and the environment variable entirely, so an unrelated key in your `.env` can no longer override your choice. It raises a `ValueError` if the endpoint is not a recognized Azure endpoint rather than minting a token for an unknown host.\n", + "\n", + "```python\n", + "target = OpenAIChatTarget(\n", + " endpoint=os.environ[\"OPENAI_CHAT_ENDPOINT\"],\n", + " auth_mode=\"identity\",\n", + ")\n", + "```\n", + "\n", + "`auth_mode` applies to the OpenAI targets, `AzureMLChatTarget`, and `PromptShieldTarget`. `AzureBlobStorageTarget` accepts the same `auth_mode=\"identity\"` to bypass its SAS token sources; it defaults to selecting a credential automatically.\n", + "\n", + "> **Migration note.** Earlier versions silently minted an Entra token when no key was found and the endpoint looked like an Azure host. That implicit fallback has been removed because it made an explicit `auth_mode=\"identity\"` indistinguishable from \"no key configured\". If you relied on it, pass `auth_mode=\"identity\"` (or a token provider as `api_key`, as the examples above do)." ] }, { diff --git a/doc/code/setup/1_configuration.py b/doc/code/setup/1_configuration.py index 2504b9d0d7..44954cb9b6 100644 --- a/doc/code/setup/1_configuration.py +++ b/doc/code/setup/1_configuration.py @@ -108,6 +108,23 @@ # ```bash # az login # ``` +# +# ### Choosing between key and Entra auth +# +# Authentication is explicit. `auth_mode="api_key"` (the default) resolves a credential in this order: a token provider callable passed as `api_key`, an explicit `api_key` string, then the target's API key environment variable. If none of those yield a key it raises a `ValueError` rather than guessing. +# +# Pass `auth_mode="identity"` when you want Entra auth. It skips the key and the environment variable entirely, so an unrelated key in your `.env` can no longer override your choice. It raises a `ValueError` if the endpoint is not a recognized Azure endpoint rather than minting a token for an unknown host. +# +# ```python +# target = OpenAIChatTarget( +# endpoint=os.environ["OPENAI_CHAT_ENDPOINT"], +# auth_mode="identity", +# ) +# ``` +# +# `auth_mode` applies to the OpenAI targets, `AzureMLChatTarget`, and `PromptShieldTarget`. `AzureBlobStorageTarget` accepts the same `auth_mode="identity"` to bypass its SAS token sources; it defaults to selecting a credential automatically. +# +# > **Migration note.** Earlier versions silently minted an Entra token when no key was found and the endpoint looked like an Azure host. That implicit fallback has been removed because it made an explicit `auth_mode="identity"` indistinguishable from "no key configured". If you relied on it, pass `auth_mode="identity"` (or a token provider as `api_key`, as the examples above do). # %% [markdown] # ## Choosing a database diff --git a/pyrit/auth/openai_auth.py b/pyrit/auth/openai_auth.py index aa281cc199..927d31b1da 100644 --- a/pyrit/auth/openai_auth.py +++ b/pyrit/auth/openai_auth.py @@ -6,6 +6,7 @@ from pyrit.auth.azure_auth import ensure_async_token_provider, get_azure_openai_auth, is_azure_openai_endpoint from pyrit.common import default_values +from pyrit.common.auth_mode import AuthMode def resolve_openai_auth( @@ -13,21 +14,37 @@ def resolve_openai_auth( endpoint: str, api_key: str | Callable[[], str | Awaitable[str]] | None, api_key_environment_variable: str, + auth_mode: AuthMode = "api_key", ) -> str | Callable[[], Awaitable[str]]: """ - Resolve OpenAI authentication from a key, environment variable, or Azure Entra fallback. + Resolve OpenAI authentication from an explicit identity choice, a key, or an environment variable. Args: endpoint (str): The OpenAI-compatible endpoint URL. api_key (str | Callable[[], str | Awaitable[str]] | None): The explicit API key or token provider. api_key_environment_variable (str): Environment variable to use when ``api_key`` is not provided. + auth_mode (AuthMode): ``"identity"`` authenticates with a Microsoft Entra ID token and ignores + ``api_key`` and its environment variable entirely. ``"api_key"`` (the default) resolves a + token-provider callable, then an explicit key, then the environment variable. Returns: str | Callable[[], Awaitable[str]]: API key string or async-compatible token provider. Raises: - ValueError: If no key is provided and the endpoint is not a recognized Azure OpenAI endpoint. + ValueError: If identity auth is requested for an endpoint that is not a recognized Azure + OpenAI endpoint, or if ``"api_key"`` auth is requested and no key is available. """ + # Identity is an explicit caller choice, so it must never be silently downgraded to a key + # that merely happens to be present in the environment. + if auth_mode == "identity": + if not is_azure_openai_endpoint(endpoint): + raise ValueError( + f"Identity-based authentication requires a recognized Azure OpenAI / AI Foundry endpoint, " + f"but got '{endpoint}'. Use api_key authentication for this endpoint, or pass your own " + "token provider callable as api_key." + ) + return get_azure_openai_auth(endpoint) + if api_key is not None and callable(api_key): return cast("str | Callable[[], Awaitable[str]]", ensure_async_token_provider(api_key)) @@ -37,10 +54,8 @@ def resolve_openai_auth( if api_key_value: return api_key_value - if is_azure_openai_endpoint(endpoint): - return get_azure_openai_auth(endpoint) - raise ValueError( - f"Environment variable {api_key_environment_variable} is required for non-Azure endpoints. " - "For recognized Azure OpenAI / AI Foundry endpoints, Entra ID authentication is used automatically." + f"No API key available for endpoint '{endpoint}'. Set the {api_key_environment_variable} environment " + 'variable, pass api_key explicitly, or pass auth_mode="identity" to authenticate with Microsoft ' + "Entra ID on a recognized Azure OpenAI / AI Foundry endpoint." ) diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index 8d1176ed07..d065f5b1f2 100644 --- a/pyrit/backend/services/target_service.py +++ b/pyrit/backend/services/target_service.py @@ -16,7 +16,7 @@ import logging import uuid from functools import lru_cache -from typing import Any, Literal +from typing import Any from pyrit.backend.mappers.target_mappers import target_object_to_instance from pyrit.backend.models.common import PaginationInfo @@ -27,6 +27,7 @@ TargetTypeResponse, ) from pyrit.common import REQUIRED_VALUE +from pyrit.common.auth_mode import AUTH_MODES, AuthMode from pyrit.models.catalog.target import TargetInstance from pyrit.models.parameter import Parameter from pyrit.registry import TargetRegistry @@ -42,6 +43,9 @@ "PromptShieldTarget": frozenset({"endpoint"}), } +# Constructor parameter through which a target accepts an explicit auth mode. +_AUTH_MODE_PARAM = "auth_mode" + class TargetService: """ @@ -138,7 +142,7 @@ def get_target_object(self, *, target_registry_name: str) -> Any | None: return self._registry.instances.get(target_registry_name) @staticmethod - def _get_supported_auth_modes(auth_modes: tuple[str, ...]) -> list[Literal["api_key", "identity"]]: + def _get_supported_auth_modes(auth_modes: tuple[str, ...]) -> list[AuthMode]: """ Validate and narrow registry authentication modes for the type response. @@ -146,17 +150,16 @@ def _get_supported_auth_modes(auth_modes: tuple[str, ...]) -> list[Literal["api_ auth_modes (tuple[str, ...]): Authentication modes declared by a target class. Returns: - list[Literal["api_key", "identity"]]: Validated authentication modes. + list[AuthMode]: Validated authentication modes. Raises: ValueError: If a target class declares an unsupported authentication mode. """ - supported_auth_modes: list[Literal["api_key", "identity"]] = [] + supported_auth_modes: list[AuthMode] = [] for auth_mode in auth_modes: - if auth_mode == "api_key" or auth_mode == "identity": - supported_auth_modes.append(auth_mode) - continue - raise ValueError(f"Unsupported target authentication mode: {auth_mode!r}") + if auth_mode not in AUTH_MODES: + raise ValueError(f"Unsupported target authentication mode: {auth_mode!r}") + supported_auth_modes.append(auth_mode) return supported_auth_modes def _project_target_parameters(self, *, target_type: str, parameters: tuple[Parameter, ...]) -> list[Parameter]: @@ -220,11 +223,14 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn reference resolution, and construction are owned by the ``TargetRegistry``. Endpoint trust and identity token minting are owned by the target classes themselves. This service only enforces the - request-level auth contract: for ``identity`` it confirms the target - supports it and omits the api_key plus any registry-flagged - identity-conflicting parameters so the target validates its own - endpoint and authenticates itself. The response is built before the - target is registered, so a failed request leaves no registered target. + request-level auth contract: it rejects an ``auth_mode`` smuggled through + ``params``, and for ``identity`` it confirms the target supports it and + omits the api_key plus any registry-flagged identity-conflicting + parameters so the target validates its own endpoint and authenticates + itself. The request-level ``auth_mode`` is authoritative and is forwarded + via ``get_auth_mode_parameters`` to targets that accept it, so the choice + is explicit rather than inferred from a missing key. The response is built + before the target is registered, so a failed request leaves no registered target. Args: request: The create target request with type, params, and auth_mode. @@ -233,8 +239,9 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn TargetInstance with the new target's details. Raises: - ValueError: If the target type is not registered or identity auth is - requested but unsupported by the target type. Construction errors + ValueError: If the target type is not registered, ``params`` carries an + ``auth_mode`` that conflicts with the request-level choice, or identity + auth is requested but unsupported by the target type. Construction errors (unknown params, incompatible inner targets, unrecognized identity endpoints) are raised by the registry / target classes. """ @@ -246,10 +253,20 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn target_cls = self._registry.get_class(request.type) params: dict[str, Any] = dict(request.params) + # auth_mode is also a constructor parameter, so the registry would otherwise accept it + # inside params as a second, competing channel that bypasses the checks below. + params_auth_mode = params.get(_AUTH_MODE_PARAM) + if params_auth_mode is not None and params_auth_mode != request.auth_mode: + raise ValueError( + f"Conflicting authentication modes: request auth_mode is '{request.auth_mode}' but " + f"params['{_AUTH_MODE_PARAM}'] is '{params_auth_mode}'. Set the request-level auth_mode only." + ) + if request.auth_mode == "identity": if "identity" not in target_cls.supported_auth_modes: raise ValueError(f"Target type '{request.type}' does not support identity-based authentication.") - # Omit any api_key so the target validates its own endpoint and authenticates itself. + # Omitting the key alone is ambiguous — every auth resolver reads the api-key env var + # when no key is passed — so also state the choice explicitly where the target accepts it. params.pop("api_key", None) # Omit any other parameter the registry metadata marks as conflicting with # identity-based auth (e.g. AzureBlobStorageTarget's sas_token), so a caller @@ -259,6 +276,7 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn for parameter in metadata.parameters: if parameter.identity_conflicting: params.pop(parameter.name, None) + params.update(target_cls.get_auth_mode_parameters(auth_mode=request.auth_mode)) # LEGACY COMPATIBILITY: The current configuration UI omits the name. diff --git a/pyrit/common/__init__.py b/pyrit/common/__init__.py index 748155ebdc..afad22871a 100644 --- a/pyrit/common/__init__.py +++ b/pyrit/common/__init__.py @@ -29,6 +29,7 @@ reset_default_values, set_default_value, ) + from pyrit.common.auth_mode import AUTH_MODES, AuthMode from pyrit.common.brick_contract import enforce_keyword_only_init, forward_init_parameters from pyrit.common.default_values import get_non_required_value, get_required_value from pyrit.common.deprecation import print_deprecation_message @@ -48,6 +49,8 @@ _LAZY_EXPORTS: dict[str, str | tuple[str, str | None]] = { "apply_defaults": "pyrit.common.apply_defaults", "apply_defaults_to_method": "pyrit.common.apply_defaults", + "AUTH_MODES": "pyrit.common.auth_mode", + "AuthMode": "pyrit.common.auth_mode", "combine_dict": "pyrit.common.utils", "combine_list": "pyrit.common.utils", "DefaultValueScope": "pyrit.common.apply_defaults", diff --git a/pyrit/common/auth_mode.py b/pyrit/common/auth_mode.py new file mode 100644 index 0000000000..806be4ef95 --- /dev/null +++ b/pyrit/common/auth_mode.py @@ -0,0 +1,29 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +""" +Canonical credential-selection mode shared by auth resolvers and targets. + +Lives in ``pyrit.common`` rather than ``pyrit.prompt_target`` because +``pyrit.auth`` resolvers consume it and must not depend on the target layer. + +Not to be confused with ``pyrit.cli._auth.AuthMode``, which names the *Azure +credential flow* the CLI should use (``"auto"`` / ``"azure_cli"`` / +``"device_code"`` / ``"none"``). +""" + +from typing import Literal + +__all__ = ["AUTH_MODES", "AuthMode"] + +#: How a component chooses its credential. +#: +#: ``api_key`` resolves a key from the explicit argument or the component's API key +#: environment variable, falling back to an ambient Azure identity only when neither +#: is available. ``identity`` is an explicit caller choice: the key and its +#: environment variable are skipped entirely and the component authenticates with an +#: ambient Azure identity (e.g. a Microsoft Entra ID token minted for its own +#: endpoint). +AuthMode = Literal["api_key", "identity"] + +AUTH_MODES: tuple[AuthMode, ...] = ("api_key", "identity") diff --git a/pyrit/embedding/openai_text_embedding.py b/pyrit/embedding/openai_text_embedding.py index 5a84a880d6..c64cd8e2fa 100644 --- a/pyrit/embedding/openai_text_embedding.py +++ b/pyrit/embedding/openai_text_embedding.py @@ -10,6 +10,7 @@ from pyrit.auth import resolve_openai_auth from pyrit.common import default_values +from pyrit.common.auth_mode import AuthMode from pyrit.models import ( EmbeddingData, EmbeddingResponse, @@ -34,15 +35,14 @@ def __init__( api_key: str | Callable[[], str | Awaitable[str]] | None = None, endpoint: str | None = None, model_name: str | None = None, + auth_mode: AuthMode = "api_key", ) -> None: """ Initialize text embedding client for Azure OpenAI or platform OpenAI. Args: api_key: The API key (string) or token provider (callable) for authentication. - For recognized Azure OpenAI / AI Foundry endpoints, if no API key is provided - (via parameter or environment variable), Entra ID authentication is used automatically. - You can also explicitly pass a token provider from pyrit.auth + You can pass a token provider from pyrit.auth (e.g., get_azure_openai_auth(endpoint) for async). Defaults to OPENAI_EMBEDDING_KEY environment variable. endpoint: The API endpoint URL. @@ -51,10 +51,15 @@ def __init__( Defaults to OPENAI_EMBEDDING_ENDPOINT environment variable. model_name: The model/deployment name (e.g., "text-embedding-3-small"). Defaults to OPENAI_EMBEDDING_MODEL environment variable. + auth_mode: ``"identity"`` authenticates with a Microsoft Entra ID token and ignores + ``api_key`` and its environment variable entirely. ``"api_key"`` (the default) + resolves a token-provider callable, then an explicit key, then the environment + variable. Raises: - ValueError: If no API key is provided (via parameter or environment variable) and the - endpoint is not a recognized Azure OpenAI / AI Foundry endpoint. + ValueError: If identity auth is requested for an endpoint that is not a recognized + Azure OpenAI / AI Foundry endpoint, or if ``"api_key"`` auth is requested and no + key is available via parameter or environment variable. """ endpoint = default_values.get_required_value( env_var_name=self.ENDPOINT_URI_ENVIRONMENT_VARIABLE, passed_value=endpoint @@ -67,6 +72,7 @@ def __init__( endpoint=endpoint, api_key=api_key, api_key_environment_variable=self.API_KEY_ENVIRONMENT_VARIABLE, + auth_mode=auth_mode, ) self._async_client = AsyncOpenAI( api_key=async_api_key, diff --git a/pyrit/prompt_target/azure_ml_chat_target.py b/pyrit/prompt_target/azure_ml_chat_target.py index ee0d83ff93..133963ed16 100644 --- a/pyrit/prompt_target/azure_ml_chat_target.py +++ b/pyrit/prompt_target/azure_ml_chat_target.py @@ -69,6 +69,7 @@ def __init__( *, endpoint: str | None = None, api_key: str | Callable[[], str | Awaitable[str]] | None = None, + auth_mode: AuthMode = "api_key", model_name: str = "", max_new_tokens: int = 400, temperature: float = 1.0, @@ -90,6 +91,10 @@ def __init__( to authenticate with Microsoft Entra ID against an AML managed online endpoint. Synchronous providers are automatically wrapped via ``ensure_async_token_provider``. Defaults to the value of the ``AZURE_ML_KEY`` environment variable. + auth_mode (AuthMode): Explicitly selects how to authenticate. ``"identity"`` mints a + Microsoft Entra ID token for the endpoint and ignores ``api_key`` and the + ``AZURE_ML_KEY`` environment variable entirely; it requires a recognized AML managed + online endpoint. Defaults to ``"api_key"``, which resolves the key as described above. model_name (str): The name of the model being used (e.g., "Llama-3.2-3B-Instruct"). Used for identification purposes. Defaults to empty string. max_new_tokens (int): The maximum number of tokens to generate in the response. @@ -124,7 +129,7 @@ def __init__( custom_configuration=custom_configuration, ) - self._initialize_vars(endpoint=endpoint, api_key=api_key) + self._initialize_vars(endpoint=endpoint, api_key=api_key, auth_mode=auth_mode) validate_temperature(temperature) validate_top_p(top_p) @@ -135,6 +140,19 @@ def __init__( self._repetition_penalty = repetition_penalty self._extra_parameters = param_kwargs + @classmethod + def get_auth_mode_parameters(cls, *, auth_mode: AuthMode) -> dict[str, object]: + """ + Preserve explicit authentication intent through target construction. + + Args: + auth_mode (AuthMode): Authentication mode selected by the caller. + + Returns: + dict[str, object]: Constructor parameters that enforce the mode. + """ + return {"auth_mode": auth_mode} + def _build_identifier(self) -> ComponentIdentifier: """ Build the identifier with Azure ML-specific parameters. @@ -153,8 +171,10 @@ def _build_identifier(self) -> ComponentIdentifier: def _initialize_vars( self, + *, endpoint: str | None = None, api_key: str | Callable[[], str | Awaitable[str]] | None = None, + auth_mode: AuthMode = "api_key", ) -> None: """ Set the endpoint and key for accessing the Azure ML model. Use this function to manually @@ -175,20 +195,38 @@ def _initialize_vars( The API key for accessing the Azure ML endpoint, or a callable which returns a bearer token, or None to fall back to the ``AZURE_ML_KEY`` env variable. + auth_mode (AuthMode): ``"identity"`` mints a Microsoft Entra ID token and ignores + ``api_key`` and the ``AZURE_ML_KEY`` environment variable entirely. ``"api_key"`` + (the default) resolves a token-provider callable, then an explicit key, then the + environment variable. Raises: - ValueError: If no api_key is supplied (via parameter or environment - variable) and the endpoint is not a recognized Azure ML managed - online endpoint for which Entra ID authentication can be used. + ValueError: If identity auth is requested for an endpoint that is not a recognized + Azure ML managed online endpoint, or if ``"api_key"`` auth is requested and no key + is available via parameter or environment variable. """ self._endpoint = default_values.get_required_value( env_var_name=self.endpoint_uri_environment_variable, passed_value=endpoint ) + self._api_key_provider: Callable[[], Awaitable[str]] | None + + # Identity is an explicit caller choice, so it must never be silently downgraded to a key + # that merely happens to be present in the environment. + if auth_mode == "identity": + if not is_azure_ml_endpoint(self._endpoint): + raise ValueError( + "Identity-based authentication requires a recognized Azure ML managed online endpoint " + f"(*.inference.ml.azure.com), but got '{self._endpoint}'. Use api_key authentication for " + "this endpoint, or pass your own token provider callable as api_key." + ) + self._api_key_provider = self._build_azure_ml_token_provider() + self._api_key = "" + return if callable(api_key): normalized = ensure_async_token_provider(api_key) provider = cast("Callable[[], Awaitable[str]]", normalized) - self._api_key_provider: Callable[[], Awaitable[str]] | None = provider + self._api_key_provider = provider self._api_key = "" return @@ -200,21 +238,23 @@ def _initialize_vars( self._api_key = api_key_value return - # No key supplied: fall back to Microsoft Entra ID, but only for a - # recognized AML managed online endpoint so a bearer token is never - # minted for an arbitrary host. - if is_azure_ml_endpoint(self._endpoint): - normalized = ensure_async_token_provider(get_azure_async_token_provider(self._AZURE_ML_SCOPE)) - self._api_key_provider = cast("Callable[[], Awaitable[str]]", normalized) - self._api_key = "" - return - raise ValueError( - f"Environment variable {self.api_key_environment_variable} is required unless the endpoint is a " - "recognized Azure ML managed online endpoint (*.inference.ml.azure.com), for which Entra ID " - "authentication is used automatically. Pass an api_key or a token provider callable instead." + f"No API key available for endpoint '{self._endpoint}'. Set the " + f"{self.api_key_environment_variable} environment variable, pass api_key (a key or a token " + 'provider callable), or pass auth_mode="identity" to authenticate with Microsoft Entra ID on a ' + "recognized Azure ML managed online endpoint (*.inference.ml.azure.com)." ) + def _build_azure_ml_token_provider(self) -> Callable[[], Awaitable[str]]: + """ + Build an async Entra ID token provider scoped to Azure Machine Learning. + + Returns: + Callable[[], Awaitable[str]]: An async-compatible bearer token provider. + """ + normalized = ensure_async_token_provider(get_azure_async_token_provider(self._AZURE_ML_SCOPE)) + return cast("Callable[[], Awaitable[str]]", normalized) + @pyrit_target_retry @limit_requests_per_minute async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Message]) -> list[Message]: diff --git a/pyrit/prompt_target/common/prompt_target.py b/pyrit/prompt_target/common/prompt_target.py index cf04d20e52..ef656aaef3 100644 --- a/pyrit/prompt_target/common/prompt_target.py +++ b/pyrit/prompt_target/common/prompt_target.py @@ -4,10 +4,11 @@ import abc import logging from collections.abc import Mapping, Sequence -from typing import Any, ClassVar, Literal, final +from typing import Any, ClassVar, final from pyrit.common.async_compatibility import legacy_sync_override from pyrit.common.attack_result_scope import get_current_attack_result_id +from pyrit.common.auth_mode import AuthMode from pyrit.common.deprecation import print_deprecation_message from pyrit.memory import CentralMemory, MemoryInterface from pyrit.message_normalizer import MessageListNormalizer @@ -34,12 +35,9 @@ logger = logging.getLogger(__name__) -# Authentication modes a target can expose to target type discovery and creation APIs. -# ``api_key`` passes a key (from params or the target's env var); ``identity`` -# omits the key so the target authenticates itself via an ambient Azure identity -# (e.g. minting a Microsoft Entra ID token for its own endpoint, or falling back -# to ``DefaultAzureCredential``). -AuthMode = Literal["api_key", "identity"] +# ``AuthMode`` is imported above (and re-exported from this module for backward +# compatibility); it is canonically defined in ``pyrit.common.auth_mode`` so +# ``pyrit.auth`` resolvers can consume it without depending on the target layer. class PromptTarget(Identifiable): diff --git a/pyrit/prompt_target/openai/openai_target.py b/pyrit/prompt_target/openai/openai_target.py index fde956649b..aa4b6810e8 100644 --- a/pyrit/prompt_target/openai/openai_target.py +++ b/pyrit/prompt_target/openai/openai_target.py @@ -87,6 +87,7 @@ def __init__( model_name: str | None = None, endpoint: str | None = None, api_key: str | Callable[[], str | Awaitable[str]] | None = None, + auth_mode: AuthMode = "api_key", headers: str | None = None, max_requests_per_minute: int | None = None, httpx_client_kwargs: dict[str, Any] | None = None, @@ -102,12 +103,15 @@ def __init__( endpoint (str, Optional): The target URL for the OpenAI service. api_key (str | Callable[[], str | Awaitable[str]], Optional): The API key for accessing the OpenAI service, or a callable that returns an access token (sync or async). - For recognized Azure OpenAI / AI Foundry endpoints, if no API key is provided - (via parameter or environment variable), Entra ID authentication is used automatically. - You can also explicitly pass a token provider from pyrit.auth + You can pass a token provider from pyrit.auth (e.g., get_azure_openai_auth(endpoint) for async, or get_azure_token_provider(scope) for sync). Synchronous token providers are automatically wrapped to work with async clients. Defaults to the target-specific API key environment variable. + auth_mode (AuthMode, Optional): Explicitly selects how to authenticate. ``"identity"`` + authenticates with a Microsoft Entra ID token for the endpoint and ignores ``api_key`` + and its environment variable entirely; it requires a recognized Azure OpenAI / + AI Foundry endpoint. Defaults to ``"api_key"``, which resolves the key as described + above. headers (str, Optional): Extra headers of the endpoint (JSON). max_requests_per_minute (int, Optional): Number of requests the target can handle per minute before hitting a rate limit. The number of requests sent to the target @@ -122,8 +126,9 @@ def __init__( this target instance. If None, uses the class-level defaults. Defaults to None. Raises: - ValueError: If no API key is provided (via parameter or environment variable) and the - endpoint is not a recognized Azure OpenAI / AI Foundry endpoint. + ValueError: If identity auth is requested for an endpoint that is not a recognized + Azure OpenAI / AI Foundry endpoint, or if ``"api_key"`` auth is requested and no + key is available via parameter or environment variable. """ self._headers: dict[str, str] = {} self._httpx_client_kwargs = httpx_client_kwargs or {} @@ -157,10 +162,24 @@ def __init__( endpoint=endpoint_value, api_key=api_key, api_key_environment_variable=self.api_key_environment_variable, + auth_mode=auth_mode, ) self._initialize_openai_client() + @classmethod + def get_auth_mode_parameters(cls, *, auth_mode: AuthMode) -> dict[str, object]: + """ + Preserve explicit authentication intent through target construction. + + Args: + auth_mode (AuthMode): Authentication mode selected by the caller. + + Returns: + dict[str, object]: Constructor parameters that enforce the mode. + """ + return {"auth_mode": auth_mode} + @staticmethod def _parse_request_headers(value: object) -> dict[str, str]: """ diff --git a/pyrit/prompt_target/prompt_shield_target.py b/pyrit/prompt_target/prompt_shield_target.py index fa9c04db81..662a5d8d49 100644 --- a/pyrit/prompt_target/prompt_shield_target.py +++ b/pyrit/prompt_target/prompt_shield_target.py @@ -68,6 +68,7 @@ def __init__( *, endpoint: str | None = None, api_key: str | Callable[[], str] | None = None, + auth_mode: AuthMode = "api_key", api_version: str | None = "2024-09-01", field: PromptShieldEntryField | None = None, max_requests_per_minute: int | None = None, @@ -87,6 +88,11 @@ def __init__( token provider, pass one from pyrit.auth (e.g., get_azure_token_provider('https://cognitiveservices.azure.com/.default')). Defaults to the `API_KEY_ENVIRONMENT_VARIABLE` environment variable. + auth_mode (AuthMode, Optional): Explicitly selects how to authenticate. ``"identity"`` + mints a Microsoft Entra ID token for the endpoint and ignores ``api_key`` and the + `API_KEY_ENVIRONMENT_VARIABLE` environment variable entirely; it requires a + recognized Azure Content Safety endpoint. Defaults to ``"api_key"``, which resolves + the key as described above. api_version (str, Optional): The version of the Azure Content Safety API. Defaults to "2024-09-01". field (PromptShieldEntryField, Optional): If "userPrompt", all input is sent to the userPrompt field. If "documents", all input is sent to the documents field. If None, the input is parsed to separate @@ -98,8 +104,9 @@ def __init__( this target instance. Defaults to None. Raises: - ValueError: If the endpoint value is not provided, or if no API key is - provided for a non-Azure Content Safety endpoint. + ValueError: If the endpoint value is not provided, if identity auth is requested for + an endpoint that is not a recognized Azure Content Safety endpoint, or if no API key + is available for ``"api_key"`` auth. """ endpoint_value = default_values.get_required_value( env_var_name=self.ENDPOINT_URI_ENVIRONMENT_VARIABLE, passed_value=endpoint @@ -114,10 +121,18 @@ def __init__( self._api_version = api_version or "2024-09-01" - # Resolve authentication: an explicit key or token-provider callable, the - # env var, or — for a recognized Azure Content Safety endpoint with no key — - # an Entra ID token provider minted for the endpoint (identity-based auth). - if api_key is not None and callable(api_key): + # Resolve authentication: an explicit key or token-provider callable, or the env var. + # Identity is an explicit caller choice, so it must never be silently downgraded + # to a key that merely happens to be present in the environment. + if auth_mode == "identity": + if not is_azure_openai_endpoint(endpoint_value): + raise ValueError( + "Identity-based authentication requires a recognized Azure Content Safety endpoint " + f"(*.cognitiveservices.azure.com), but got '{endpoint_value}'. Use api_key authentication " + "for this endpoint, or pass your own token provider callable as api_key." + ) + self._api_key = get_azure_token_provider(get_default_azure_scope(endpoint_value)) + elif api_key is not None and callable(api_key): self._api_key = api_key else: api_key_value = default_values.get_non_required_value( @@ -125,17 +140,29 @@ def __init__( ) if api_key_value: self._api_key = api_key_value - elif is_azure_openai_endpoint(endpoint_value): - self._api_key = get_azure_token_provider(get_default_azure_scope(endpoint_value)) else: raise ValueError( - "API key is required for non-Azure Content Safety endpoints. For recognized Azure " - "endpoints (*.cognitiveservices.azure.com), identity-based authentication is used " - "automatically." + f"No API key available for endpoint '{endpoint_value}'. Set the " + f"{self.API_KEY_ENVIRONMENT_VARIABLE} environment variable, pass api_key (a key or a " + 'token provider callable), or pass auth_mode="identity" to authenticate with Microsoft ' + "Entra ID on a recognized Azure Content Safety endpoint (*.cognitiveservices.azure.com)." ) self._force_entry_field: PromptShieldEntryField = field + @classmethod + def get_auth_mode_parameters(cls, *, auth_mode: AuthMode) -> dict[str, object]: + """ + Preserve explicit authentication intent through target construction. + + Args: + auth_mode (AuthMode): Authentication mode selected by the caller. + + Returns: + dict[str, object]: Constructor parameters that enforce the mode. + """ + return {"auth_mode": auth_mode} + def _build_identifier(self) -> ComponentIdentifier: """ Build the identifier with Prompt Shield-specific parameters. diff --git a/tests/unit/auth/test_openai_auth.py b/tests/unit/auth/test_openai_auth.py new file mode 100644 index 0000000000..a411c68523 --- /dev/null +++ b/tests/unit/auth/test_openai_auth.py @@ -0,0 +1,137 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import os +from unittest.mock import patch + +import pytest + +from pyrit.auth.openai_auth import resolve_openai_auth + +AZURE_ENDPOINT = "https://my-resource.openai.azure.com/openai/v1" +NON_AZURE_ENDPOINT = "https://api.openai.com/v1" +API_KEY_ENV_VAR = "OPENAI_CHAT_API_KEY" + + +@pytest.fixture +def minted_provider(): + async def _provider() -> str: + return "entra-token" + + with patch("pyrit.auth.openai_auth.get_azure_openai_auth", return_value=_provider) as mock_auth: + yield _provider, mock_auth + + +def test_identity_ignores_env_var_api_key(minted_provider): + """An explicit identity choice must not be downgraded to a key sitting in the environment.""" + provider, mock_auth = minted_provider + with patch.dict(os.environ, {API_KEY_ENV_VAR: "sk-SECRET-FROM-DOTENV"}): + resolved = resolve_openai_auth( + endpoint=AZURE_ENDPOINT, + api_key=None, + api_key_environment_variable=API_KEY_ENV_VAR, + auth_mode="identity", + ) + + assert resolved is provider + mock_auth.assert_called_once_with(AZURE_ENDPOINT) + + +def test_identity_ignores_explicit_api_key(minted_provider): + """Identity wins over a key passed alongside it rather than silently using the key.""" + provider, _ = minted_provider + resolved = resolve_openai_auth( + endpoint=AZURE_ENDPOINT, + api_key="sk-explicit", + api_key_environment_variable=API_KEY_ENV_VAR, + auth_mode="identity", + ) + + assert resolved is provider + + +def test_identity_raises_for_non_azure_endpoint(): + with pytest.raises(ValueError, match="Identity-based authentication requires a recognized Azure"): + resolve_openai_auth( + endpoint=NON_AZURE_ENDPOINT, + api_key=None, + api_key_environment_variable=API_KEY_ENV_VAR, + auth_mode="identity", + ) + + +def test_default_auth_mode_uses_env_var(): + with patch.dict(os.environ, {API_KEY_ENV_VAR: "sk-from-env"}): + resolved = resolve_openai_auth( + endpoint=AZURE_ENDPOINT, + api_key=None, + api_key_environment_variable=API_KEY_ENV_VAR, + ) + + assert resolved == "sk-from-env" + + +def test_api_key_mode_prefers_explicit_key_over_env_var(): + with patch.dict(os.environ, {API_KEY_ENV_VAR: "sk-from-env"}): + resolved = resolve_openai_auth( + endpoint=AZURE_ENDPOINT, + api_key="sk-explicit", + api_key_environment_variable=API_KEY_ENV_VAR, + auth_mode="api_key", + ) + + assert resolved == "sk-explicit" + + +def test_api_key_mode_wraps_callable_before_reading_env_var(): + def sync_provider() -> str: + return "callable-token" + + with patch.dict(os.environ, {API_KEY_ENV_VAR: "sk-from-env"}): + resolved = resolve_openai_auth( + endpoint=AZURE_ENDPOINT, + api_key=sync_provider, + api_key_environment_variable=API_KEY_ENV_VAR, + ) + + assert callable(resolved) + assert resolved is not sync_provider + + +def test_api_key_mode_raises_when_no_key_available(minted_provider): + """api_key mode no longer mints an Entra token just because the endpoint looks like Azure.""" + _, mock_auth = minted_provider + with patch.dict(os.environ, {API_KEY_ENV_VAR: ""}): + with pytest.raises(ValueError, match="No API key available"): + resolve_openai_auth( + endpoint=AZURE_ENDPOINT, + api_key=None, + api_key_environment_variable=API_KEY_ENV_VAR, + ) + + mock_auth.assert_not_called() + + +def test_api_key_mode_raises_for_non_azure_endpoint_without_key(): + with patch.dict(os.environ, {API_KEY_ENV_VAR: ""}): + with pytest.raises(ValueError, match="No API key available"): + resolve_openai_auth( + endpoint=NON_AZURE_ENDPOINT, + api_key=None, + api_key_environment_variable=API_KEY_ENV_VAR, + ) + + +def test_api_key_mode_error_names_the_identity_migration(): + """The break is only safe if the error tells the caller how to opt into identity.""" + with patch.dict(os.environ, {API_KEY_ENV_VAR: ""}): + with pytest.raises(ValueError) as exc_info: + resolve_openai_auth( + endpoint=AZURE_ENDPOINT, + api_key=None, + api_key_environment_variable=API_KEY_ENV_VAR, + ) + + message = str(exc_info.value) + assert 'auth_mode="identity"' in message + assert API_KEY_ENV_VAR in message diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index e6803cc181..925885161c 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -727,9 +727,140 @@ async def test_create_openai_target_with_identity_non_azure_endpoint_raises(self auth_mode="identity", ) - with pytest.raises(ValueError, match="non-Azure endpoints"): + with pytest.raises(ValueError, match="Identity-based authentication requires a recognized Azure"): await service.create_target_async(request=request) + async def test_create_openai_target_with_identity_ignores_env_api_key(self, sqlite_instance) -> None: + """Regression: a key in the environment must not override an explicit identity choice.""" + + with patch.dict(os.environ, {"OPENAI_CHAT_KEY": "sk-SECRET-FROM-DOTENV"}): + with patch( + "pyrit.auth.openai_auth.get_azure_openai_auth", + return_value=_test_token_provider, + ): + service = TargetService() + + request = CreateTargetRequest( + type="OpenAIChatTarget", + params={ + "endpoint": "https://test.openai.azure.com/", + "model_name": "gpt-4o", + }, + auth_mode="identity", + ) + + result = await service.create_target_async(request=request) + + target_obj = service.get_target_object(target_registry_name=result.target_registry_name) + assert target_obj is not None + assert target_obj._api_key is _test_token_provider # type: ignore[attr-defined] + + async def test_create_azureml_target_with_identity_ignores_env_api_key(self, sqlite_instance) -> None: + """Regression: AZURE_ML_KEY must not override an explicit identity choice.""" + + with patch.dict(os.environ, {"AZURE_ML_KEY": "key-from-dotenv"}): + with patch( + "pyrit.prompt_target.azure_ml_chat_target.get_azure_async_token_provider", + return_value=_test_token_provider, + ): + service = TargetService() + + request = CreateTargetRequest( + type="AzureMLChatTarget", + params={"endpoint": "https://my-aml.region.inference.ml.azure.com/score"}, + auth_mode="identity", + ) + + result = await service.create_target_async(request=request) + + target_obj = service.get_target_object(target_registry_name=result.target_registry_name) + assert target_obj is not None + assert target_obj._api_key_provider is _test_token_provider # type: ignore[attr-defined] + assert target_obj._api_key == "" # type: ignore[attr-defined] + + async def test_create_target_with_api_key_mode_preserves_env_var_fallback(self, sqlite_instance) -> None: + """api_key requests keep the historical resolution order, including the env-var fallback.""" + with patch.dict(os.environ, {"OPENAI_CHAT_KEY": "sk-from-env"}): + service = TargetService() + + request = CreateTargetRequest( + type="OpenAIChatTarget", + params={ + "endpoint": "https://test.openai.azure.com/", + "model_name": "gpt-4o", + }, + ) + + result = await service.create_target_async(request=request) + + target_obj = service.get_target_object(target_registry_name=result.target_registry_name) + assert target_obj is not None + assert target_obj._api_key == "sk-from-env" # type: ignore[attr-defined] + + async def test_create_target_params_auth_mode_conflicting_with_api_key_request_raises( + self, sqlite_instance + ) -> None: + """params must not be a second channel that silently overrides a supplied key with identity.""" + service = TargetService() + + request = CreateTargetRequest( + type="OpenAIChatTarget", + params={ + "endpoint": "https://test.openai.azure.com/", + "model_name": "gpt-4o", + "api_key": "sk-user-supplied", + "auth_mode": "identity", + }, + auth_mode="api_key", + ) + + with pytest.raises(ValueError, match="Conflicting authentication modes"): + await service.create_target_async(request=request) + + async def test_create_target_params_auth_mode_conflicting_with_identity_request_raises( + self, sqlite_instance + ) -> None: + """The opposite conflict direction is rejected too, rather than silently resolved.""" + service = TargetService() + + request = CreateTargetRequest( + type="OpenAIChatTarget", + params={ + "endpoint": "https://test.openai.azure.com/", + "model_name": "gpt-4o", + "auth_mode": "api_key", + }, + auth_mode="identity", + ) + + with pytest.raises(ValueError, match="Conflicting authentication modes"): + await service.create_target_async(request=request) + + async def test_create_target_params_auth_mode_matching_request_is_accepted(self, sqlite_instance) -> None: + """A redundant but agreeing params auth_mode is harmless.""" + with patch.dict(os.environ, {"OPENAI_CHAT_KEY": "sk-SECRET-FROM-DOTENV"}): + with patch( + "pyrit.auth.openai_auth.get_azure_openai_auth", + return_value=_test_token_provider, + ): + service = TargetService() + + request = CreateTargetRequest( + type="OpenAIChatTarget", + params={ + "endpoint": "https://test.openai.azure.com/", + "model_name": "gpt-4o", + "auth_mode": "identity", + }, + auth_mode="identity", + ) + + result = await service.create_target_async(request=request) + + target_obj = service.get_target_object(target_registry_name=result.target_registry_name) + assert target_obj is not None + assert target_obj._api_key is _test_token_provider # type: ignore[attr-defined] + async def test_create_target_identity_unsupported_type_raises(self, sqlite_instance) -> None: """Identity-based auth is only supported for targets that declare it.""" service = TargetService() diff --git a/tests/unit/embedding/test_azure_text_embedding.py b/tests/unit/embedding/test_azure_text_embedding.py index d2f1ba4b63..e87ef00d69 100644 --- a/tests/unit/embedding/test_azure_text_embedding.py +++ b/tests/unit/embedding/test_azure_text_embedding.py @@ -29,7 +29,7 @@ def test_valid_init_env(): def test_invalid_key_raises(): """An empty API key on a non-Azure endpoint raises ValueError (no Entra fallback).""" os.environ[OpenAITextEmbedding.API_KEY_ENVIRONMENT_VARIABLE] = "" - with pytest.raises(ValueError, match="required for non-Azure endpoints"): + with pytest.raises(ValueError, match="No API key available"): OpenAITextEmbedding( api_key="", endpoint="https://api.openai.com/v1", @@ -111,10 +111,16 @@ def _build_embedding( endpoint: str = _AZURE_ENDPOINT, api_key: str | Callable[[], str | Awaitable[str]] | None = "test-key", model_name: str = "text-embedding-3-small", + auth_mode: str = "api_key", ) -> OpenAITextEmbedding: """Build an OpenAITextEmbedding with a cleared environment so env vars don't leak in.""" with patch.dict(os.environ, {}, clear=True): - return OpenAITextEmbedding(api_key=api_key, endpoint=endpoint, model_name=model_name) + return OpenAITextEmbedding( + api_key=api_key, + endpoint=endpoint, + model_name=model_name, + auth_mode=auth_mode, # type: ignore[arg-type] + ) @patch("pyrit.embedding.openai_text_embedding.AsyncOpenAI") @@ -141,19 +147,52 @@ async def async_provider() -> str: @patch("pyrit.embedding.openai_text_embedding.AsyncOpenAI") -def test_no_key_azure_endpoint_falls_back_to_entra(mock_async_openai): - """A recognized Azure endpoint with no key mints an Entra token provider.""" +def test_no_key_azure_endpoint_raises(mock_async_openai): + """A recognized Azure endpoint no longer auto-mints a token; identity must be explicit.""" + mock_async_openai.return_value = MagicMock() + + with patch("pyrit.auth.openai_auth.get_azure_openai_auth") as mock_get_auth: + with pytest.raises(ValueError, match="No API key available"): + _build_embedding(api_key=None, endpoint=_AZURE_ENDPOINT) + + mock_get_auth.assert_not_called() + + +@patch("pyrit.embedding.openai_text_embedding.AsyncOpenAI") +def test_identity_auth_mode_ignores_env_key(mock_async_openai): + """An explicit identity choice must not be downgraded to the embedding key env var.""" mock_async_openai.return_value = MagicMock() mock_auth = AsyncMock(return_value="entra-token") - with patch("pyrit.auth.openai_auth.get_azure_openai_auth", return_value=mock_auth) as mock_get_auth: - _build_embedding(api_key=None, endpoint=_AZURE_ENDPOINT) + with ( + patch.dict( + os.environ, + {OpenAITextEmbedding.API_KEY_ENVIRONMENT_VARIABLE: "sk-SECRET-FROM-DOTENV"}, + clear=True, + ), + patch("pyrit.auth.openai_auth.get_azure_openai_auth", return_value=mock_auth) as mock_get_auth, + ): + OpenAITextEmbedding( + api_key=None, + endpoint=_AZURE_ENDPOINT, + model_name="text-embedding-3-small", + auth_mode="identity", + ) mock_get_auth.assert_called_once_with(_AZURE_ENDPOINT) assert mock_async_openai.call_args.kwargs["api_key"] is mock_auth +@patch("pyrit.embedding.openai_text_embedding.AsyncOpenAI") +def test_identity_auth_mode_non_azure_endpoint_raises(mock_async_openai): + """Identity must never mint a token for an unrecognized host.""" + mock_async_openai.return_value = MagicMock() + + with pytest.raises(ValueError, match="Identity-based authentication requires"): + _build_embedding(api_key=None, endpoint=_NON_AZURE_ENDPOINT, auth_mode="identity") + + def test_no_key_non_azure_endpoint_raises(): """A non-Azure endpoint with no key raises ValueError (no Entra fallback).""" - with pytest.raises(ValueError, match="required for non-Azure endpoints"): + with pytest.raises(ValueError, match="No API key available"): _build_embedding(api_key=None, endpoint=_NON_AZURE_ENDPOINT) diff --git a/tests/unit/prompt_target/target/test_azure_ml_chat_target.py b/tests/unit/prompt_target/target/test_azure_ml_chat_target.py index 6c39a42b4a..2fad20fd36 100644 --- a/tests/unit/prompt_target/target/test_azure_ml_chat_target.py +++ b/tests/unit/prompt_target/target/test_azure_ml_chat_target.py @@ -55,27 +55,65 @@ def test_initialization_with_no_api_raises(): AzureMLChatTarget(api_key="xxxxx") -def test_no_key_recognized_aml_endpoint_auto_mints_entra(patch_central_database): - """With no key and a recognized *.inference.ml.azure.com endpoint, the target - auto-mints an Entra token provider for the AML scope.""" +def test_no_key_recognized_aml_endpoint_raises(patch_central_database): + """A recognized AML endpoint no longer auto-mints a token; identity must be explicit.""" + with ( + patch.dict(os.environ, {AzureMLChatTarget.api_key_environment_variable: ""}), + patch( + "pyrit.prompt_target.azure_ml_chat_target.get_azure_async_token_provider", + ) as mock_provider, + ): + with pytest.raises(ValueError, match="No API key available"): + AzureMLChatTarget(endpoint="https://my-aml.region.inference.ml.azure.com/score") + + mock_provider.assert_not_called() + + +def test_identity_auth_mode_ignores_env_key(patch_central_database): + """An explicit identity choice must not be downgraded to the AZURE_ML_KEY env var.""" async def _provider() -> str: return "aml-entra-token" with ( - patch.dict(os.environ, {AzureMLChatTarget.api_key_environment_variable: ""}), + patch.dict(os.environ, {AzureMLChatTarget.api_key_environment_variable: "key-from-dotenv"}), patch( "pyrit.prompt_target.azure_ml_chat_target.get_azure_async_token_provider", return_value=_provider, - ) as mock_provider, + ), + ): + target = AzureMLChatTarget( + endpoint="https://my-aml.region.inference.ml.azure.com/score", + auth_mode="identity", + ) + + assert target._api_key_provider is _provider + assert target._api_key == "" + + +def test_identity_auth_mode_ignores_explicit_key(patch_central_database): + async def _provider() -> str: + return "aml-entra-token" + + with patch( + "pyrit.prompt_target.azure_ml_chat_target.get_azure_async_token_provider", + return_value=_provider, ): - target = AzureMLChatTarget(endpoint="https://my-aml.region.inference.ml.azure.com/score") + target = AzureMLChatTarget( + endpoint="https://my-aml.region.inference.ml.azure.com/score", + api_key="key-passed-anyway", + auth_mode="identity", + ) - mock_provider.assert_called_once_with(AzureMLChatTarget._AZURE_ML_SCOPE) assert target._api_key_provider is _provider assert target._api_key == "" +def test_identity_auth_mode_non_aml_endpoint_raises(patch_central_database): + with pytest.raises(ValueError, match="Identity-based authentication requires a recognized Azure ML"): + AzureMLChatTarget(endpoint="http://aml-test-endpoint.com", auth_mode="identity") + + def test_no_key_non_aml_endpoint_raises(patch_central_database): """With no key and an endpoint that is not a recognized AML host, the target refuses to mint a bearer token.""" diff --git a/tests/unit/prompt_target/target/test_openai_chat_target.py b/tests/unit/prompt_target/target/test_openai_chat_target.py index 589b645268..429326c207 100644 --- a/tests/unit/prompt_target/target/test_openai_chat_target.py +++ b/tests/unit/prompt_target/target/test_openai_chat_target.py @@ -847,33 +847,25 @@ def test_set_auth_with_api_key(patch_central_database): assert target._api_key == "test_api_key_456" -def test_no_key_recognized_azure_endpoint_auto_mints_entra(patch_central_database): - """With no key and a recognized Azure OpenAI endpoint, the target auto-mints - an Entra token provider for that endpoint.""" - - async def _provider() -> str: - return "aoai-entra-token" - +def test_no_key_recognized_azure_endpoint_raises(patch_central_database): + """A recognized Azure OpenAI endpoint no longer auto-mints a token; identity must be explicit.""" with ( patch.dict(os.environ, {}, clear=True), - patch( - "pyrit.auth.openai_auth.get_azure_openai_auth", - return_value=_provider, - ) as mock_get_auth, + patch("pyrit.auth.openai_auth.get_azure_openai_auth") as mock_get_auth, ): - target = OpenAIChatTarget( - model_name="gpt-4", - endpoint="https://test.openai.azure.com/", - ) + with pytest.raises(ValueError, match="No API key available"): + OpenAIChatTarget( + model_name="gpt-4", + endpoint="https://test.openai.azure.com/", + ) - mock_get_auth.assert_called_once_with("https://test.openai.azure.com/") - assert target._api_key is _provider + mock_get_auth.assert_not_called() def test_no_key_non_azure_endpoint_raises(patch_central_database): """With no key and a non-Azure endpoint, the target refuses to mint a token.""" with patch.dict(os.environ, {}, clear=True): - with pytest.raises(ValueError, match="non-Azure endpoints"): + with pytest.raises(ValueError, match="No API key available"): OpenAIChatTarget(model_name="gpt-4", endpoint="https://api.openai.com/") @@ -881,7 +873,7 @@ def test_no_key_substring_lookalike_endpoint_raises(patch_central_database): """A hostname merely containing 'azure' (but not a recognized suffix) must not trigger auto-Entra minting (loose->strict hardening).""" with patch.dict(os.environ, {}, clear=True): - with pytest.raises(ValueError, match="non-Azure endpoints"): + with pytest.raises(ValueError, match="No API key available"): OpenAIChatTarget(model_name="gpt-4", endpoint="https://evil-azure.example.com/") diff --git a/tests/unit/prompt_target/target/test_openai_target_auth.py b/tests/unit/prompt_target/target/test_openai_target_auth.py index b2b98e0d73..6ab4b21b73 100644 --- a/tests/unit/prompt_target/target/test_openai_target_auth.py +++ b/tests/unit/prompt_target/target/test_openai_target_auth.py @@ -10,6 +10,7 @@ import pytest from pyrit.auth import ensure_async_token_provider +from pyrit.common.auth_mode import AuthMode from pyrit.prompt_target.openai.openai_target import OpenAITarget @@ -42,6 +43,7 @@ def _build_target( endpoint: str = "https://test.openai.azure.com/openai/v1", api_key: str | Callable | None = "test-key", env_vars: dict[str, str] | None = None, + auth_mode: AuthMode = "api_key", ) -> _ConcreteOpenAITarget: """Helper to build a _ConcreteOpenAITarget with controlled env.""" env = {"TEST_MODEL": "gpt-4", "TEST_ENDPOINT": endpoint} @@ -52,6 +54,7 @@ def _build_target( model_name="gpt-4", endpoint=endpoint, api_key=api_key, + auth_mode=auth_mode, ) @@ -71,22 +74,22 @@ def test_env_var_api_key_used_when_no_param(self): def test_non_azure_endpoint_without_key_raises(self): """Non-Azure endpoints must have an API key; otherwise ValueError is raised.""" - with pytest.raises(ValueError, match="TEST_API_KEY is required for non-Azure endpoints"): + with pytest.raises(ValueError, match="No API key available"): _build_target( endpoint="https://api.openai.com/v1", api_key=None, ) - def test_azure_endpoint_falls_back_to_entra(self): - """Azure endpoints without a key fall back to get_azure_openai_auth.""" - mock_auth = AsyncMock(return_value="entra-token") - with patch("pyrit.auth.openai_auth.get_azure_openai_auth", return_value=mock_auth): - target = _build_target( - endpoint="https://myresource.openai.azure.com/openai/v1", - api_key=None, - ) - # The api_key should be the async callable returned by get_azure_openai_auth - assert target._api_key is mock_auth + def test_azure_endpoint_without_key_raises(self): + """Azure endpoints no longer fall back to Entra implicitly; identity must be explicit.""" + with patch("pyrit.auth.openai_auth.get_azure_openai_auth") as mock_auth: + with pytest.raises(ValueError, match="No API key available"): + _build_target( + endpoint="https://myresource.openai.azure.com/openai/v1", + api_key=None, + ) + + mock_auth.assert_not_called() def test_callable_token_provider_bypasses_env_lookup(self): """A callable api_key is used directly without checking env vars.""" @@ -121,6 +124,31 @@ def test_param_api_key_takes_precedence_over_env_var(self): target = _build_target(api_key="param-key", env_vars={"TEST_API_KEY": "env-key"}) assert target._api_key == "param-key" + def test_identity_auth_mode_ignores_env_var_key(self): + """An explicit identity choice must not be downgraded to the key in the environment.""" + mock_auth = AsyncMock(return_value="entra-token") + with patch("pyrit.auth.openai_auth.get_azure_openai_auth", return_value=mock_auth): + target = _build_target( + api_key=None, + env_vars={"TEST_API_KEY": "env-key"}, + auth_mode="identity", + ) + assert target._api_key is mock_auth + + def test_identity_auth_mode_ignores_explicit_key(self): + mock_auth = AsyncMock(return_value="entra-token") + with patch("pyrit.auth.openai_auth.get_azure_openai_auth", return_value=mock_auth): + target = _build_target(api_key="param-key", auth_mode="identity") + assert target._api_key is mock_auth + + def test_identity_auth_mode_non_azure_endpoint_raises(self): + with pytest.raises(ValueError, match="Identity-based authentication requires a recognized Azure"): + _build_target( + endpoint="https://api.openai.com/v1", + api_key=None, + auth_mode="identity", + ) + class TestEnsureAsyncTokenProvider: """Tests for the ensure_async_token_provider helper function.""" diff --git a/tests/unit/prompt_target/target/test_prompt_shield_target.py b/tests/unit/prompt_target/target/test_prompt_shield_target.py index fc885ee44b..8f45e28535 100644 --- a/tests/unit/prompt_target/target/test_prompt_shield_target.py +++ b/tests/unit/prompt_target/target/test_prompt_shield_target.py @@ -163,25 +163,59 @@ def test_init_raises_when_no_api_key_and_non_azure_endpoint(sqlite_instance): """No key + a non-Azure endpoint raises (identity auth only works for Azure endpoints).""" with patch.dict(os.environ, {}, clear=False): os.environ.pop("AZURE_CONTENT_SAFETY_API_KEY", None) - with pytest.raises(ValueError, match="API key is required for non-Azure"): + with pytest.raises(ValueError, match="No API key available"): PromptShieldTarget(endpoint="https://test.endpoint.com", api_key=None) -def test_init_uses_identity_token_provider_for_azure_endpoint(sqlite_instance): - """No key + a recognized Azure Content Safety endpoint falls back to an Entra ID token provider.""" - token_provider = MagicMock(return_value="minted-token") +def test_init_raises_when_no_api_key_on_azure_endpoint(sqlite_instance): + """A recognized Azure endpoint no longer auto-mints a token; identity must be explicit.""" with patch.dict(os.environ, {}, clear=False): os.environ.pop("AZURE_CONTENT_SAFETY_API_KEY", None) with patch( "pyrit.prompt_target.prompt_shield_target.get_azure_token_provider", - return_value=token_provider, ) as mock_provider: - target = PromptShieldTarget(endpoint="https://myresource.cognitiveservices.azure.com", api_key=None) + with pytest.raises(ValueError, match="No API key available"): + PromptShieldTarget(endpoint="https://myresource.cognitiveservices.azure.com", api_key=None) - mock_provider.assert_called_once_with("https://cognitiveservices.azure.com/.default") - assert target._api_key is token_provider + mock_provider.assert_not_called() def test_supported_auth_modes_includes_identity(): """Prompt Shield advertises identity-based auth alongside api_key.""" assert PromptShieldTarget.supported_auth_modes == ("api_key", "identity") + + +def test_identity_auth_mode_ignores_env_key(sqlite_instance): + """An explicit identity choice must not be downgraded to the content safety key env var.""" + token_provider = MagicMock(return_value="minted-token") + with patch.dict(os.environ, {"AZURE_CONTENT_SAFETY_API_KEY": "key-from-dotenv"}): + with patch( + "pyrit.prompt_target.prompt_shield_target.get_azure_token_provider", + return_value=token_provider, + ): + target = PromptShieldTarget( + endpoint="https://myresource.cognitiveservices.azure.com", + auth_mode="identity", + ) + + assert target._api_key is token_provider + + +def test_identity_auth_mode_ignores_explicit_key(sqlite_instance): + token_provider = MagicMock(return_value="minted-token") + with patch( + "pyrit.prompt_target.prompt_shield_target.get_azure_token_provider", + return_value=token_provider, + ): + target = PromptShieldTarget( + endpoint="https://myresource.cognitiveservices.azure.com", + api_key="key-passed-anyway", + auth_mode="identity", + ) + + assert target._api_key is token_provider + + +def test_identity_auth_mode_non_azure_endpoint_raises(sqlite_instance): + with pytest.raises(ValueError, match="Identity-based authentication requires a recognized Azure Content Safety"): + PromptShieldTarget(endpoint="https://test.endpoint.com", auth_mode="identity") diff --git a/tests/unit/registry/test_target_registry.py b/tests/unit/registry/test_target_registry.py index 00b8260500..ffc0fbe1b7 100644 --- a/tests/unit/registry/test_target_registry.py +++ b/tests/unit/registry/test_target_registry.py @@ -446,3 +446,15 @@ def test_credential_parameters_are_sensitive(self, registry: TargetRegistry) -> ) if looks_like_credential: assert parameter.sensitive, f"{name}.{parameter.name} looks like a credential" + + def test_identity_targets_accept_an_explicit_auth_mode(self, registry: TargetRegistry) -> None: + # Advertising identity support without accepting the explicit mode would silently + # fall back to inferring auth from a missing key, which is the ambiguity the + # explicit auth_mode contract exists to remove. + for name in registry.get_class_names(): + target_cls = registry.get_class(name) + if "identity" not in target_cls.supported_auth_modes: + continue + assert target_cls.get_auth_mode_parameters(auth_mode="identity") == {"auth_mode": "identity"}, ( + f"{name} advertises identity support but does not forward an explicit auth_mode" + )