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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 18 additions & 1 deletion doc/code/setup/1_configuration.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -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)."
]
},
{
Expand Down
17 changes: 17 additions & 0 deletions doc/code/setup/1_configuration.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
29 changes: 22 additions & 7 deletions pyrit/auth/openai_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,28 +6,45 @@

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(
*,
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))

Expand All @@ -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."
)
50 changes: 34 additions & 16 deletions pyrit/backend/services/target_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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:
"""
Expand Down Expand Up @@ -138,25 +142,24 @@ 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.

Args:
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]:
Expand Down Expand Up @@ -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.
Expand All @@ -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.
"""
Expand All @@ -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
Expand All @@ -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.
Expand Down
3 changes: 3 additions & 0 deletions pyrit/common/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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",
Expand Down
29 changes: 29 additions & 0 deletions pyrit/common/auth_mode.py
Original file line number Diff line number Diff line change
@@ -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")
16 changes: 11 additions & 5 deletions pyrit/embedding/openai_text_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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.
Expand All @@ -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
Expand All @@ -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,
Expand Down
Loading
Loading