From 0abe8fe2b17badd1383e5faf6fa429c9e1d93397 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Tue, 6 Oct 2026 15:47:59 -0400 Subject: [PATCH 1/6] FIX: Make identity auth an explicit choice instead of an inferred one Selecting identity-based authentication was signalled by deleting the api_key, but "no api_key" is ambiguous: every auth resolver interprets it as "read the key from the environment variable", so an explicit identity choice was silently downgraded to api-key auth whenever the env var was set. Thread an explicit auth_mode through resolve_openai_auth, OpenAITarget, AzureMLChatTarget and PromptShieldTarget. auth_mode defaults to "api_key", so the existing callable -> explicit key -> env var -> Entra fallback chain is unchanged; only an explicit auth_mode="identity" short-circuits to a token provider. Identity still refuses to mint tokens for unrecognized hosts and now raises a clear ValueError instead. AuthMode moves to pyrit/common/auth_mode.py so pyrit.auth can reference it without depending on the target layer; it is re-exported from pyrit.prompt_target.common.prompt_target for backward compatibility. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/auth/openai_auth.py | 21 ++- pyrit/backend/services/target_service.py | 27 ++-- pyrit/common/__init__.py | 3 + pyrit/common/auth_mode.py | 29 +++++ pyrit/prompt_target/azure_ml_chat_target.py | 59 ++++++++- pyrit/prompt_target/common/prompt_target.py | 12 +- pyrit/prompt_target/openai/openai_target.py | 26 +++- pyrit/prompt_target/prompt_shield_target.py | 36 +++++- tests/unit/auth/test_openai_auth.py | 120 ++++++++++++++++++ tests/unit/backend/test_target_service.py | 69 +++++++++- .../target/test_azure_ml_chat_target.py | 45 +++++++ .../target/test_openai_target_auth.py | 28 ++++ .../target/test_prompt_shield_target.py | 36 ++++++ 13 files changed, 481 insertions(+), 30 deletions(-) create mode 100644 pyrit/common/auth_mode.py create mode 100644 tests/unit/auth/test_openai_auth.py diff --git a/pyrit/auth/openai_auth.py b/pyrit/auth/openai_auth.py index aa281cc199..17d98a82f7 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,6 +14,7 @@ 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. @@ -21,13 +23,30 @@ def resolve_openai_auth( 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) keeps the + historical resolution order: token-provider callable, explicit key, environment variable, + then an Entra ID fallback for recognized Azure endpoints. 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 no key is provided and the endpoint is not a recognized Azure + OpenAI endpoint. """ + # 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)) diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index 8d1176ed07..6ad9b99042 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, cast 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(cast("AuthMode", auth_mode)) return supported_auth_modes def _project_target_parameters(self, *, target_type: str, parameters: tuple[Parameter, ...]) -> list[Parameter]: @@ -223,7 +226,10 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn 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 + endpoint and authenticates itself. The selected mode is also passed + explicitly to targets that accept it, via + ``get_auth_mode_parameters``, 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: @@ -249,7 +255,8 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn 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 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/prompt_target/azure_ml_chat_target.py b/pyrit/prompt_target/azure_ml_chat_target.py index ee0d83ff93..c9f1d91f94 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) keeps the historical resolution order. Raises: - ValueError: If no api_key is supplied (via parameter or environment - variable) and the endpoint is not a recognized Azure ML managed + ValueError: If identity auth is requested for an endpoint that is not a recognized + Azure ML managed online endpoint, or 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. """ 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 @@ -204,8 +242,7 @@ def _initialize_vars( # 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_provider = self._build_azure_ml_token_provider() self._api_key = "" return @@ -215,6 +252,16 @@ def _initialize_vars( "authentication is used automatically. Pass an api_key or a token provider callable instead." ) + 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..2dfb28720d 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, @@ -108,6 +109,11 @@ def __init__( (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 +128,10 @@ 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 no API key is provided (via parameter or + environment variable) and the endpoint is not a recognized Azure OpenAI / + AI Foundry endpoint. """ self._headers: dict[str, str] = {} self._httpx_client_kwargs = httpx_client_kwargs or {} @@ -157,10 +165,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..b1c9dcb53e 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 provided for a non-Azure Content Safety endpoint. """ endpoint_value = default_values.get_required_value( env_var_name=self.ENDPOINT_URI_ENVIRONMENT_VARIABLE, passed_value=endpoint @@ -117,7 +124,17 @@ def __init__( # 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): + # 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( @@ -136,6 +153,19 @@ def __init__( 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..7eaa9cbee7 --- /dev/null +++ b/tests/unit/auth/test_openai_auth.py @@ -0,0 +1,120 @@ +# 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_falls_back_to_entra_when_no_key(minted_provider): + provider, _ = minted_provider + with patch.dict(os.environ, {API_KEY_ENV_VAR: ""}): + resolved = resolve_openai_auth( + endpoint=AZURE_ENDPOINT, + api_key=None, + api_key_environment_variable=API_KEY_ENV_VAR, + ) + + assert resolved is provider + + +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="is required for non-Azure endpoints"): + resolve_openai_auth( + endpoint=NON_AZURE_ENDPOINT, + api_key=None, + api_key_environment_variable=API_KEY_ENV_VAR, + ) diff --git a/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index e6803cc181..f6bc0a6340 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -727,9 +727,76 @@ 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_does_not_pass_auth_mode(self, sqlite_instance) -> None: + """api_key requests keep the historical params exactly, 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_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/prompt_target/target/test_azure_ml_chat_target.py b/tests/unit/prompt_target/target/test_azure_ml_chat_target.py index 6c39a42b4a..dce6876ef7 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 @@ -76,6 +76,51 @@ async def _provider() -> str: assert target._api_key == "" +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: "key-from-dotenv"}), + 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", + 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", + api_key="key-passed-anyway", + auth_mode="identity", + ) + + 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_target_auth.py b/tests/unit/prompt_target/target/test_openai_target_auth.py index b2b98e0d73..2ac38424e1 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, ) @@ -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..94952ff8cc 100644 --- a/tests/unit/prompt_target/target/test_prompt_shield_target.py +++ b/tests/unit/prompt_target/target/test_prompt_shield_target.py @@ -185,3 +185,39 @@ def test_init_uses_identity_token_provider_for_azure_endpoint(sqlite_instance): 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") From 9e3c184aa5c970f3f0ea270467836e64f43527b3 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Wed, 7 Oct 2026 11:54:38 -0400 Subject: [PATCH 2/6] FIX: Make request-level auth_mode authoritative over params auth_mode is also a constructor parameter, so the registry accepted it inside params as a second, competing channel. A request with auth_mode="api_key" and params["auth_mode"]="identity" selected identity, silently ignoring a supplied api_key and bypassing the service's supported_auth_modes check; the opposite conflict was silently resolved in favor of the top-level value. Reject conflicting values and forward the request-level auth_mode for both modes so it is authoritative in either direction. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/backend/services/target_service.py | 31 +++++++---- tests/unit/backend/test_target_service.py | 68 ++++++++++++++++++++++- 2 files changed, 87 insertions(+), 12 deletions(-) diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index 6ad9b99042..2d419113a2 100644 --- a/pyrit/backend/services/target_service.py +++ b/pyrit/backend/services/target_service.py @@ -223,14 +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 selected mode is also passed - explicitly to targets that accept it, via - ``get_auth_mode_parameters``, 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. + 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. @@ -239,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. """ @@ -252,6 +253,15 @@ 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.") @@ -266,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/tests/unit/backend/test_target_service.py b/tests/unit/backend/test_target_service.py index f6bc0a6340..925885161c 100644 --- a/tests/unit/backend/test_target_service.py +++ b/tests/unit/backend/test_target_service.py @@ -778,8 +778,8 @@ async def test_create_azureml_target_with_identity_ignores_env_api_key(self, sql 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_does_not_pass_auth_mode(self, sqlite_instance) -> None: - """api_key requests keep the historical params exactly, including the env-var fallback.""" + 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() @@ -797,6 +797,70 @@ async def test_create_target_with_api_key_mode_does_not_pass_auth_mode(self, sql 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() From 0c780de5d7673b8118959dc1edc1bd5450a7eec2 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Wed, 7 Oct 2026 11:56:14 -0400 Subject: [PATCH 3/6] DOC: Document explicit auth_mode="identity" for Entra auth The configuration guide described Entra auth only as the implicit fallback. Document the explicit mode, the resolution order it bypasses, and the targets that accept it. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/setup/1_configuration.ipynb | 17 ++++++++++++++++- doc/code/setup/1_configuration.py | 15 +++++++++++++++ 2 files changed, 31 insertions(+), 1 deletion(-) diff --git a/doc/code/setup/1_configuration.ipynb b/doc/code/setup/1_configuration.ipynb index fc18179efd..3bf38f01af 100644 --- a/doc/code/setup/1_configuration.ipynb +++ b/doc/code/setup/1_configuration.ipynb @@ -197,7 +197,22 @@ "\n", " ```bash\n", " az login\n", - " ```" + " ```\n", + "\n", + "### Choosing Entra auth explicitly\n", + "\n", + "By default a target resolves its credential in this order: a token provider callable passed as `api_key`, an explicit `api_key` string, the target's API key environment variable, and finally — for recognized Azure endpoints only — an Entra token. That last step is a *fallback*, so it is skipped whenever a key happens to be set in your `.env`.\n", + "\n", + "Pass `auth_mode=\"identity\"` when you want Entra auth regardless of what is in the environment. It skips the key and the environment variable entirely, and 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` defaults to `\"api_key\"`, which preserves the resolution order above. It is supported by the OpenAI targets, `AzureMLChatTarget`, and `PromptShieldTarget`." ] }, { diff --git a/doc/code/setup/1_configuration.py b/doc/code/setup/1_configuration.py index 2504b9d0d7..3f061baa76 100644 --- a/doc/code/setup/1_configuration.py +++ b/doc/code/setup/1_configuration.py @@ -108,6 +108,21 @@ # ```bash # az login # ``` +# +# ### Choosing Entra auth explicitly +# +# By default a target resolves its credential in this order: a token provider callable passed as `api_key`, an explicit `api_key` string, the target's API key environment variable, and finally — for recognized Azure endpoints only — an Entra token. That last step is a *fallback*, so it is skipped whenever a key happens to be set in your `.env`. +# +# Pass `auth_mode="identity"` when you want Entra auth regardless of what is in the environment. It skips the key and the environment variable entirely, and 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` defaults to `"api_key"`, which preserves the resolution order above. It is supported by the OpenAI targets, `AzureMLChatTarget`, and `PromptShieldTarget`. # %% [markdown] # ## Choosing a database From 9e51bbdccecdd8ea0eff07eabe61831670f4b74c Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Wed, 7 Oct 2026 16:56:03 -0400 Subject: [PATCH 4/6] FIX: Reconcile explicit auth_mode with the get_auth_mode_parameters hook PR #2846 landed the AzureBlobStorageTarget half of this bug using a get_auth_mode_parameters classmethod. Adopt that hook as the single mechanism for carrying auth intent into target construction and drop the _accepts_auth_mode signature introspection, which only existed because AzureBlobStorageTarget lacked the parameter. Every target advertising identity support now overrides the hook, guarded by a registry-wide contract test so a future identity target cannot silently fall back to inferring auth from a missing key. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/setup/1_configuration.ipynb | 2 +- doc/code/setup/1_configuration.py | 2 +- tests/unit/registry/test_target_registry.py | 12 ++++++++++++ 3 files changed, 14 insertions(+), 2 deletions(-) diff --git a/doc/code/setup/1_configuration.ipynb b/doc/code/setup/1_configuration.ipynb index 3bf38f01af..d7597e993d 100644 --- a/doc/code/setup/1_configuration.ipynb +++ b/doc/code/setup/1_configuration.ipynb @@ -212,7 +212,7 @@ ")\n", "```\n", "\n", - "`auth_mode` defaults to `\"api_key\"`, which preserves the resolution order above. It is supported by the OpenAI targets, `AzureMLChatTarget`, and `PromptShieldTarget`." + "On the OpenAI targets, `AzureMLChatTarget`, and `PromptShieldTarget`, `auth_mode` defaults to `\"api_key\"`, which preserves the resolution order above. `AzureBlobStorageTarget` accepts the same `auth_mode=\"identity\"` to bypass its SAS token sources; it defaults to selecting a credential automatically." ] }, { diff --git a/doc/code/setup/1_configuration.py b/doc/code/setup/1_configuration.py index 3f061baa76..f965514650 100644 --- a/doc/code/setup/1_configuration.py +++ b/doc/code/setup/1_configuration.py @@ -122,7 +122,7 @@ # ) # ``` # -# `auth_mode` defaults to `"api_key"`, which preserves the resolution order above. It is supported by the OpenAI targets, `AzureMLChatTarget`, and `PromptShieldTarget`. +# On the OpenAI targets, `AzureMLChatTarget`, and `PromptShieldTarget`, `auth_mode` defaults to `"api_key"`, which preserves the resolution order above. `AzureBlobStorageTarget` accepts the same `auth_mode="identity"` to bypass its SAS token sources; it defaults to selecting a credential automatically. # %% [markdown] # ## Choosing a database 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" + ) From 19f2f8031e15c949039f12f1ab8842948347eec9 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Wed, 7 Oct 2026 17:56:26 -0400 Subject: [PATCH 5/6] FIX: Drop redundant AuthMode cast flagged by ty AUTH_MODES is typed tuple[AuthMode, ...], so the membership check already narrows auth_mode to AuthMode and the cast tripped ty's redundant-cast rule. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- pyrit/backend/services/target_service.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index 2d419113a2..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, cast +from typing import Any from pyrit.backend.mappers.target_mappers import target_object_to_instance from pyrit.backend.models.common import PaginationInfo @@ -159,7 +159,7 @@ def _get_supported_auth_modes(auth_modes: tuple[str, ...]) -> list[AuthMode]: for auth_mode in auth_modes: if auth_mode not in AUTH_MODES: raise ValueError(f"Unsupported target authentication mode: {auth_mode!r}") - supported_auth_modes.append(cast("AuthMode", auth_mode)) + supported_auth_modes.append(auth_mode) return supported_auth_modes def _project_target_parameters(self, *, target_type: str, parameters: tuple[Parameter, ...]) -> list[Parameter]: From da72cf550f8afa0f44ca5d8ea698242cdaa7a732 Mon Sep 17 00:00:00 2001 From: Copilot <223556219+Copilot@users.noreply.github.com> Date: Thu, 8 Oct 2026 15:28:30 -0400 Subject: [PATCH 6/6] FIX: remove the implicit Entra fallback so auth intent is always explicit Per review feedback, drop the "no key found and the endpoint looks like Azure, so mint a token" fallback from the OpenAI resolver and the two inlined copies in AzureMLChatTarget and PromptShieldTarget. That fallback is what made an explicit auth_mode="identity" indistinguishable from "no key configured", so the two paths could not be told apart. api_key mode now requires a key or an explicit token provider and fails with an error that names the env var, api_key, and auth_mode="identity". identity mode ignores keys entirely. Also threads auth_mode through OpenAITextEmbedding, which is a fourth consumer of resolve_openai_auth and would otherwise have lost its only ergonomic path to identity auth. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- doc/code/setup/1_configuration.ipynb | 10 ++-- doc/code/setup/1_configuration.py | 10 ++-- pyrit/auth/openai_auth.py | 18 +++---- pyrit/embedding/openai_text_embedding.py | 16 ++++-- pyrit/prompt_target/azure_ml_chat_target.py | 23 +++----- pyrit/prompt_target/openai/openai_target.py | 9 ++-- pyrit/prompt_target/prompt_shield_target.py | 15 +++--- tests/unit/auth/test_openai_auth.py | 35 ++++++++---- .../embedding/test_azure_text_embedding.py | 53 ++++++++++++++++--- .../target/test_azure_ml_chat_target.py | 17 ++---- .../target/test_openai_chat_target.py | 30 ++++------- .../target/test_openai_target_auth.py | 22 ++++---- .../target/test_prompt_shield_target.py | 14 +++-- 13 files changed, 152 insertions(+), 120 deletions(-) diff --git a/doc/code/setup/1_configuration.ipynb b/doc/code/setup/1_configuration.ipynb index d7597e993d..1c2b8fc018 100644 --- a/doc/code/setup/1_configuration.ipynb +++ b/doc/code/setup/1_configuration.ipynb @@ -199,11 +199,11 @@ " az login\n", " ```\n", "\n", - "### Choosing Entra auth explicitly\n", + "### Choosing between key and Entra auth\n", "\n", - "By default a target resolves its credential in this order: a token provider callable passed as `api_key`, an explicit `api_key` string, the target's API key environment variable, and finally — for recognized Azure endpoints only — an Entra token. That last step is a *fallback*, so it is skipped whenever a key happens to be set in your `.env`.\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 regardless of what is in the environment. It skips the key and the environment variable entirely, and raises a `ValueError` if the endpoint is not a recognized Azure endpoint rather than minting a token for an unknown host.\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", @@ -212,7 +212,9 @@ ")\n", "```\n", "\n", - "On the OpenAI targets, `AzureMLChatTarget`, and `PromptShieldTarget`, `auth_mode` defaults to `\"api_key\"`, which preserves the resolution order above. `AzureBlobStorageTarget` accepts the same `auth_mode=\"identity\"` to bypass its SAS token sources; it defaults to selecting a credential automatically." + "`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 f965514650..44954cb9b6 100644 --- a/doc/code/setup/1_configuration.py +++ b/doc/code/setup/1_configuration.py @@ -109,11 +109,11 @@ # az login # ``` # -# ### Choosing Entra auth explicitly +# ### Choosing between key and Entra auth # -# By default a target resolves its credential in this order: a token provider callable passed as `api_key`, an explicit `api_key` string, the target's API key environment variable, and finally — for recognized Azure endpoints only — an Entra token. That last step is a *fallback*, so it is skipped whenever a key happens to be set in your `.env`. +# 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 regardless of what is in the environment. It skips the key and the environment variable entirely, and raises a `ValueError` if the endpoint is not a recognized Azure endpoint rather than minting a token for an unknown host. +# 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( @@ -122,7 +122,9 @@ # ) # ``` # -# On the OpenAI targets, `AzureMLChatTarget`, and `PromptShieldTarget`, `auth_mode` defaults to `"api_key"`, which preserves the resolution order above. `AzureBlobStorageTarget` accepts the same `auth_mode="identity"` to bypass its SAS token sources; it defaults to selecting a credential automatically. +# `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 17d98a82f7..927d31b1da 100644 --- a/pyrit/auth/openai_auth.py +++ b/pyrit/auth/openai_auth.py @@ -17,24 +17,22 @@ def resolve_openai_auth( 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) keeps the - historical resolution order: token-provider callable, explicit key, environment variable, - then an Entra ID fallback for recognized Azure endpoints. + ``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 identity auth is requested for an endpoint that is not a recognized Azure - OpenAI endpoint, or if no key is provided and the endpoint is not a recognized Azure - OpenAI endpoint. + 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. @@ -56,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/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 c9f1d91f94..133963ed16 100644 --- a/pyrit/prompt_target/azure_ml_chat_target.py +++ b/pyrit/prompt_target/azure_ml_chat_target.py @@ -197,13 +197,13 @@ def _initialize_vars( ``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) keeps the historical resolution order. + (the default) resolves a token-provider callable, then an explicit key, then the + environment variable. Raises: ValueError: If identity auth is requested for an endpoint that is not a recognized - Azure ML managed online endpoint, or 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. + 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 @@ -238,18 +238,11 @@ 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): - self._api_key_provider = self._build_azure_ml_token_provider() - 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]]: diff --git a/pyrit/prompt_target/openai/openai_target.py b/pyrit/prompt_target/openai/openai_target.py index 2dfb28720d..aa4b6810e8 100644 --- a/pyrit/prompt_target/openai/openai_target.py +++ b/pyrit/prompt_target/openai/openai_target.py @@ -103,9 +103,7 @@ 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. @@ -129,9 +127,8 @@ def __init__( Raises: ValueError: If identity auth is requested for an endpoint that is not a recognized - Azure OpenAI / AI Foundry endpoint, or if no API key is provided (via parameter or - environment variable) and the endpoint is not a recognized Azure OpenAI / - AI Foundry endpoint. + 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 {} diff --git a/pyrit/prompt_target/prompt_shield_target.py b/pyrit/prompt_target/prompt_shield_target.py index b1c9dcb53e..662a5d8d49 100644 --- a/pyrit/prompt_target/prompt_shield_target.py +++ b/pyrit/prompt_target/prompt_shield_target.py @@ -106,7 +106,7 @@ def __init__( Raises: 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 provided for a non-Azure Content Safety endpoint. + is available for ``"api_key"`` auth. """ endpoint_value = default_values.get_required_value( env_var_name=self.ENDPOINT_URI_ENVIRONMENT_VARIABLE, passed_value=endpoint @@ -121,9 +121,7 @@ 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). + # 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": @@ -142,13 +140,12 @@ 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 diff --git a/tests/unit/auth/test_openai_auth.py b/tests/unit/auth/test_openai_auth.py index 7eaa9cbee7..a411c68523 100644 --- a/tests/unit/auth/test_openai_auth.py +++ b/tests/unit/auth/test_openai_auth.py @@ -98,23 +98,40 @@ def sync_provider() -> str: assert resolved is not sync_provider -def test_api_key_mode_falls_back_to_entra_when_no_key(minted_provider): - provider, _ = minted_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: ""}): - resolved = resolve_openai_auth( - endpoint=AZURE_ENDPOINT, - api_key=None, - api_key_environment_variable=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, + ) - assert resolved is provider + 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="is required for non-Azure endpoints"): + 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/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 dce6876ef7..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,25 +55,18 @@ 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.""" - - async def _provider() -> str: - return "aml-entra-token" - +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", - return_value=_provider, ) as mock_provider, ): - target = AzureMLChatTarget(endpoint="https://my-aml.region.inference.ml.azure.com/score") + with pytest.raises(ValueError, match="No API key available"): + AzureMLChatTarget(endpoint="https://my-aml.region.inference.ml.azure.com/score") - mock_provider.assert_called_once_with(AzureMLChatTarget._AZURE_ML_SCOPE) - assert target._api_key_provider is _provider - assert target._api_key == "" + mock_provider.assert_not_called() def test_identity_auth_mode_ignores_env_key(patch_central_database): 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 2ac38424e1..6ab4b21b73 100644 --- a/tests/unit/prompt_target/target/test_openai_target_auth.py +++ b/tests/unit/prompt_target/target/test_openai_target_auth.py @@ -74,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.""" 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 94952ff8cc..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,23 +163,21 @@ 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():