feat(microsoft_365_copilot): add Microsoft 365 Copilot chat provider with OAuth token exchange (#45158)

* feat(microsoft_365_copilot): add Microsoft 365 Copilot chat provider with OAuth token exchange

* fix(ui): create a credential for credential-only auth types in Add Model

* fix(microsoft_365_copilot): accept max_tokens and list the chat model

* fix(microsoft_365_copilot): register provider model set

* fix(ui): hide other auth types' fields when credential-only types are filtered

* fix(health): do not inherit stored auth settings when a connection test brings its own api_key

* docs(health): restore connection-test inheritance docstrings

* fix(lint): suppress justified provider boundary casts

* fix(lint): satisfy M365 type-discipline gate

* fix(types): eliminate type-check gate regressions

* style: shorten suppression reasons so ruff format leaves them on one line

* refactor(types): drop shared-file type widening unrelated to the Copilot provider

* fix(health): pass caller headers as secret_fields to connection-test probes

* fix(types): type connection-test probe and token counter locally

* fix(lint): drop cast import from connection-test header getter

* chore: remove explanatory comments flagged by review

* fix(proxy): gate OAuth client credentials behind proxy admin

* fix(router): skip cooldown on caller-scoped OAuth auth failures

* fix(router): skip fallback cooldown on caller-scoped OAuth auth failures

* fix(router): type caller-scoped cooldown lookups

* fix(ci): include Microsoft 365 Copilot tests in unit shard

* fix(proxy): keep saved endpoint when test_connection retests a deployment by id

* fix(m365): collapse doubled Graph replies

* fix(m365): preserve trailing system messages

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Yassin Kortam 2026-10-09 15:19:07 -07:00 • committed by GitHub
parent 2248fcefda
commit fbf4af301c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
59 changed files with 4925 additions and 144 deletions

View file

@ -300,6 +300,7 @@ jobs:
tests/unit/llms/lm_studio
tests/unit/llms/manus
tests/unit/llms/meta_llama
tests/unit/llms/microsoft_365_copilot
tests/unit/llms/minimax
tests/unit/llms/mistral
tests/unit/llms/modelscope

View file

@ -711,6 +711,7 @@ docker_model_runner_models: Set = set()
amazon_nova_models: Set = set()
stability_models: Set = set()
github_copilot_models: Set = set()
microsoft_365_copilot_models: Set[str] = set()
chatgpt_models: Set = set()
minimax_models: Set = set()
aws_polly_models: Set = set()
@ -992,6 +993,8 @@ def _populate_provider_model_sets(model_cost_map: Dict) -> None:
stability_models.add(key)
elif value.get("litellm_provider") == "github_copilot":
github_copilot_models.add(key)
elif value.get("litellm_provider") == "microsoft_365_copilot":
microsoft_365_copilot_models.add(key)
elif value.get("litellm_provider") == "chatgpt":
chatgpt_models.add(key)
elif value.get("litellm_provider") == "minimax":
@ -1248,6 +1251,7 @@ def _build_models_by_provider() -> dict:
"amazon_nova": amazon_nova_models,
"stability": stability_models,
"github_copilot": github_copilot_models,
"microsoft_365_copilot": microsoft_365_copilot_models,
"chatgpt": chatgpt_models,
"minimax": minimax_models,
"aws_polly": aws_polly_models,

View file

@ -21,14 +21,14 @@ import time
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Generator, Mapping
from typing import TYPE_CHECKING, Any, Final, Optional, TypeVar
from pydantic import ConfigDict, SkipValidation, ValidationError
from pydantic import ConfigDict, SkipValidation, TypeAdapter, ValidationError
import litellm
from litellm._internal_context import post_response_phase
from litellm._logging import print_verbose, verbose_logger
from litellm.caching import InMemoryCache
from litellm.caching.caching import S3Cache, response_cache_phase
from litellm.constants import CACHE_WRITE_SHUTDOWN_FLUSH_TIMEOUT_SECONDS
from litellm.constants import CACHE_WRITE_SHUTDOWN_FLUSH_TIMEOUT_SECONDS, RESPONSE_CACHE_EXCLUDED_PROVIDERS
from litellm.litellm_core_utils.hidden_params import get_hidden_params
from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
update_response_metadata,
@ -106,6 +106,14 @@ def _drop_logging_obj_from_kwargs(request_kwargs: dict[str, object]) -> dict[str
return {k: v for k, v in request_kwargs.items() if k != "litellm_logging_obj"}
def _is_response_cache_excluded(model: str | None, kwargs: Mapping[str, object]) -> bool:
custom_llm_provider: Final = kwargs.get("custom_llm_provider")
model_provider: Final = model.split("/", maxsplit=1)[0] if model is not None else None
return (
isinstance(custom_llm_provider, str) and custom_llm_provider in RESPONSE_CACHE_EXCLUDED_PROVIDERS
) or model_provider in RESPONSE_CACHE_EXCLUDED_PROVIDERS
def _is_chat_completion_cached_dict(cached_result: dict) -> bool:
cached_id: Final = cached_result.get("id")
if isinstance(cached_id, str) and cached_id.startswith("chatcmpl"):
@ -308,6 +316,9 @@ class LLMCachingHandler:
Raises:
None
"""
if _is_response_cache_excluded(model=model, kwargs=kwargs):
return None
# Check if caching should be performed BEFORE doing expensive operations
if (
(kwargs.get("caching", None) is None and litellm.cache is not None) or kwargs.get("caching", False) is True
@ -432,6 +443,8 @@ class LLMCachingHandler:
cached_result: Any | None = None
# Check if caching should be performed BEFORE doing expensive kwargs copy
if _is_response_cache_excluded(model=model, kwargs=kwargs):
return CachingHandlerResponse(cached_result=None)
if litellm.cache is not None and self._is_call_type_supported_by_cache(original_function=original_function):
args = args or ()
# Now that we confirmed caching will happen, prepare kwargs
@ -852,6 +865,10 @@ class LLMCachingHandler:
)
if new_kwargs.get("metadata") is None:
new_kwargs.pop("metadata", None)
model_value: Final = new_kwargs.get("model")
model: Final = model_value if isinstance(model_value, str) else None
if _is_response_cache_excluded(model=model, kwargs=new_kwargs):
return None
if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs:
new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs)
self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs)
@ -1197,15 +1214,28 @@ class LLMCachingHandler:
return
def should_store_result_in_cache(self, original_function: Callable[..., object], kwargs: dict[str, Any]) -> bool:
def should_store_result_in_cache(
self, original_function: Callable[..., object], kwargs: Mapping[str, object]
) -> bool:
"""
Helper function to determine if the result should be stored in the cache.
Returns:
bool: True if the result should be stored in the cache, False otherwise.
"""
return self._is_call_type_supported_by_cache(original_function=original_function) and (
kwargs.get("cache", {}).get("no-store", False) is not True
model_value: Final = kwargs.get("model")
model: Final = model_value if isinstance(model_value, str) else None
cache_options_value: Final = kwargs.get("cache")
cache_options: Final[Mapping[str, object]] = (
TypeAdapter(Mapping[str, object]).validate_python(cache_options_value)
if isinstance(cache_options_value, Mapping)
else {}
)
no_store: Final = cache_options.get("no-store", False)
return (
not _is_response_cache_excluded(model=model, kwargs=kwargs)
and self._is_call_type_supported_by_cache(original_function=original_function)
and no_store is not True
)
_should_store_result_in_cache = should_store_result_in_cache
@ -1232,7 +1262,7 @@ class LLMCachingHandler:
def _is_call_type_supported_by_cache(
self,
original_function: Callable,
original_function: Callable[..., object],
) -> bool:
"""
Helper function to determine if the call type is supported by the cache.
@ -1260,6 +1290,11 @@ class LLMCachingHandler:
"""
if not self.should_store_result_in_cache(
original_function=self.original_function,
kwargs=self.request_kwargs,
):
return
complete_streaming_response: Final[ModelResponse | TextCompletionResponse | None] = (
assemble_complete_response_from_streaming_chunks(
result=processed_chunk,
@ -1284,6 +1319,11 @@ class LLMCachingHandler:
"""
Sync internal method to add the streaming response to the cache
"""
if not self.should_store_result_in_cache(
original_function=self.original_function,
kwargs=self.request_kwargs,
):
return
complete_streaming_response: Final[ModelResponse | TextCompletionResponse | None] = (
assemble_complete_response_from_streaming_chunks(
result=processed_chunk,

View file

@ -6,6 +6,13 @@ from typing import Final, Literal
from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_in_range, get_env_int_or_none
MICROSOFT_GRAPH_BETA_BASE: Final = "https://graph.microsoft.com/beta"
OAUTH_TOKEN_EXCHANGE_CACHE_SAFETY_MARGIN_SECONDS: Final = 60
MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_PROFILE: Final = "jwt_bearer_obo"
MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_SCOPE: Final = "https://graph.microsoft.com/.default"
# Replies depend on the caller's delegated identity, which response-cache keys do not include.
RESPONSE_CACHE_EXCLUDED_PROVIDERS: Final = frozenset({"microsoft_365_copilot"})
MICROSOFT_365_COPILOT_DEFAULT_TIME_ZONE: Final = "UTC"
SERVER_STREAMING_CLASSIFICATION_KEY: Final = "litellm_server_streaming_classification"
@ -14,7 +21,6 @@ class ServerStreamingClassification(str, Enum):
SERVER_STREAMING_CLASSIFICATION_MARKER: Final = ServerStreamingClassification.MARKER
DEFER_PYDANTIC_BUILD: Final = os.getenv("DEFER_PYDANTIC_BUILD", "true") in ("true", "1", "on")
DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"))
AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))

View file

@ -32,6 +32,14 @@ AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
)
PROVIDER_AFFINITY_HEADER_KWARG_KEY: Final = "provider_affinity_header"
OAUTH_TOKEN_EXCHANGE_KWARGS_KEYS: Final = frozenset(
{
"token_exchange_endpoint",
"token_exchange_profile",
"token_exchange_scope",
"token_exchange_audience",
}
)
# Pre-define optional kwargs keys as frozenset for O(1) lookups
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
@ -74,6 +82,7 @@ OPTIONAL_KWARGS_KEYS: Final = (
| AWS_CREDENTIAL_KWARGS_KEYS
| ANTHROPIC_WIF_KWARGS_KEYS
| OPENAI_WIF_KWARGS_KEYS
| OAUTH_TOKEN_EXCHANGE_KWARGS_KEYS
| frozenset(CustomPricingLiteLLMParams.model_fields)
)

View file

@ -0,0 +1,307 @@
import hashlib
import json
import re
from collections.abc import Mapping
from dataclasses import dataclass
from typing import (
Final,
Literal,
Protocol,
TypeAlias,
cast, # noqa: TID251 # adapter protocols cover untyped cache and HTTP handler methods
)
import httpx
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import OAUTH_TOKEN_EXCHANGE_CACHE_SAFETY_MARGIN_SECONDS
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
TokenExchangeProfile: TypeAlias = Literal["rfc8693", "jwt_bearer_obo"]
SUBJECT_TOKEN_TYPE_ACCESS_TOKEN: Final = "urn:ietf:params:oauth:token-type:access_token"
_RFC8693_GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:token-exchange"
_JWT_BEARER_GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:jwt-bearer"
_TOKEN_CACHE: Final = InMemoryCache(max_size_in_memory=1000, default_ttl=600)
class _TokenCache(Protocol):
def set_cache(self, *, key: str, value: str, ttl: int) -> None: ...
def get_cache(self, *, key: str) -> object: ...
class _SyncTokenExchangeClient(Protocol):
def post(
self,
url: str,
*,
data: Mapping[str, str],
timeout: float | httpx.Timeout | None,
) -> httpx.Response: ...
class _AsyncTokenExchangeClient(Protocol):
async def post(
self,
url: str,
*,
data: Mapping[str, str],
timeout: float | httpx.Timeout | None,
) -> httpx.Response: ...
class OAuthTokenExchangeError(Exception):
def __init__(self, status_code: int, message: str) -> None:
self.status_code: Final = status_code
self.message: Final = message
super().__init__(message)
class _OAuthTokenResponse(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
access_token: str
expires_in: int = Field(gt=0)
class _OAuthErrorResponse(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
error: str | None = None
error_description: str | None = None
_TOKEN_RESPONSE_ADAPTER: Final = TypeAdapter(_OAuthTokenResponse)
_ERROR_RESPONSE_ADAPTER: Final = TypeAdapter(_OAuthErrorResponse)
def _is_valid_token_endpoint(value: str) -> bool:
try:
endpoint: Final = httpx.URL(value)
except httpx.InvalidURL:
return False
return endpoint.scheme.casefold() == "https" and bool(endpoint.host)
@dataclass(frozen=True, slots=True, repr=False)
class OAuthTokenExchangeConfig:
token_endpoint: str
client_id: str
client_secret: str
profile: TokenExchangeProfile
scope: str | None
audience: str | None
def __post_init__(self) -> None:
if not _is_valid_token_endpoint(self.token_endpoint):
raise OAuthTokenExchangeError(
status_code=400,
message="token_exchange_endpoint must be an HTTPS URL with a host",
)
if self.profile not in ("rfc8693", "jwt_bearer_obo"):
raise OAuthTokenExchangeError(status_code=400, message="Unsupported OAuth token exchange profile")
if self.profile == "jwt_bearer_obo" and (self.scope is None or not self.scope.strip()):
raise OAuthTokenExchangeError(
status_code=400,
message="jwt_bearer_obo requires a non-empty token exchange scope",
)
def redact_sensitive_values(message: str, sensitive_values: tuple[str, ...]) -> str:
redacted_values: Final = tuple(sorted((value for value in sensitive_values if value), key=len, reverse=True))
if not redacted_values:
return message
pattern: Final = re.compile("|".join(re.escape(value) for value in redacted_values))
return pattern.sub("[REDACTED]", message)
def _cache_key(config: OAuthTokenExchangeConfig, subject_token: str) -> str:
key_material: Final = json.dumps(
{
"token_endpoint": config.token_endpoint,
"client_id": config.client_id,
"client_secret": config.client_secret,
"profile": config.profile,
"scope": config.scope,
"audience": config.audience,
"subject_token": subject_token,
},
ensure_ascii=True,
separators=(",", ":"),
sort_keys=True,
)
return hashlib.sha256(key_material.encode("utf-8")).hexdigest()
def build_token_exchange_form(
config: OAuthTokenExchangeConfig,
subject_token: str,
) -> dict[str, str]: # mutable-ok: the form builder contract requires an HTTP form dictionary
match config.profile:
case "rfc8693":
return {
"grant_type": _RFC8693_GRANT_TYPE,
"client_id": config.client_id,
"client_secret": config.client_secret,
"subject_token": subject_token,
"subject_token_type": SUBJECT_TOKEN_TYPE_ACCESS_TOKEN,
**({"scope": config.scope} if config.scope else {}),
**({"audience": config.audience} if config.audience else {}),
}
case "jwt_bearer_obo":
if config.scope is None or not config.scope.strip():
raise OAuthTokenExchangeError(
status_code=400,
message="jwt_bearer_obo requires a non-empty token exchange scope",
)
return {
"grant_type": _JWT_BEARER_GRANT_TYPE,
"client_id": config.client_id,
"client_secret": config.client_secret,
"assertion": subject_token,
"scope": config.scope,
"requested_token_use": "on_behalf_of",
}
def _error_message(response: httpx.Response, sensitive_values: tuple[str, ...]) -> str:
try:
error_response: Final = _ERROR_RESPONSE_ADAPTER.validate_python(response.json())
except (ValidationError, ValueError):
return redact_sensitive_values(
"OAuth token exchange failed: unknown_error: invalid response",
sensitive_values,
)
error: Final = error_response.error or "unknown_error"
description: Final = error_response.error_description or "no description provided"
return redact_sensitive_values(
f"OAuth token exchange failed: {error}: {description}",
sensitive_values,
)
def _response_token(response: httpx.Response, config: OAuthTokenExchangeConfig, subject_token: str) -> str:
try:
token_response: Final = _TOKEN_RESPONSE_ADAPTER.validate_python(response.json())
except (ValidationError, ValueError):
raise OAuthTokenExchangeError(
status_code=502,
message="OAuth token exchange returned an invalid token response",
) from None
effective_ttl: Final = token_response.expires_in - OAUTH_TOKEN_EXCHANGE_CACHE_SAFETY_MARGIN_SECONDS
if effective_ttl > 0:
cache: Final = cast(_TokenCache, _TOKEN_CACHE) # cast-ok: InMemoryCache exposes untyped cache methods
cache.set_cache(
key=_cache_key(config, subject_token),
value=token_response.access_token,
ttl=effective_ttl,
)
return token_response.access_token
def _error_status_code(status_code: int) -> int:
if status_code == 429:
return status_code
if 400 <= status_code < 500:
return 401
return status_code
def _post(
client: HTTPHandler,
config: OAuthTokenExchangeConfig,
request_data: Mapping[str, str],
timeout: float | httpx.Timeout | None,
) -> httpx.Response:
try:
http_client: Final = cast( # cast-ok: HTTPHandler.post has an untyped response contract
_SyncTokenExchangeClient,
client,
)
return http_client.post(config.token_endpoint, data=dict(request_data), timeout=timeout)
except httpx.HTTPStatusError as error:
return error.response
except httpx.HTTPError:
raise OAuthTokenExchangeError(
status_code=502,
message="OAuth token exchange request failed",
) from None
async def _apost(
client: AsyncHTTPHandler,
config: OAuthTokenExchangeConfig,
request_data: Mapping[str, str],
timeout: float | httpx.Timeout | None,
) -> httpx.Response:
try:
http_client: Final = cast( # cast-ok: AsyncHTTPHandler.post has an untyped response contract
_AsyncTokenExchangeClient,
client,
)
return await http_client.post(config.token_endpoint, data=dict(request_data), timeout=timeout)
except httpx.HTTPStatusError as error:
return error.response
except httpx.HTTPError:
raise OAuthTokenExchangeError(
status_code=502,
message="OAuth token exchange request failed",
) from None
def exchange_token(
client: HTTPHandler,
config: OAuthTokenExchangeConfig,
subject_token: str,
timeout: float | httpx.Timeout | None = None,
) -> str:
cache_key: Final = _cache_key(config, subject_token)
cache: Final = cast(_TokenCache, _TOKEN_CACHE) # cast-ok: InMemoryCache exposes untyped cache methods
cached_token: Final = cache.get_cache(key=cache_key)
if isinstance(cached_token, str):
return cached_token
response: Final = _post(
client=client,
config=config,
request_data=build_token_exchange_form(config, subject_token),
timeout=timeout,
)
if response.status_code != 200:
raise OAuthTokenExchangeError(
status_code=_error_status_code(response.status_code),
message=_error_message(
response=response,
sensitive_values=(config.client_secret, subject_token),
),
)
return _response_token(response=response, config=config, subject_token=subject_token)
async def aexchange_token(
client: AsyncHTTPHandler,
config: OAuthTokenExchangeConfig,
subject_token: str,
timeout: float | httpx.Timeout | None = None,
) -> str:
cache_key: Final = _cache_key(config, subject_token)
cache: Final = cast(_TokenCache, _TOKEN_CACHE) # cast-ok: InMemoryCache exposes untyped cache methods
cached_token: Final = cache.get_cache(key=cache_key)
if isinstance(cached_token, str):
return cached_token
response: Final = await _apost(
client=client,
config=config,
request_data=build_token_exchange_form(config, subject_token),
timeout=timeout,
)
if response.status_code != 200:
raise OAuthTokenExchangeError(
status_code=_error_status_code(response.status_code),
message=_error_message(
response=response,
sensitive_values=(config.client_secret, subject_token),
),
)
return _response_token(response=response, config=config, subject_token=subject_token)

View file

@ -0,0 +1,3 @@
from .chat.transformation import Microsoft365CopilotChatConfig
__all__ = ["Microsoft365CopilotChatConfig"]

View file

@ -0,0 +1,3 @@
from .transformation import Microsoft365CopilotChatConfig
__all__ = ["Microsoft365CopilotChatConfig"]

View file

@ -0,0 +1,580 @@
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import (
Final,
Protocol,
cast, # noqa: TID251 # adapter protocols cover pluggable logging and untyped HTTP methods
)
from urllib.parse import quote
import httpx
from aiohttp import ClientSession
from pydantic import TypeAdapter, ValidationError
from litellm import LlmProviders
from litellm.constants import (
MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_PROFILE,
MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_SCOPE,
MICROSOFT_GRAPH_BETA_BASE,
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.oauth_token_exchange import (
OAuthTokenExchangeConfig,
OAuthTokenExchangeError,
TokenExchangeProfile,
aexchange_token,
exchange_token,
)
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.llms.custom_httpx import http_handler
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
)
from litellm.llms.custom_httpx.llm_http_handler import MockResponseIterator
from litellm.llms.microsoft_365_copilot.chat.transformation import (
GraphChatRequest,
build_chat_request,
extract_graph_error_message,
map_graph_response,
parse_graph_conversation_id,
)
from litellm.llms.microsoft_365_copilot.common_utils import (
Microsoft365CopilotError,
extract_caller_assertion,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
from litellm.types.utils import ModelResponse
@dataclass(frozen=True, slots=True, repr=False)
class _TokenExchangeCredentials:
config: OAuthTokenExchangeConfig
subject_token: str
class _CopilotLogging(Protocol):
def pre_call(
self,
*,
input: Sequence[AllMessageValues],
api_key: str,
model: str,
additional_args: Mapping[str, object],
) -> object: ...
def post_call(
self,
*,
original_response: object,
input: Sequence[AllMessageValues],
api_key: str,
additional_args: Mapping[str, object],
) -> object: ...
class _SyncGraphClient(Protocol):
def post(
self,
url: str,
*,
headers: Mapping[str, str],
json: Mapping[str, object],
timeout: float | httpx.Timeout | None,
) -> httpx.Response: ...
class _AsyncGraphClient(Protocol):
async def post(
self,
url: str,
*,
headers: Mapping[str, str],
json: Mapping[str, object],
timeout: float | httpx.Timeout | None,
) -> httpx.Response: ...
class _HTTPClientFactory(Protocol):
def get_httpx_client(self, params: Mapping[str, object] | None = None) -> HTTPHandler: ...
def get_async_httpx_client(
self,
llm_provider: LlmProviders | str,
params: Mapping[str, object] | None = None,
shared_session: ClientSession | None = None,
) -> AsyncHTTPHandler: ...
_JSON_VALUE_ADAPTER: Final = TypeAdapter(object)
_HTTP_CLIENT_FACTORY: Final = cast( # cast-ok: shared cached-client methods have partially typed signatures
_HTTPClientFactory, http_handler
)
def _get_sync_http_client(params: Mapping[str, object] | None) -> HTTPHandler:
return _HTTP_CLIENT_FACTORY.get_httpx_client(dict(params) if params is not None else None)
def _get_async_http_client(
provider: LlmProviders | str,
params: Mapping[str, object] | None,
shared_session: ClientSession | None,
) -> AsyncHTTPHandler:
return _HTTP_CLIENT_FACTORY.get_async_httpx_client(
provider,
dict(params) if params is not None else None,
shared_session,
)
def _token_exchange_profile(value: object) -> TokenExchangeProfile:
if value is None:
return MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_PROFILE
if value == "rfc8693":
return "rfc8693"
if value == "jwt_bearer_obo":
return "jwt_bearer_obo"
raise Microsoft365CopilotError(
status_code=400,
message="microsoft_365_copilot token_exchange_profile must be 'rfc8693' or 'jwt_bearer_obo'",
)
def _get_token_exchange_credentials(
litellm_params: Mapping[str, object],
secret_fields: SecretFields | None,
) -> _TokenExchangeCredentials | None:
token_endpoint: Final[object] = litellm_params.get("token_exchange_endpoint")
client_id: Final[object] = litellm_params.get("client_id")
client_secret: Final[object] = litellm_params.get("client_secret")
credential_values: Final = (token_endpoint, client_id, client_secret)
configured_values: Final = tuple(value is not None and value != "" for value in credential_values)
if not any(configured_values):
return None
if not all(configured_values):
raise Microsoft365CopilotError(
status_code=400,
message="microsoft_365_copilot requires token_exchange_endpoint, client_id, and client_secret together",
)
if not isinstance(token_endpoint, str) or not isinstance(client_id, str) or not isinstance(client_secret, str):
raise Microsoft365CopilotError(
status_code=400,
message=(
"microsoft_365_copilot requires token_exchange_endpoint, client_id, and client_secret "
"as non-empty strings"
),
)
if not token_endpoint.strip() or not client_id.strip() or not client_secret.strip():
raise Microsoft365CopilotError(
status_code=400,
message="microsoft_365_copilot requires token_exchange_endpoint, client_id, and client_secret together",
)
profile: Final = _token_exchange_profile(litellm_params.get("token_exchange_profile"))
scope_value: Final[object] = litellm_params.get("token_exchange_scope")
if scope_value is not None and not isinstance(scope_value, str):
raise Microsoft365CopilotError(
status_code=400,
message="microsoft_365_copilot token_exchange_scope must be a string",
)
audience_value: Final[object] = litellm_params.get("token_exchange_audience")
if audience_value is not None and not isinstance(audience_value, str):
raise Microsoft365CopilotError(
status_code=400,
message="microsoft_365_copilot token_exchange_audience must be a string",
)
scope: Final = MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_SCOPE if scope_value is None else scope_value
try:
config: Final = OAuthTokenExchangeConfig(
token_endpoint=token_endpoint,
client_id=client_id,
client_secret=client_secret,
profile=profile,
scope=scope,
audience=audience_value,
)
except OAuthTokenExchangeError as error:
raise Microsoft365CopilotError(status_code=error.status_code, message=error.message) from error
subject_token: Final = extract_caller_assertion(secret_fields)
if subject_token is None:
raise Microsoft365CopilotError(
status_code=401,
message=(
"microsoft_365_copilot with OAuth token exchange requires the caller's "
"IdP-issued access token in the Authorization header"
),
)
return _TokenExchangeCredentials(config=config, subject_token=subject_token)
def _resolve_access_token(
client: HTTPHandler,
api_key: str | None,
litellm_params: Mapping[str, object],
secret_fields: SecretFields | None,
timeout: float | httpx.Timeout | None,
) -> str:
exchange_credentials: Final = _get_token_exchange_credentials(
litellm_params=litellm_params,
secret_fields=secret_fields,
)
if exchange_credentials is not None:
try:
return exchange_token(
client=client,
config=exchange_credentials.config,
subject_token=exchange_credentials.subject_token,
timeout=timeout,
)
except OAuthTokenExchangeError as error:
raise Microsoft365CopilotError(status_code=error.status_code, message=error.message) from error
if isinstance(api_key, str) and api_key:
return api_key
raise Microsoft365CopilotError(
status_code=401,
message="microsoft_365_copilot requires api_key or a complete token exchange configuration",
)
async def _aresolve_access_token(
client: AsyncHTTPHandler,
api_key: str | None,
litellm_params: Mapping[str, object],
secret_fields: SecretFields | None,
timeout: float | httpx.Timeout | None,
) -> str:
exchange_credentials: Final = _get_token_exchange_credentials(
litellm_params=litellm_params,
secret_fields=secret_fields,
)
if exchange_credentials is not None:
try:
return await aexchange_token(
client=client,
config=exchange_credentials.config,
subject_token=exchange_credentials.subject_token,
timeout=timeout,
)
except OAuthTokenExchangeError as error:
raise Microsoft365CopilotError(status_code=error.status_code, message=error.message) from error
if isinstance(api_key, str) and api_key:
return api_key
raise Microsoft365CopilotError(
status_code=401,
message="microsoft_365_copilot requires api_key or a complete token exchange configuration",
)
def _json_body(response: httpx.Response) -> object:
try:
return _JSON_VALUE_ADAPTER.validate_json(response.content)
except ValidationError:
return {}
def _graph_error(
response: httpx.Response,
sensitive_values: tuple[str, ...],
) -> Microsoft365CopilotError:
return Microsoft365CopilotError(
status_code=response.status_code,
message=extract_graph_error_message(
response_body=_json_body(response),
sensitive_values=sensitive_values,
),
)
def _pre_call(
logging_obj: Logging | None,
messages: Sequence[AllMessageValues],
model: str,
request_data: GraphChatRequest,
) -> None:
if logging_obj is None:
return
typed_logging: Final = cast( # cast-ok: Logging callback arguments are provider-pluggable
_CopilotLogging,
logging_obj,
)
typed_logging.pre_call(
input=messages,
api_key="",
model=model,
additional_args={
"api_base": MICROSOFT_GRAPH_BETA_BASE,
"complete_input_dict": request_data,
},
)
def _post_call(
logging_obj: Logging | None,
messages: Sequence[AllMessageValues],
request_data: GraphChatRequest,
status_code: int,
) -> None:
if logging_obj is None:
return
typed_logging: Final = cast( # cast-ok: Logging callback arguments are provider-pluggable
_CopilotLogging,
logging_obj,
)
typed_logging.post_call(
original_response={"status_code": status_code},
input=messages,
api_key="",
additional_args={
"api_base": MICROSOFT_GRAPH_BETA_BASE,
"complete_input_dict": request_data,
},
)
def _post_graph_request(
client: HTTPHandler,
url: str,
access_token: str,
request_data: GraphChatRequest | Mapping[str, object],
timeout: float | httpx.Timeout | None,
) -> httpx.Response:
try:
http_client: Final = cast( # cast-ok: HTTPHandler.post has an untyped response contract
_SyncGraphClient,
client,
)
request_mapping: Final[Mapping[str, object]] = request_data
return http_client.post(
url,
headers={"Authorization": f"Bearer {access_token}"},
json=dict(request_mapping),
timeout=timeout,
)
except httpx.HTTPStatusError as error:
return error.response
except httpx.HTTPError:
raise Microsoft365CopilotError(status_code=502, message="Microsoft Graph request failed") from None
async def _apost_graph_request(
client: AsyncHTTPHandler,
url: str,
access_token: str,
request_data: GraphChatRequest | Mapping[str, object],
timeout: float | httpx.Timeout | None,
) -> httpx.Response:
try:
http_client: Final = cast( # cast-ok: AsyncHTTPHandler.post has an untyped response contract
_AsyncGraphClient,
client,
)
request_mapping: Final[Mapping[str, object]] = request_data
return await http_client.post(
url,
headers={"Authorization": f"Bearer {access_token}"},
json=dict(request_mapping),
timeout=timeout,
)
except httpx.HTTPStatusError as error:
return error.response
except httpx.HTTPError:
raise Microsoft365CopilotError(status_code=502, message="Microsoft Graph request failed") from None
def _stream_response(
response: ModelResponse,
model: str,
logging_obj: Logging | None,
) -> CustomStreamWrapper:
if logging_obj is None:
raise Microsoft365CopilotError(
status_code=500,
message="Microsoft 365 Copilot streaming requires a logging context",
)
return CustomStreamWrapper(
completion_stream=MockResponseIterator(model_response=response),
model=model,
custom_llm_provider=LlmProviders.MICROSOFT_365_COPILOT.value,
logging_obj=logging_obj,
)
def _raise_graph_error_if_needed(
response: httpx.Response,
sensitive_values: tuple[str, ...],
logging_obj: Logging | None,
messages: Sequence[AllMessageValues],
request_data: GraphChatRequest,
) -> None:
if response.is_success:
return
_post_call(
logging_obj=logging_obj,
messages=messages,
request_data=request_data,
status_code=response.status_code,
)
raise _graph_error(response=response, sensitive_values=sensitive_values)
def _timeout_value(timeout: float | str | httpx.Timeout | None) -> float | httpx.Timeout | None:
if not isinstance(timeout, str):
return timeout
try:
return float(timeout)
except ValueError:
raise Microsoft365CopilotError(status_code=400, message="timeout must be a number") from None
def completion(
model: str,
messages: Sequence[AllMessageValues],
api_key: str | None,
litellm_params: Mapping[str, object],
optional_params: Mapping[str, object],
secret_fields: SecretFields | None = None,
stream: bool = False,
timeout: float | str | httpx.Timeout | None = None,
logging_obj: Logging | None = None,
client: HTTPHandler | None = None,
) -> ModelResponse | CustomStreamWrapper:
request_data: Final = build_chat_request(
messages=messages,
optional_params=optional_params,
)
http_client: Final = client if client is not None else _get_sync_http_client({})
timeout_value: Final = _timeout_value(timeout)
_pre_call(logging_obj=logging_obj, messages=messages, model=model, request_data=request_data)
access_token: Final = _resolve_access_token(
client=http_client,
api_key=api_key,
litellm_params=litellm_params,
secret_fields=secret_fields,
timeout=timeout_value,
)
assertion: Final = extract_caller_assertion(secret_fields)
configured_client_secret: Final[object] = litellm_params.get("client_secret")
sensitive_client_secret: Final = configured_client_secret if isinstance(configured_client_secret, str) else ""
sensitive_values: Final = (access_token, api_key or "", assertion or "", sensitive_client_secret)
conversation_url: Final = f"{MICROSOFT_GRAPH_BETA_BASE}/copilot/conversations"
conversation_response: Final = _post_graph_request(
client=http_client,
url=conversation_url,
access_token=access_token,
request_data={},
timeout=timeout_value,
)
_raise_graph_error_if_needed(
response=conversation_response,
sensitive_values=sensitive_values,
logging_obj=logging_obj,
messages=messages,
request_data=request_data,
)
conversation_id: Final = parse_graph_conversation_id(_json_body(conversation_response))
chat_url: Final = f"{conversation_url}/{quote(conversation_id, safe='')}/chat"
chat_response: Final = _post_graph_request(
client=http_client,
url=chat_url,
access_token=access_token,
request_data=request_data,
timeout=timeout_value,
)
_raise_graph_error_if_needed(
response=chat_response,
sensitive_values=sensitive_values,
logging_obj=logging_obj,
messages=messages,
request_data=request_data,
)
_post_call(
logging_obj=logging_obj,
messages=messages,
request_data=request_data,
status_code=chat_response.status_code,
)
response: Final = map_graph_response(
graph_response=_json_body(chat_response),
model=model,
messages=messages,
)
return _stream_response(response, model, logging_obj) if stream else response
async def acompletion(
model: str,
messages: Sequence[AllMessageValues],
api_key: str | None,
litellm_params: Mapping[str, object],
optional_params: Mapping[str, object],
secret_fields: SecretFields | None = None,
stream: bool = False,
timeout: float | str | httpx.Timeout | None = None,
logging_obj: Logging | None = None,
client: AsyncHTTPHandler | None = None,
shared_session: ClientSession | None = None,
) -> ModelResponse | CustomStreamWrapper:
request_data: Final = build_chat_request(
messages=messages,
optional_params=optional_params,
)
http_client: Final = (
client if client is not None else _get_async_http_client(LlmProviders.MICROSOFT_365_COPILOT, {}, shared_session)
)
timeout_value: Final = _timeout_value(timeout)
_pre_call(logging_obj=logging_obj, messages=messages, model=model, request_data=request_data)
access_token: Final = await _aresolve_access_token(
client=http_client,
api_key=api_key,
litellm_params=litellm_params,
secret_fields=secret_fields,
timeout=timeout_value,
)
assertion: Final = extract_caller_assertion(secret_fields)
configured_client_secret: Final[object] = litellm_params.get("client_secret")
sensitive_client_secret: Final = configured_client_secret if isinstance(configured_client_secret, str) else ""
sensitive_values: Final = (access_token, api_key or "", assertion or "", sensitive_client_secret)
conversation_url: Final = f"{MICROSOFT_GRAPH_BETA_BASE}/copilot/conversations"
conversation_response: Final = await _apost_graph_request(
client=http_client,
url=conversation_url,
access_token=access_token,
request_data={},
timeout=timeout_value,
)
_raise_graph_error_if_needed(
response=conversation_response,
sensitive_values=sensitive_values,
logging_obj=logging_obj,
messages=messages,
request_data=request_data,
)
conversation_id: Final = parse_graph_conversation_id(_json_body(conversation_response))
chat_url: Final = f"{conversation_url}/{quote(conversation_id, safe='')}/chat"
chat_response: Final = await _apost_graph_request(
client=http_client,
url=chat_url,
access_token=access_token,
request_data=request_data,
timeout=timeout_value,
)
_raise_graph_error_if_needed(
response=chat_response,
sensitive_values=sensitive_values,
logging_obj=logging_obj,
messages=messages,
request_data=request_data,
)
_post_call(
logging_obj=logging_obj,
messages=messages,
request_data=request_data,
status_code=chat_response.status_code,
)
response: Final = map_graph_response(
graph_response=_json_body(chat_response),
model=model,
messages=messages,
)
return _stream_response(response, model, logging_obj) if stream else response

View file

@ -0,0 +1,332 @@
from collections.abc import Mapping, Sequence
from typing import (
TYPE_CHECKING,
Final,
Literal,
Protocol,
cast, # noqa: TID251 # AllMessageValues variants are read-only TypedDict mappings
)
import httpx
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm.utils as litellm_utils
from litellm.constants import MICROSOFT_365_COPILOT_DEFAULT_TIME_ZONE
from litellm.litellm_core_utils.oauth_token_exchange import redact_sensitive_values
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.microsoft_365_copilot.common_utils import Microsoft365CopilotError
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import Choices, Message, ModelResponse, Usage
class _TokenCounter(Protocol):
def __call__(
self,
model: str = "",
*,
messages: Sequence[AllMessageValues | Message] | None = None,
text: str | None = None,
count_response_tokens: bool | None = False,
) -> int: ...
_TOKEN_COUNTER: Final = cast( # cast-ok: only the typed token-counter arguments used here are narrowed locally
_TokenCounter, litellm_utils.token_counter
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer
class _TextPart(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
type: Literal["text"]
text: str
class _GraphConversation(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
id: str
class _GraphMessage(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
text: str | None = None
class _GraphConversationResponse(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
messages: Sequence[_GraphMessage] | None = None
class _GraphError(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
message: str | None = None
class _GraphErrorResponse(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
error: _GraphError | None = None
class _GraphRequestMessage(TypedDict):
text: ReadOnly[str]
class _GraphAdditionalContext(TypedDict):
text: ReadOnly[str]
description: ReadOnly[str]
class _GraphLocationHint(TypedDict):
timeZone: ReadOnly[str]
class GraphChatRequest(TypedDict):
message: ReadOnly[_GraphRequestMessage]
locationHint: ReadOnly[_GraphLocationHint]
additionalContext: NotRequired[ReadOnly[Sequence[_GraphAdditionalContext]]]
_TEXT_PART_ADAPTER: Final = TypeAdapter(_TextPart)
_GRAPH_CONVERSATION_ADAPTER: Final = TypeAdapter(_GraphConversation)
_GRAPH_RESPONSE_ADAPTER: Final = TypeAdapter(_GraphConversationResponse)
_GRAPH_ERROR_ADAPTER: Final = TypeAdapter(_GraphErrorResponse)
_GRAPH_CHAT_REQUEST_ADAPTER: Final = TypeAdapter(GraphChatRequest)
_CONTENT_PARTS_ADAPTER: Final = TypeAdapter(list[object])
_JSON_VALUE_ADAPTER: Final = TypeAdapter(object)
class Microsoft365CopilotChatConfig(BaseConfig):
def get_supported_openai_params(self, model: str) -> list[str]: # mutable-ok: BaseConfig requires a list return
return ["stream", "max_tokens", "max_completion_tokens"]
def transform_request(
self,
model: str,
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
headers: Mapping[str, object],
) -> dict[str, object]: # mutable-ok: BaseConfig requires a mutable JSON dictionary
return {**build_chat_request(messages=messages, optional_params=optional_params)}
def transform_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: "Logging",
request_data: Mapping[str, object],
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
encoding: "Tokenizer | None",
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:
return map_graph_response(
graph_response=_response_json(raw_response.content),
model=model,
messages=messages,
)
def get_error_class(
self,
error_message: str,
status_code: int,
headers: Mapping[str, object] | httpx.Headers,
) -> BaseLLMException:
return Microsoft365CopilotError(status_code=status_code, message=error_message)
def map_openai_params(
self,
non_default_params: Mapping[str, object],
optional_params: Mapping[str, object],
model: str,
drop_params: bool,
) -> dict[str, object]: # mutable-ok: BaseConfig requires a mutable parameter dictionary
return dict(optional_params)
def validate_environment(
self,
headers: Mapping[str, object],
model: str,
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
) -> dict[str, object]: # mutable-ok: BaseConfig requires a mutable header dictionary
return dict(headers)
def _content_to_text(content: object) -> str:
if isinstance(content, str):
return content
if not isinstance(content, list):
raise Microsoft365CopilotError(status_code=400, message="only text content is supported")
try:
content_parts: Final = _CONTENT_PARTS_ADAPTER.validate_python(content)
text_parts: Final = tuple(_TEXT_PART_ADAPTER.validate_python(part) for part in content_parts)
except ValidationError:
raise Microsoft365CopilotError(status_code=400, message="only text content is supported") from None
return "\n".join(part.text for part in text_parts)
def _message_content(message: AllMessageValues) -> object:
message_fields: Final = cast( # cast-ok: AllMessageValues variants are read-only TypedDict mappings
Mapping[str, object],
message,
)
return message_fields.get("content")
def _response_json(response_content: bytes) -> object:
try:
return _JSON_VALUE_ADAPTER.validate_json(response_content)
except ValidationError:
return {}
def build_chat_request(
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
) -> GraphChatRequest:
last_user_index: Final = next(
(index for index in reversed(range(len(messages))) if messages[index]["role"] == "user"),
None,
)
if last_user_index is None:
raise Microsoft365CopilotError(status_code=400, message="at least one message must have role 'user'")
last_user_message: Final = messages[last_user_index]
history: Final = tuple(
{
"text": _content_to_text(_message_content(message)),
"description": f"{message['role']} message",
}
for index, message in enumerate(messages)
if index != last_user_index
)
requested_time_zone: Final = optional_params.get("time_zone")
time_zone: Final = (
requested_time_zone
if isinstance(requested_time_zone, str) and requested_time_zone
else MICROSOFT_365_COPILOT_DEFAULT_TIME_ZONE
)
request_data: Final = {
"message": {"text": _content_to_text(_message_content(last_user_message))},
"locationHint": {"timeZone": time_zone},
**({"additionalContext": list(history)} if history else {}),
}
try:
return _GRAPH_CHAT_REQUEST_ADAPTER.validate_python(request_data)
except ValidationError:
raise Microsoft365CopilotError(status_code=400, message="only text content is supported") from None
def _collapse_doubled_reply(text: str) -> str:
if not text or len(text) % 2 != 0:
return text
half: Final = len(text) // 2
return text[:half] if text[:half] == text[half:] else text
def map_graph_response(
graph_response: object,
model: str,
messages: Sequence[AllMessageValues],
) -> ModelResponse:
try:
parsed_response: Final = _GRAPH_RESPONSE_ADAPTER.validate_python(graph_response)
except ValidationError:
raise Microsoft365CopilotError(
status_code=502,
message="Microsoft Graph returned an invalid Copilot conversation response",
) from None
if not parsed_response.messages:
raise Microsoft365CopilotError(
status_code=502,
message="Microsoft Graph returned a Copilot response without messages",
)
reply: Final = parsed_response.messages[-1].text
if reply is None:
raise Microsoft365CopilotError(
status_code=502,
message="Microsoft Graph returned a Copilot response without reply text",
)
normalized_reply: Final = _collapse_doubled_reply(reply)
prompt_tokens: Final = _TOKEN_COUNTER(model=model, messages=messages)
completion_tokens: Final = _TOKEN_COUNTER(
model=model,
text=normalized_reply,
count_response_tokens=True,
)
return ModelResponse(
model=model,
choices=[
Choices(
index=0,
message=Message(content=normalized_reply, role="assistant"),
finish_reason="stop",
)
],
usage=Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
),
)
def parse_graph_conversation_id(graph_conversation: object) -> str:
try:
conversation: Final = _GRAPH_CONVERSATION_ADAPTER.validate_python(graph_conversation)
except ValidationError:
raise Microsoft365CopilotError(
status_code=502,
message="Microsoft Graph returned an invalid Copilot conversation response",
) from None
if not conversation.id:
raise Microsoft365CopilotError(
status_code=502,
message="Microsoft Graph returned a Copilot conversation without an id",
)
return conversation.id
def _last_message_text(graph_response: object) -> str | None:
try:
parsed_response: Final = _GRAPH_RESPONSE_ADAPTER.validate_python(graph_response)
except ValidationError:
return None
if not parsed_response.messages:
return None
return parsed_response.messages[-1].text
def extract_graph_error_message(
response_body: object,
sensitive_values: tuple[str, ...] = (),
) -> str:
try:
parsed_error: Final = _GRAPH_ERROR_ADAPTER.validate_python(response_body)
except ValidationError:
return "Microsoft Graph request failed"
message: Final = parsed_error.error.message if parsed_error.error is not None else None
if message is None:
return "Microsoft Graph request failed"
try:
nested_response: Final = _JSON_VALUE_ADAPTER.validate_json(message)
except ValidationError:
return redact_sensitive_values(message, sensitive_values)
nested_reply: Final = _last_message_text(nested_response)
return redact_sensitive_values(nested_reply if nested_reply is not None else message, sensitive_values)

View file

@ -0,0 +1,47 @@
from collections.abc import Mapping
from typing import Final
from pydantic import TypeAdapter, ValidationError
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
class Microsoft365CopilotError(BaseLLMException):
pass
_RAW_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str])
_SECRET_FIELDS_ADAPTER: Final = TypeAdapter(SecretFields)
def extract_caller_assertion(secret_fields: Mapping[str, object] | None) -> str | None:
if secret_fields is None:
return None
raw_headers: Final = secret_fields.get("raw_headers")
try:
headers: Final = _RAW_HEADERS_ADAPTER.validate_python(raw_headers)
except ValidationError:
return None
authorization_values: Final = tuple(value for name, value in headers.items() if name.casefold() == "authorization")
if len(authorization_values) != 1:
return None
authorization_parts: Final = authorization_values[0].split(maxsplit=1)
if len(authorization_parts) != 2 or authorization_parts[0].casefold() != "bearer":
return None
assertion: Final = authorization_parts[1].strip()
assertion_parts: Final = tuple(assertion.split("."))
if (
len(assertion_parts) != 3
or any(not part for part in assertion_parts)
or any(character.isspace() for character in assertion)
):
return None
return assertion
def as_secret_fields(secret_fields: object) -> SecretFields | None:
try:
return _SECRET_FIELDS_ADAPTER.validate_python(secret_fields)
except ValidationError:
return None

View file

@ -83,6 +83,7 @@ from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.get_litellm_params import (
AWS_CREDENTIAL_KWARGS_KEYS,
OAUTH_TOKEN_EXCHANGE_KWARGS_KEYS,
OPTIONAL_KWARGS_KEYS,
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
InvalidControlOption,
@ -140,6 +141,7 @@ from litellm.types.completion import (
_CompletionDispatchResult,
)
from litellm.types.litellm_params import ControlOptions, RetryStrategy
from litellm.types.llms.openai import AllMessageValues
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
CustomPricingLiteLLMParams,
@ -5102,6 +5104,56 @@ def _complete_langgraph(ctx: CompletionDispatchContext) -> _CompletionDispatchRe
)
def _complete_microsoft_365_copilot(ctx: CompletionDispatchContext) -> _CompletionDispatchResult:
from litellm.llms.microsoft_365_copilot.chat.handler import acompletion, completion
from litellm.llms.microsoft_365_copilot.common_utils import as_secret_fields
messages: Final[Sequence[AllMessageValues]] = cast( # cast-ok: dispatch messages are prevalidated
Sequence[AllMessageValues],
ctx.messages,
)
optional_params_value: Final[object] = cast( # cast-ok: dispatch params are validated below
object,
ctx.optional_params,
)
optional_params: Final = TypeAdapter(Mapping[str, object]).validate_python(optional_params_value)
kwargs_value: Final[object] = cast( # cast-ok: dispatch kwargs are validated below
object,
ctx.kwargs,
)
kwargs: Final = TypeAdapter(Mapping[str, object]).validate_python(kwargs_value)
secret_fields: Final = as_secret_fields(kwargs.get("secret_fields"))
client: Final = _dispatch_client_http(ctx)
if ctx.acompletion:
async_client: Final = client if isinstance(client, AsyncHTTPHandler) else None
return acompletion(
model=ctx.model,
messages=messages,
api_key=ctx.api_key,
litellm_params=ctx.litellm_params,
optional_params=optional_params,
secret_fields=secret_fields,
stream=ctx.stream is True,
timeout=ctx.timeout,
logging_obj=ctx.logging,
client=async_client,
shared_session=ctx.shared_session,
)
sync_client: Final = client if isinstance(client, HTTPHandler) else None
return completion(
model=ctx.model,
messages=messages,
api_key=ctx.api_key,
litellm_params=ctx.litellm_params,
optional_params=optional_params,
secret_fields=secret_fields,
stream=ctx.stream is True,
timeout=ctx.timeout,
logging_obj=ctx.logging,
client=sync_client,
)
def _complete_langflow(ctx: CompletionDispatchContext) -> _CompletionDispatchResult:
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
@ -5737,6 +5789,7 @@ def completion(
*AWS_CREDENTIAL_KWARGS_KEYS,
*ANTHROPIC_WIF_KWARGS_KEYS,
*OPENAI_WIF_KWARGS_KEYS,
*OAUTH_TOKEN_EXCHANGE_KWARGS_KEYS,
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
"fireworks_forward_user_id",
)
@ -6091,6 +6144,9 @@ def completion(
# LangGraph - Agent Runtime Provider
response = _complete_langgraph(_dispatch_ctx)
elif custom_llm_provider == "microsoft_365_copilot":
response = _complete_microsoft_365_copilot(_dispatch_ctx) # rebind-ok: shared provider return
elif custom_llm_provider == "langflow":
# LangFlow - Visual AI Agent Platform
response = _complete_langflow(_dispatch_ctx)

View file

@ -78859,6 +78859,13 @@
"prompt_cache_min_tokens": 512,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing"
},
"microsoft_365_copilot/chat": {
"litellm_provider": "microsoft_365_copilot",
"mode": "chat",
"supports_mid_conversation_system": true,
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0
},
"openrouter/inclusionai/ling-3.0-flash-sante": {
"cache_read_input_token_cost": 8.4e-09,
"input_cost_per_token": 4.2e-08,

View file

@ -269,7 +269,8 @@ def reject_federated_credential_reference(body: Mapping[str, object]) -> None:
if wif_fields:
raise ValueError(
f"Rejected Request: litellm_credential_name={named!r} names a credential configured for "
f"workload identity federation ({wif_fields[0]}), which a request body cannot choose. "
f"workload identity federation or OAuth token exchange ({wif_fields[0]}), which a request body "
"cannot choose. "
"A proxy admin attaches it to a deployment."
)

View file

@ -28,6 +28,8 @@ _LITELLM_PROVIDER_IDS: Final = frozenset(provider.value for provider in LlmProvi
_FEDERATION_SURFACE_FIELDS: Final = frozenset(
(
*clientside_credential_keys,
"client_id",
"client_secret",
"configurable_clientside_auth_params",
"litellm_credential_name",
*server_owned_wif_litellm_params,
@ -36,15 +38,17 @@ _FEDERATION_SURFACE_FIELDS: Final = frozenset(
def write_touches_federation_surface(incoming: Mapping[str, object] | None) -> bool:
"""Whether this write can move or re-scope the token a federated deployment mints.
"""Whether this write can move or re-scope a deployment's delegated token.
Three groups of fields can. The federation parameters choose which server-side secret is read
and what the minted token is scoped to. ``litellm_credential_name`` resolves to those same
parameters by reference. ``api_key``, ``api_base``, ``base_url``, and the
and what the minted token is scoped to. OAuth token-exchange parameters select the server-side
token endpoint and grant profile, with ``client_id`` and ``client_secret`` authenticating to it.
``litellm_credential_name`` resolves to those same parameters
by reference. ``api_key``, ``api_base``, ``base_url``, and the
``configurable_clientside_auth_params`` that let a caller override them decide where the
resulting token is sent. A write setting none of them leaves the federation configuration
exactly as the proxy admin left it, so renaming a federated deployment or changing its rpm
stays an ordinary team-admin edit.
resulting token is sent. A write setting none of them leaves the delegated-token configuration
exactly as the proxy admin left it, so renaming a deployment or changing its rpm stays an
ordinary team-admin edit.
"""
return incoming is not None and not _FEDERATION_SURFACE_FIELDS.isdisjoint(incoming.keys())

View file

@ -53,15 +53,16 @@ def _reject_non_admin_wif_fields(
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""A credential referenced by ``litellm_credential_name`` feeds its values into the same
workload identity federation resolution as a deployment's own ``litellm_params``. Only proxy
admins may touch a server-owned WIF field, whether they write it, drop it, or edit a stored
credential that already carries one.
server-owned federation or OAuth token-exchange configuration as a deployment's own
``litellm_params``. Only proxy admins may touch one, whether they write it, drop it, or edit a
stored credential that already carries it.
"""
if not wif_fields or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
return
raise ProxyException(
message=(
f"Only proxy admins can change {wif_fields[0]!r}, a server-owned workload identity federation parameter."
f"Only proxy admins can change {wif_fields[0]!r}, a server-owned workload identity federation "
"or OAuth token exchange parameter."
),
type=ProxyErrorTypes.auth_error.value,
code=status.HTTP_403_FORBIDDEN,

View file

@ -34,9 +34,9 @@ from litellm.router_utils.auto_router_model_naming import (
)
from litellm.types.utils import secret_bearing_wif_litellm_params, server_owned_wif_litellm_params
# Provider routing and workload identity federation fields. Allowed for proxy admins so they can
# see which region/version a deployment is checking and which identity it federates as; gated at
# the endpoint layer for non-admin callers (see _strip_admin_only_fields_from_health_result).
# Provider routing and server-owned federation or OAuth token-exchange fields. Allowed for proxy
# admins so they can see which region/version a deployment is checking and which identity it uses;
# gated at the endpoint layer for non-admin callers (see _strip_admin_only_fields_from_health_result).
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS: Final = (
"api_base",
"api_version",

View file

@ -46,6 +46,7 @@ from litellm.proxy.auth.auth_utils import (
)
from litellm.proxy.auth.model_checks import get_key_models
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils import http_parsing_utils
from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.db.health_check_latest import (
@ -78,6 +79,7 @@ from litellm.router_utils.clientside_credential_handler import (
clientside_credential_keys,
)
from litellm.secret_managers.main import get_secret_bool
from litellm.types.proxy.litellm_pre_call_utils import RedactedDict, SecretFields
#### Health ENDPOINTS ####
@ -135,20 +137,25 @@ def _request_inherits_config_credentials(
config_params: Mapping[str, object],
request_params: Mapping[str, object],
allow_client_side_credentials: bool,
*,
selected_by_id: bool,
) -> bool:
"""Whether the configuration's credentials are this request's to be probed with.
The configuration reached here by matching the request's model string, which
also matches wildcard routes and unrelated deployments that merely serve the
same model, so a request naming a stored credential of its own has already
said where its credentials come from and does not borrow that one's. A blank
name is no name: ``load_credentials_from_list`` resolves nothing from it, so
it must not cost the request the credentials it would otherwise be probed
with.
said where its credentials come from and does not borrow that one's. A request
bringing its own ``api_key`` describes its own connection when matched only by
model name. A blank name or key is no name or key: ``load_credentials_from_list``
resolves nothing from it, so it must not cost the request the credentials it
would otherwise be probed with.
"""
requested_credential: Final = request_params.get("litellm_credential_name")
if requested_credential and requested_credential != config_params.get("litellm_credential_name"):
return False
if not selected_by_id and request_params.get("api_key"):
return False
if allow_client_side_credentials:
return True
return not any(param in request_params for param in _BANNED_REQUEST_BODY_PARAMS)
@ -158,22 +165,28 @@ def _config_base_for_health_check(
config_params: Mapping[str, object],
request_params: Mapping[str, object],
allow_client_side_credentials: bool = False,
*,
selected_by_id: bool,
) -> dict[str, object]:
"""Return the configured parameters to merge under a connection-test request.
A request that sets its own connection fields, or names its own stored
credential, describes a connection of its own, so the configuration's
credentials are not carried into it: they belong to the endpoint the
configuration names. Anything the request does not set still comes from the
configuration, which is what lets a request name a configured model and test
it as configured.
A request that sets its own connection fields, brings an ``api_key`` when
matched only by model name, or names its own stored credential describes a
connection of its own, so the configuration's credentials are not carried
into it. Anything the request does not set still comes from the configuration,
which is what lets a request name a configured model and test it as configured.
``litellm_credential_name`` is dropped alongside the literal credential
fields: it names a stored credential that ``load_credentials_from_list``
resolves into the same secrets further down the call, so leaving it in place
would reintroduce them by reference.
"""
if _request_inherits_config_credentials(config_params, request_params, allow_client_side_credentials):
if _request_inherits_config_credentials(
config_params,
request_params,
allow_client_side_credentials,
selected_by_id=selected_by_id,
):
return dict(config_params)
return {key: value for key, value in config_params.items() if key not in _CONFIG_CONNECTION_FIELDS}
@ -978,11 +991,10 @@ def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
def _strip_admin_only_fields_from_health_result(result: dict) -> dict:
"""
Return a copy of the /health response with the admin-only fields (provider routing plus the
workload identity federation params naming the identity a deployment mints as) removed from
each healthy/unhealthy endpoint entry. Used to hide those fields from non-admin callers while
still showing them which deployments they own and whether each one is
healthy. Proxy admins receive the unmodified result.
Return a copy of the /health response with the admin-only fields (provider routing plus
server-owned federation or OAuth token-exchange params) removed from each healthy/unhealthy
endpoint entry. Used to hide those fields from non-admin callers while still showing them which
deployments they own and whether each one is healthy. Proxy admins receive the unmodified result.
"""
out: Final = dict(result)
drop: Final = set(ADMIN_ONLY_HEALTH_DISPLAY_PARAMS)
@ -2180,6 +2192,7 @@ async def test_model_connection(
# This gets the litellm_params from proxy config (with resolved env vars)
config_litellm_params: dict = {}
loaded_model_info: dict | None = None
deployment_by_id: Deployment | None = None
if llm_router is not None:
# Prefer disambiguation by deployment id (`model_info.id`) when
# the caller supplies it. This is required when multiple
@ -2192,7 +2205,6 @@ async def test_model_connection(
request_model_info: Final = model_info or {}
request_model_id: Final = request_model_info.get("id")
try:
deployment_by_id = None
if request_model_id:
deployment_by_id = llm_router.get_deployment(model_id=request_model_id)
@ -2226,6 +2238,7 @@ async def test_model_connection(
"Could not find model %s in router: %s. Proceeding with request params only.", model_name, e
)
selected_by_id: Final = deployment_by_id is not None
reject_server_owned_wif_params(request_litellm_params)
# Merge: config params (from proxy config) as base, request params override
litellm_params = {
@ -2233,6 +2246,7 @@ async def test_model_connection(
config_litellm_params,
request_litellm_params,
allow_client_side_credentials=general_settings.get("allow_client_side_credentials") is True,
selected_by_id=selected_by_id,
),
**request_litellm_params,
}
@ -2270,9 +2284,17 @@ async def test_model_connection(
or resolve_health_check_mode(probe_model_info, _OBJECT_MAPPING.validate_python(litellm_params))
)
raw_headers: Final[dict[str, str]] = TypeAdapter(dict[str, str]).validate_python(
http_parsing_utils.safe_get_request_headers(request)
)
health_check_params: Final = {
**_OBJECT_MAPPING.validate_python(litellm_params),
"secret_fields": SecretFields(raw_headers=RedactedDict(raw_headers)),
}
result: Final = await run_with_timeout(
litellm.ahealth_check(
model_params=litellm_params,
model_params=health_check_params,
mode=probe_mode,
prompt="test from litellm",
input=["test from litellm"],
@ -2281,7 +2303,7 @@ async def test_model_connection(
)
# Clean the result for display
cleaned_result: Final = clean_endpoint_data({**litellm_params, **result}, details=True)
cleaned_result: Final = clean_endpoint_data({**health_check_params, **result}, details=True)
return {
"status": "error" if "error" in result else "success",

View file

@ -2190,7 +2190,7 @@ class ModelManagementAuthChecks:
raise ProxyException(
message=(
f"Only proxy admins can change the credentials of a deployment configured for "
f"workload identity federation ({wif_fields[0]!r})."
f"workload identity federation or OAuth token exchange ({wif_fields[0]!r})."
),
type=ProxyErrorTypes.auth_error.value,
code=status.HTTP_403_FORBIDDEN,

View file

@ -2135,6 +2135,84 @@
],
"default_model_placeholder": "gpt-3.5-turbo"
},
{
"provider": "MICROSOFT_365_COPILOT",
"provider_display_name": "Microsoft 365 Copilot",
"litellm_provider": "microsoft_365_copilot",
"credential_fields": [
{
"key": "token_exchange_endpoint",
"label": "Token Endpoint URL",
"placeholder": "https://login.microsoftonline.com/<tenant-id>/oauth2/v2.0/token",
"tooltip": "Your IdP's OAuth 2.0 token endpoint. For Microsoft 365 Copilot this must be Microsoft Entra.",
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "token_exchange_profile",
"label": "Exchange Grant",
"placeholder": null,
"tooltip": "jwt_bearer_obo: on-behalf-of with the JWT bearer grant (Microsoft Entra). rfc8693: standard OAuth 2.0 token exchange.",
"required": false,
"field_type": "select",
"options": ["jwt_bearer_obo", "rfc8693"],
"default_value": "jwt_bearer_obo"
},
{
"key": "client_id",
"label": "Client ID",
"placeholder": null,
"tooltip": null,
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "client_secret",
"label": "Client Secret",
"placeholder": null,
"tooltip": null,
"required": false,
"field_type": "password",
"options": null,
"default_value": null
},
{
"key": "token_exchange_scope",
"label": "Scope",
"placeholder": null,
"tooltip": null,
"required": false,
"field_type": "text",
"options": null,
"default_value": "https://graph.microsoft.com/.default"
},
{
"key": "token_exchange_audience",
"label": "Audience",
"placeholder": null,
"tooltip": "Optional. Sent only with the rfc8693 grant.",
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "api_key",
"label": "Delegated Access Token",
"placeholder": null,
"tooltip": null,
"required": false,
"field_type": "password",
"options": null,
"default_value": null
}
],
"default_model_placeholder": "microsoft_365_copilot/chat"
},
{
"provider": "MistralAI",
"provider_display_name": "Mistral AI",

View file

@ -210,6 +210,7 @@ from litellm.router_utils.cooldown_handlers import (
get_cooldown_deployments,
is_advisor_orchestration_failure,
is_background_response_cost_poll_not_found,
is_caller_scoped_auth_failure,
is_caller_timeout_408,
set_cooldown_deployments,
)
@ -8603,6 +8604,24 @@ class Router:
)
return False
caller_failure_model_info: Final = (
_MODEL_INFO_ADAPTER.validate_python(_model_info) if isinstance(_model_info, dict) else None
)
caller_failure_model_id: Final = (
caller_failure_model_info.get("id") if caller_failure_model_info is not None else None
)
caller_failure_deployment: Final = (
self.get_deployment(model_id=caller_failure_model_id)
if isinstance(caller_failure_model_id, str)
else None
)
if is_caller_scoped_auth_failure(caller_failure_deployment, exception_status):
verbose_router_logger.debug(
"Router: Exiting 'deployment_callback_on_failure' without cooldown. "
"Caller-scoped OAuth authentication failed, not the deployment."
)
return False
exception_headers: Final = litellm.litellm_core_utils.exception_mapping_utils._get_response_headers(
original_exception=exception
)

View file

@ -21,7 +21,7 @@ clientside_credential_keys: Final = ["api_key", "api_base", "base_url"]
# mint a federation token there even when WIF is configured only through ANTHROPIC_* env vars (which
# cannot be cleared from litellm_params).
DISABLE_WORKLOAD_IDENTITY_PARAM: Final = "anthropic_disable_workload_identity_federation"
_WIF_CLEAR_ON_BASE_OVERRIDE: Final = tuple(sorted(server_owned_wif_litellm_params))
_SERVER_OWNED_IDENTITY_CLEAR_ON_BASE_OVERRIDE: Final = tuple(sorted(server_owned_wif_litellm_params))
def _admin_config_fields_to_clear_on_base_override() -> list[str]:
@ -67,14 +67,13 @@ def _admin_config_fields_to_clear_on_base_override() -> list[str]:
# ``api_base`` for the same reason as the OCI entries above.
"nvcf_function_id",
"use_ssl",
# Workload-identity federation minting fields, restated here from
# Server-owned federation and OAuth token-exchange fields, restated here from
# server_owned_wif_litellm_params the same way azure_ad_token above is restated
# despite also being declared on CredentialLiteLLMParams (hence covered by
# typed_fields too): a federation token minted for a client-redirected api_base
# would send the workload's OIDC assertion, and then the minted bearer, to the
# caller-chosen host, so this list must stay correct even if a field is ever
# dropped from the typed model.
*_WIF_CLEAR_ON_BASE_OVERRIDE,
# typed_fields too): tokens minted for a client-redirected api_base could send an
# assertion and bearer to the caller-chosen host, so this list must stay correct even
# if a field is ever dropped from the typed model.
*_SERVER_OWNED_IDENTITY_CLEAR_ON_BASE_OVERRIDE,
]
return typed_fields + kwargs_only_fields

View file

@ -13,6 +13,8 @@ from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
from pydantic import TypeAdapter
import litellm
from litellm._internal_context import service_target
from litellm._logging import verbose_router_logger
@ -24,8 +26,10 @@ from litellm.constants import (
INTERNAL_CALL_ORIGIN_METADATA_KEY,
SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD,
)
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.router_utils.cooldown_cache import ROUTER_COOLDOWNS_TARGET, CooldownCacheValue
from litellm.router_utils.cooldown_callbacks import router_cooldown_event_callback
from litellm.types.router import Deployment
from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN
from .router_callbacks.track_deployment_metrics import (
@ -45,6 +49,7 @@ else:
Span = Any
_ADVISOR_ORCHESTRATION_FAILURE_ATTR: Final = "_litellm_advisor_orchestration_failure"
_CREDENTIAL_VALUES_ADAPTER: Final = TypeAdapter(Mapping[str, object])
def mark_advisor_orchestration_failure(exception: BaseException) -> None:
@ -688,3 +693,22 @@ def is_caller_timeout_408(
if not isinstance(timeout, (int, float)) or not isinstance(started, datetime) or not isinstance(finished, datetime):
return False
return (finished - started).total_seconds() >= timeout
def is_caller_scoped_auth_failure(deployment: Deployment | None, exception_status: str | int) -> bool:
if cast_exception_status_to_int(exception_status) not in (401, 403) or deployment is None:
return False
token_exchange_endpoint: Final = deployment.litellm_params.token_exchange_endpoint
if isinstance(token_exchange_endpoint, str) and token_exchange_endpoint:
return True
credential_name: Final = deployment.litellm_params.litellm_credential_name
if credential_name is None:
return False
credential_values: Final[Mapping[str, object]] = _CREDENTIAL_VALUES_ADAPTER.validate_python(
CredentialAccessor.get_credential_values(credential_name)
)
credential_token_exchange_endpoint: Final = credential_values.get("token_exchange_endpoint")
return isinstance(credential_token_exchange_endpoint, str) and bool(credential_token_exchange_endpoint)

View file

@ -22,6 +22,7 @@ from litellm.router_utils.cooldown_handlers import (
_first_present, # pyright: ignore[reportPrivateUsage] - shared internal helper, used across router_utils
cast_exception_status_to_int,
is_advisor_orchestration_failure,
is_caller_scoped_auth_failure,
is_caller_timeout_408,
set_cooldown_deployments,
)
@ -108,6 +109,14 @@ def _trigger_cooldown_for_failed_deployment(
verbose_router_logger.debug("Cannot trigger cooldown for fallback: no failed_deployment_id on exception")
return
deployment: Final = litellm_router.get_deployment(model_id=deployment_id)
if is_caller_scoped_auth_failure(deployment, exception_status):
verbose_router_logger.debug(
"Not triggering cooldown for fallback deployment %s: caller-scoped OAuth authentication failure.",
deployment_id,
)
return
# Priority: deployment config > response header > router default, matching
# Router.deployment_callback_on_failure's precedence for the primary path.
deployment_dict: Final = litellm_router.get_model_info(id=deployment_id)

View file

@ -87,6 +87,10 @@ class ProviderConnection:
litellm_credential_name: str | None = None
configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None
use_xai_oauth: bool | None = None
token_exchange_endpoint: str | None = None
token_exchange_profile: str | None = None
token_exchange_scope: str | None = None
token_exchange_audience: str | None = None
fireworks_forward_user_id: bool | None = None

View file

@ -338,6 +338,10 @@ class ModelInfo(MirroredPricingParams):
class CredentialLiteLLMParams(LiteLLMBaseModel):
if TYPE_CHECKING:
def __init__(self, /, **data: object) -> None: ... # kwargs-ok: credential fields vary by provider
api_key: str | None = None
api_base: str | None = None
api_version: str | None = None
@ -421,28 +425,32 @@ class CredentialLiteLLMParams(LiteLLMBaseModel):
openai_identity_provider_id: str | None = None
openai_service_account_id: str | None = None
openai_identity_token_file: str | None = None
token_exchange_endpoint: str | None = None
token_exchange_profile: str | None = None
token_exchange_scope: str | None = None
token_exchange_audience: str | None = None
def server_owned_wif_fields_present(fields: Mapping[str, object]) -> tuple[str, ...]:
"""Server-owned workload identity federation field names set in ``fields``.
"""Server-owned federation or OAuth token-exchange field names set in ``fields``.
``fields`` is a ``litellm_params`` dict (or a credential's ``credential_values`` mapping,
which feeds the same resolution when referenced by name). Derived from
``server_owned_wif_litellm_params`` rather than hand-copied, so a persistence gate built on
this stays correct when a new WIF field is added there.
this stays correct when a server-owned field is added there.
"""
return tuple(name for name in _server_owned_wif_litellm_params if fields.get(name) is not None)
def server_owned_wif_fields_named(keys: Container[str]) -> tuple[str, ...]:
"""Server-owned workload identity federation field names that appear in ``keys``, whatever
"""Server-owned federation or OAuth token-exchange field names appearing in ``keys``, whatever
value they carry.
The write gates on credentials need this key-based sibling of ``server_owned_wif_fields_present``:
``get_litellm_params`` forwards a WIF kwarg on key presence and the federation resolver rejects
a foreign variant's field by key, so a persisted ``{"anthropic_issuer_url": None}`` wedges every
deployment that references the credential even though no value is set. Pass a mapping (its keys
are tested) or a plain collection of key names.
``get_litellm_params`` forwards a server-owned field on key presence, so a persisted
``{"anthropic_issuer_url": None}`` can wedge every deployment that references the credential
even though no value is set. Pass a mapping (its keys are tested) or a plain collection of key
names.
"""
return tuple(name for name in _server_owned_wif_litellm_params if name in keys)
@ -464,6 +472,10 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
LiteLLM Params without 'model' arg (used across completion / assistants api)
"""
if TYPE_CHECKING:
def __init__(self, /, **data: object) -> None: ... # kwargs-ok: provider parameters vary by deployment
custom_llm_provider: str | None = None
tpm: int | None = None
rpm: int | None = None
@ -609,6 +621,10 @@ class LiteLLM_Params(GenericLiteLLMParams):
LiteLLM Params with 'model' requirement - used for completions
"""
if TYPE_CHECKING:
def __init__(self, *, model: str, **data: object) -> None: ... # kwargs-ok: provider fields vary by model
model: str
model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True)
@ -632,6 +648,10 @@ class LiteLLM_Params(GenericLiteLLMParams):
class updateLiteLLMParams(GenericLiteLLMParams):
# This class is used to update the LiteLLM_Params
# only differece is model is optional
if TYPE_CHECKING:
def __init__(self, *, model: str | None = None, **data: object) -> None: ... # kwargs-ok: per-provider
model: str | None = None
@ -1338,18 +1358,20 @@ class AdaptiveRouterPreferences(LiteLLMBaseModel):
def reject_server_owned_wif_params(body: Mapping[str, object]) -> None:
"""Raise ``ValueError`` if a mapping that did not come from deployment config carries a
server-owned workload identity federation field.
server-owned workload identity federation or OAuth token-exchange field.
These are never settable inline on a client surface, with or without a client-side credential
opt-in. Naming a stored credential that already holds them is the other way in and has its own
gate: ``_check_banned_params`` resolves ``litellm_credential_name`` and refuses a federated one.
gate: ``_check_banned_params`` resolves ``litellm_credential_name`` and refuses a credential
carrying one of these server-owned fields.
This lives here rather than under ``litellm.proxy`` so the router can call it on a
post-authentication merge without core importing from the proxy package.
"""
for param in _server_owned_wif_litellm_params:
if param in body:
raise ValueError(
f"Rejected Request: {param} is a server-owned workload identity federation parameter "
f"Rejected Request: {param} is a server-owned workload identity federation or OAuth token exchange "
"parameter "
"and cannot be set in a request body. A proxy admin configures it on the deployment "
"or on a stored credential."
)

View file

@ -4054,7 +4054,15 @@ ADDRESSED_RESPONSE_ID_FIELD: Final = _litellm_params.ADDRESSED_RESPONSE_ID_FIELD
anthropic_wif_litellm_params: Final = tuple(sorted(ANTHROPIC_WIF_KWARGS_KEYS))
openai_wif_litellm_params: Final = tuple(sorted(OPENAI_WIF_KWARGS_KEYS))
server_owned_wif_litellm_params: Final = anthropic_wif_litellm_params + openai_wif_litellm_params
oauth_token_exchange_litellm_params: Final = (
"token_exchange_audience",
"token_exchange_endpoint",
"token_exchange_profile",
"token_exchange_scope",
)
server_owned_wif_litellm_params: Final = (
anthropic_wif_litellm_params + openai_wif_litellm_params + oauth_token_exchange_litellm_params
)
secret_bearing_wif_litellm_params: Final = tuple(sorted(WIF_SECRET_BEARING_KEYS))
all_litellm_params = [ # rebind-ok: two star imports in litellm/__init__.py re-bind it
@ -4250,6 +4258,7 @@ class LlmProviders(str, Enum):
A2A_AGENT = "a2a_agent"
LANGGRAPH = "langgraph"
LANGFLOW = "langflow"
MICROSOFT_365_COPILOT = "microsoft_365_copilot"
MINIMAX = "minimax"
SYNTHETIC = "synthetic"
APERTIS = "apertis"

View file

@ -8671,6 +8671,10 @@ class ProviderConfigManager:
lambda: ProviderConfigManager._get_langgraph_config(),
False,
),
LlmProviders.MICROSOFT_365_COPILOT: (
lambda: ProviderConfigManager._get_microsoft_365_copilot_config(),
False,
),
LlmProviders.SAIL: (ProviderConfigManager._get_sail_chat_config, False),
LlmProviders.LANGFLOW: (
lambda: ProviderConfigManager._get_langflow_config(),
@ -8765,6 +8769,12 @@ class ProviderConfigManager:
return LangGraphConfig()
@staticmethod
def _get_microsoft_365_copilot_config() -> BaseConfig:
from litellm.llms.microsoft_365_copilot.chat.transformation import Microsoft365CopilotChatConfig
return Microsoft365CopilotChatConfig()
@staticmethod
def _get_langflow_config() -> BaseConfig:
"""Get LangFlow config."""

View file

@ -78859,6 +78859,13 @@
"prompt_cache_min_tokens": 512,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing"
},
"microsoft_365_copilot/chat": {
"litellm_provider": "microsoft_365_copilot",
"mode": "chat",
"supports_mid_conversation_system": true,
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0
},
"openrouter/inclusionai/ling-3.0-flash-sante": {
"cache_read_input_token_cost": 8.4e-09,
"input_cost_per_token": 4.2e-08,

View file

@ -2871,6 +2871,24 @@
"interactions": true
}
},
"microsoft_365_copilot": {
"display_name": "Microsoft 365 Copilot (`microsoft_365_copilot`)",
"url": "https://docs.litellm.ai/docs/providers/microsoft_365_copilot",
"endpoints": {
"chat_completions": true,
"messages": false,
"responses": false,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false,
"interactions": false
}
},
"langflow": {
"display_name": "LangFlow (`langflow`)",
"url": "https://docs.litellm.ai/docs/providers/langflow",

View file

@ -10,6 +10,7 @@ import hashlib
import random
import pytest
from pydantic import TypeAdapter
import litellm
from litellm import aembedding, completion, embedding, aresponses, responses
@ -39,9 +40,9 @@ from litellm.types.utils import (
Embedding,
)
from litellm.types.llms.openai import ResponsesAPIResponse
from collections.abc import Awaitable, Callable
from collections.abc import Awaitable, Callable, Mapping
from datetime import timedelta, datetime
from typing import Final
from typing import Final, cast
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm._logging import verbose_logger
@ -52,6 +53,7 @@ import respx
from fastapi.testclient import TestClient
from litellm._internal_context import current_service_target, in_post_response_phase
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
def setup_cache():
@ -61,6 +63,143 @@ def setup_cache():
return cache
@pytest.fixture
def copilot_response_cache(monkeypatch: pytest.MonkeyPatch) -> Cache:
cache: Final = Cache(type=LiteLLMCacheType.LOCAL)
monkeypatch.setattr(litellm, "cache", cache)
return cache
class _CopilotCacheAsyncTransport(httpx.AsyncBaseTransport):
def __init__(self, responder: Callable[[httpx.Request], httpx.Response]) -> None:
self._responder = responder
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
return self._responder(request)
def _copilot_cache_responder(requests: list[httpx.Request]) -> Callable[[httpx.Request], httpx.Response]:
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
if request.url.path == "/beta/copilot/conversations":
return httpx.Response(status_code=201, json={"id": "cache-conversation"}, request=request)
if request.url.path.endswith("/chat"):
return httpx.Response(
status_code=200,
json={"messages": [{"text": "cache reply"}]},
request=request,
)
return httpx.Response(status_code=404, json={}, request=request)
return respond
def _sync_copilot_cache_client(
requests: list[httpx.Request],
) -> HTTPHandler:
return HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_copilot_cache_responder(requests))))
def _async_copilot_cache_client(
requests: list[httpx.Request],
) -> AsyncHTTPHandler:
return AsyncHTTPHandler(transport=_CopilotCacheAsyncTransport(_copilot_cache_responder(requests)))
def _assert_copilot_cache_empty(cache: Cache) -> None:
cache_backend: Final = cache.cache
assert isinstance(cache_backend, InMemoryCache)
cache_contents: Final[Mapping[str, object]] = cast( # cast-ok: in-memory cache exposes an untyped mapping
Mapping[str, object],
cache_backend.cache_dict,
)
assert cache_contents == {}
def _sync_copilot_cache_call(client: HTTPHandler, stream: bool) -> None:
response: Final = litellm.completion(
model="microsoft_365_copilot/chat",
messages=[{"role": "user", "content": "cache isolation prompt"}],
api_key="delegated-cache-token",
client=client,
stream=stream,
caching=True,
)
if stream:
assert isinstance(response, CustomStreamWrapper)
tuple(response)
async def _async_copilot_cache_call(client: AsyncHTTPHandler, stream: bool) -> None:
response: Final = await litellm.acompletion(
model="microsoft_365_copilot/chat",
messages=[{"role": "user", "content": "cache isolation prompt"}],
api_key="delegated-cache-token",
client=client,
stream=stream,
caching=True,
)
if stream:
assert isinstance(response, CustomStreamWrapper)
tuple([chunk async for chunk in response])
@pytest.mark.parametrize("stream", [False, True])
def test_sync_copilot_requests_bypass_response_cache(copilot_response_cache: Cache, stream: bool) -> None:
requests: Final[list[httpx.Request]] = []
client: Final = _sync_copilot_cache_client(requests)
for _ in range(2):
_sync_copilot_cache_call(client=client, stream=stream)
graph_chat_requests: Final = tuple(request for request in requests if request.url.path.endswith("/chat"))
assert len(graph_chat_requests) == 2
_assert_copilot_cache_empty(copilot_response_cache)
@pytest.mark.asyncio
@pytest.mark.parametrize("stream", [False, True])
async def test_async_copilot_requests_bypass_response_cache(copilot_response_cache: Cache, stream: bool) -> None:
requests: Final[list[httpx.Request]] = []
client: Final = _async_copilot_cache_client(requests)
async with client.client:
for _ in range(2):
await _async_copilot_cache_call(client=client, stream=stream)
await asyncio.gather(*_PENDING_CACHE_WRITES)
graph_chat_requests: Final = tuple(request for request in requests if request.url.path.endswith("/chat"))
assert len(graph_chat_requests) == 2
_assert_copilot_cache_empty(copilot_response_cache)
@pytest.mark.asyncio
async def test_existing_provider_response_cache_still_hits(copilot_response_cache: Cache) -> None:
messages: Final = [{"role": "user", "content": "cached provider control"}]
first: Final = await litellm.acompletion(
model="gpt-4o",
messages=messages,
mock_response="first cached response",
caching=True,
)
await asyncio.gather(*_PENDING_CACHE_WRITES)
second: Final = await litellm.acompletion(
model="gpt-4o",
messages=messages,
mock_response="second response should not be used",
caching=True,
)
assert isinstance(first, ModelResponse)
assert isinstance(second, ModelResponse)
hidden_params_value: Final[object] = cast( # cast-ok: validate the response's dynamic metadata
object,
second.hidden_params,
)
hidden_params: Final[Mapping[str, object]] = TypeAdapter(Mapping[str, object]).validate_python(hidden_params_value)
assert second.choices[0].message.content == first.choices[0].message.content
assert hidden_params.get("cache_hit") is True
chat_completion_response = litellm.ModelResponse(
id=str(uuid.uuid4()),
choices=[

View file

@ -0,0 +1,289 @@
from collections.abc import Callable
from dataclasses import replace
from typing import Final
from urllib.parse import parse_qs
import httpx
import pytest
from litellm.litellm_core_utils import oauth_token_exchange
from litellm.litellm_core_utils.oauth_token_exchange import (
OAuthTokenExchangeConfig,
OAuthTokenExchangeError,
TokenExchangeProfile,
aexchange_token,
build_token_exchange_form,
exchange_token,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
class _AsyncCaptureTransport(httpx.AsyncBaseTransport):
def __init__(self, responder: Callable[[httpx.Request], httpx.Response]) -> None:
self._responder = responder
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
return self._responder(request)
def _config(
*,
token_endpoint: str = "https://identity.example.com/oauth2/token",
client_id: str = "client-9306",
client_secret: str = "secret-9306",
profile: TokenExchangeProfile = "rfc8693",
scope: str | None = "graph.scope",
audience: str | None = "graph-api",
) -> OAuthTokenExchangeConfig:
return OAuthTokenExchangeConfig(
token_endpoint=token_endpoint,
client_id=client_id,
client_secret=client_secret,
profile=profile,
scope=scope,
audience=audience,
)
def _sync_client(
status_code: int,
payload: object,
) -> tuple[HTTPHandler, list[httpx.Request]]:
requests: Final[list[httpx.Request]] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(status_code=status_code, json=payload, request=request)
return HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond))), requests
def _decode_form(request: httpx.Request) -> dict[str, str]:
return {name: values[0] for name, values in parse_qs(request.content.decode("utf-8")).items()}
def test_rfc8693_form_contains_exact_fields_and_optional_scope_and_audience() -> None:
config: Final = _config()
form: Final = build_token_exchange_form(config, "subject-token-9306")
assert form == {
"grant_type": "urn:ietf:params:oauth:grant-type:token-exchange",
"client_id": "client-9306",
"client_secret": "secret-9306",
"subject_token": "subject-token-9306",
"subject_token_type": "urn:ietf:params:oauth:token-type:access_token",
"scope": "graph.scope",
"audience": "graph-api",
}
@pytest.mark.parametrize(
("scope", "audience", "expected_optional_fields"),
[
(None, None, {}),
("", "", {}),
("graph.scope", None, {"scope": "graph.scope"}),
(None, "graph-api", {"audience": "graph-api"}),
],
)
def test_rfc8693_form_omits_empty_or_unset_optional_fields(
scope: str | None,
audience: str | None,
expected_optional_fields: dict[str, str],
) -> None:
config: Final = _config(scope=scope, audience=audience)
form: Final = build_token_exchange_form(config, "subject-token-9306")
assert form == {
"grant_type": "urn:ietf:params:oauth:grant-type:token-exchange",
"client_id": "client-9306",
"client_secret": "secret-9306",
"subject_token": "subject-token-9306",
"subject_token_type": "urn:ietf:params:oauth:token-type:access_token",
**expected_optional_fields,
}
def test_jwt_bearer_obo_form_requires_scope_and_ignores_audience() -> None:
config: Final = _config(profile="jwt_bearer_obo", audience="ignored-audience")
form: Final = build_token_exchange_form(config, "subject-token-9306")
assert form == {
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
"client_id": "client-9306",
"client_secret": "secret-9306",
"assertion": "subject-token-9306",
"scope": "graph.scope",
"requested_token_use": "on_behalf_of",
}
@pytest.mark.parametrize("scope", [None, "", " "])
def test_jwt_bearer_obo_rejects_empty_scope(scope: str | None) -> None:
with pytest.raises(OAuthTokenExchangeError) as error:
_config(profile="jwt_bearer_obo", scope=scope)
assert error.value.status_code == 400
assert error.value.message == "jwt_bearer_obo requires a non-empty token exchange scope"
@pytest.mark.parametrize(
"endpoint",
[
"http://identity.example.com/token",
"https:///token",
"https://identity.example.com:invalid/token",
"not-a-url",
],
)
def test_config_rejects_invalid_token_endpoint(endpoint: str) -> None:
with pytest.raises(OAuthTokenExchangeError) as error:
_config(token_endpoint=endpoint)
assert error.value.status_code == 400
assert error.value.message == "token_exchange_endpoint must be an HTTPS URL with a host"
def test_cache_separates_each_config_field_and_subject_token() -> None:
base_config: Final = _config(
token_endpoint="https://identity.example.com/cache-separation-9306",
client_id="cache-client-9306",
client_secret="cache-secret-9306",
profile="rfc8693",
scope="cache-scope-9306",
audience="cache-audience-9306",
)
config_and_subjects: Final = (
(base_config, "cache-subject-9306"),
(replace(base_config, token_endpoint="https://identity.example.com/other"), "cache-subject-9306"),
(replace(base_config, client_id="other-client"), "cache-subject-9306"),
(replace(base_config, client_secret="other-secret"), "cache-subject-9306"),
(replace(base_config, profile="jwt_bearer_obo"), "cache-subject-9306"),
(replace(base_config, scope="other-scope"), "cache-subject-9306"),
(replace(base_config, audience="other-audience"), "cache-subject-9306"),
(base_config, "other-subject"),
)
response_count: Final[list[int]] = []
def respond(request: httpx.Request) -> httpx.Response:
response_count.append(len(response_count) + 1)
return httpx.Response(
status_code=200,
json={"access_token": f"cached-token-{len(response_count)}", "expires_in": 3600},
request=request,
)
client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond)))
returned_tokens: Final = tuple(
exchange_token(client, config, subject_token) for config, subject_token in config_and_subjects
)
repeated_tokens: Final = tuple(
exchange_token(client, config, subject_token) for config, subject_token in config_and_subjects
)
assert len(response_count) == len(config_and_subjects)
assert returned_tokens == repeated_tokens
def test_exchange_applies_cache_safety_margin() -> None:
cache_margin: Final = oauth_token_exchange.OAUTH_TOKEN_EXCHANGE_CACHE_SAFETY_MARGIN_SECONDS
no_cache_client, no_cache_requests = _sync_client(
200,
{"access_token": "ttl-boundary-token-9306", "expires_in": cache_margin},
)
no_cache_config: Final = _config(token_endpoint="https://identity.example.com/ttl-boundary-9306")
cache_client, cache_requests = _sync_client(
200,
{"access_token": "ttl-cached-token-9306", "expires_in": cache_margin + 60},
)
cache_config: Final = _config(token_endpoint="https://identity.example.com/ttl-cached-9306")
exchange_token(no_cache_client, no_cache_config, "ttl-boundary-subject-9306")
exchange_token(no_cache_client, no_cache_config, "ttl-boundary-subject-9306")
exchange_token(cache_client, cache_config, "ttl-cached-subject-9306")
exchange_token(cache_client, cache_config, "ttl-cached-subject-9306")
assert len(no_cache_requests) == 2
assert len(cache_requests) == 1
@pytest.mark.parametrize(
("status_code", "expected_status"),
[(400, 401), (429, 429), (503, 503)],
)
def test_sync_exchange_maps_error_status_and_redacts_secrets(status_code: int, expected_status: int) -> None:
subject_token: Final = "header.secret.subject"
client_secret: Final = "secret-error-9306"
client, _ = _sync_client(
status_code,
{
"error": "invalid_grant",
"error_description": f"subject={subject_token}, secret={client_secret}",
},
)
with pytest.raises(OAuthTokenExchangeError) as error:
exchange_token(client, _config(client_secret=client_secret), subject_token)
assert error.value.status_code == expected_status
assert error.value.message == "OAuth token exchange failed: invalid_grant: subject=[REDACTED], secret=[REDACTED]"
assert client_secret not in error.value.message
assert subject_token not in error.value.message
def test_sync_exchange_posts_to_endpoint_with_exact_form() -> None:
config: Final = _config(token_endpoint="https://identity.example.com/sync-9306")
client, requests = _sync_client(200, {"access_token": "sync-token-9306", "expires_in": 3600})
token: Final = exchange_token(client, config, "sync-subject-9306")
assert str(requests[0].url) == config.token_endpoint
assert _decode_form(requests[0]) == build_token_exchange_form(config, "sync-subject-9306")
assert token == "sync-token-9306"
@pytest.mark.asyncio
async def test_async_exchange_posts_to_endpoint_with_exact_form() -> None:
config: Final = _config(token_endpoint="https://identity.example.com/async-9306")
requests: Final[list[httpx.Request]] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(
status_code=200,
json={"access_token": "async-token-9306", "expires_in": 3600},
request=request,
)
client: Final = AsyncHTTPHandler(transport=_AsyncCaptureTransport(respond))
async with client.client:
token: Final = await aexchange_token(client, config, "async-subject-9306")
assert str(requests[0].url) == config.token_endpoint
assert _decode_form(requests[0]) == build_token_exchange_form(config, "async-subject-9306")
assert token == "async-token-9306"
@pytest.mark.asyncio
@pytest.mark.parametrize(
("status_code", "expected_status"),
[(400, 401), (429, 429), (503, 503)],
)
async def test_async_exchange_maps_error_status(status_code: int, expected_status: int) -> None:
client: Final = AsyncHTTPHandler(
transport=_AsyncCaptureTransport(
lambda request: httpx.Response(
status_code=status_code,
json={"error": "invalid_grant", "error_description": "request failed"},
request=request,
)
)
)
async with client.client:
with pytest.raises(OAuthTokenExchangeError) as error:
await aexchange_token(client, _config(), "async-status-subject-9306")
assert error.value.status_code == expected_status

View file

@ -3761,7 +3761,10 @@ class TestWifServerOwnedParamsAreUnconditional:
router.py merges request kwargs OVER deployment params, so it also beat the configured one."""
from litellm.proxy.auth.auth_utils import is_request_body_safe
with pytest.raises(Exception, match="server-owned workload identity federation parameter"):
with pytest.raises(
Exception,
match="server-owned workload identity federation or OAuth token exchange parameter",
):
is_request_body_safe(
request_body={"model": "claude-sonnet-5", "anthropic_federation_workspace_id": "wrkspc_abc"},
general_settings={},

View file

@ -0,0 +1,677 @@
import json
from collections.abc import Callable, Mapping
from copy import deepcopy
from typing import Final, cast
from urllib.parse import parse_qs
import httpx
import pytest
from pydantic import TypeAdapter
import litellm
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.microsoft_365_copilot.chat.handler import acompletion, completion
from litellm.llms.microsoft_365_copilot.common_utils import Microsoft365CopilotError
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
from litellm.types.utils import ModelResponse
from litellm.utils import token_counter
_MODEL: Final = "microsoft_365_copilot/chat"
_MESSAGES_ADAPTER: Final = TypeAdapter(list[AllMessageValues])
_DEFAULT_CHAT_RESPONSE: Final = {"messages": [{"text": "prompt echo"}, {"text": "Copilot response"}]}
_DEFAULT_OAUTH_RESPONSE: Final = {"access_token": "graph-access-token", "expires_in": 3600}
class _AsyncCaptureTransport(httpx.AsyncBaseTransport):
def __init__(self, responder: Callable[[httpx.Request], httpx.Response]) -> None:
self._responder = responder
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
return self._responder(request)
def _secret_fields(assertion: str) -> SecretFields:
return SecretFields(raw_headers={"Authorization": f"bEaReR {assertion}"})
def _messages() -> list[AllMessageValues]:
return _MESSAGES_ADAPTER.validate_python([{"role": "user", "content": "hello Copilot"}])
def _sync_client(
oauth_status: int = 200,
oauth_payload: object = _DEFAULT_OAUTH_RESPONSE,
graph_status: int = 200,
graph_payload: object = _DEFAULT_CHAT_RESPONSE,
conversation_id: str = "conversation/42",
) -> tuple[HTTPHandler, list[httpx.Request]]:
requests: Final[list[httpx.Request]] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
if request.url.host == "identity.example.com":
return httpx.Response(
status_code=oauth_status,
json=oauth_payload,
request=request,
)
if request.url.path == "/beta/copilot/conversations":
return httpx.Response(
status_code=201,
json={"id": conversation_id, "state": "active"},
request=request,
)
if request.url.path.endswith("/chat"):
return httpx.Response(
status_code=graph_status,
json=graph_payload,
request=request,
)
return httpx.Response(status_code=404, json={}, request=request)
client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond)))
return client, requests
def _async_client(
oauth_status: int = 200,
oauth_payload: object = _DEFAULT_OAUTH_RESPONSE,
graph_status: int = 200,
graph_payload: object = _DEFAULT_CHAT_RESPONSE,
conversation_id: str = "conversation/42",
) -> tuple[AsyncHTTPHandler, list[httpx.Request]]:
requests: Final[list[httpx.Request]] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
if request.url.host == "identity.example.com":
return httpx.Response(
status_code=oauth_status,
json=oauth_payload,
request=request,
)
if request.url.path == "/beta/copilot/conversations":
return httpx.Response(
status_code=201,
json={"id": conversation_id, "state": "active"},
request=request,
)
if request.url.path.endswith("/chat"):
return httpx.Response(
status_code=graph_status,
json=graph_payload,
request=request,
)
return httpx.Response(status_code=404, json={}, request=request)
client: Final = AsyncHTTPHandler(transport=_AsyncCaptureTransport(respond))
return client, requests
def _token_exchange_params(
exchange_name: str,
client_id: str,
client_secret: str,
) -> Mapping[str, object]:
return {
"token_exchange_endpoint": f"https://identity.example.com/{exchange_name}/oauth2/v2.0/token",
"client_id": client_id,
"client_secret": client_secret,
"token_exchange_profile": "jwt_bearer_obo",
"token_exchange_scope": "https://graph.microsoft.com/.default",
}
def _token_requests(requests: list[httpx.Request]) -> tuple[httpx.Request, ...]:
return tuple(request for request in requests if request.url.host == "identity.example.com")
def test_sync_jwt_bearer_obo_uses_exact_exchange_and_sends_graph_token() -> None:
assertion: Final = "header-sync-9306.payload-sync-9306.signature-sync-9306"
params: Final = _token_exchange_params(
"sync-handler-9306",
"client-sync-handler-9306",
"secret-sync-handler-9306",
)
client, requests = _sync_client()
response: Final = completion(
model=_MODEL,
messages=_messages(),
api_key=None,
litellm_params=params,
optional_params={},
secret_fields=_secret_fields(assertion),
client=client,
)
token_request: Final = _token_requests(requests)[0]
token_form: Final = {name: values[0] for name, values in parse_qs(token_request.content.decode("utf-8")).items()}
graph_requests: Final = tuple(request for request in requests if request.url.host == "graph.microsoft.com")
assert str(token_request.url) == "https://identity.example.com/sync-handler-9306/oauth2/v2.0/token"
assert token_form == {
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
"client_id": "client-sync-handler-9306",
"client_secret": "secret-sync-handler-9306",
"assertion": assertion,
"scope": "https://graph.microsoft.com/.default",
"requested_token_use": "on_behalf_of",
}
assert graph_requests[0].method == "POST"
assert graph_requests[0].url.path == "/beta/copilot/conversations"
assert json.loads(graph_requests[0].content) == {}
assert graph_requests[0].headers["Authorization"] == "Bearer graph-access-token"
assert graph_requests[1].url.raw_path == b"/beta/copilot/conversations/conversation%2F42/chat"
assert json.loads(graph_requests[1].content) == {
"message": {"text": "hello Copilot"},
"locationHint": {"timeZone": "UTC"},
}
assert graph_requests[1].headers["Authorization"] == "Bearer graph-access-token"
assert assertion not in graph_requests[0].headers["Authorization"]
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "Copilot response"
assert response.model == _MODEL
@pytest.mark.asyncio
async def test_async_jwt_bearer_obo_uses_exact_exchange_and_sends_graph_token() -> None:
assertion: Final = "header-async-9306.payload-async-9306.signature-async-9306"
params: Final = _token_exchange_params(
"async-handler-9306",
"client-async-handler-9306",
"secret-async-handler-9306",
)
client, requests = _async_client()
async with client.client:
response: Final = await acompletion(
model=_MODEL,
messages=_messages(),
api_key=None,
litellm_params=params,
optional_params={},
secret_fields=_secret_fields(assertion),
client=client,
)
token_request: Final = _token_requests(requests)[0]
token_form: Final = {name: values[0] for name, values in parse_qs(token_request.content.decode("utf-8")).items()}
graph_requests: Final = tuple(request for request in requests if request.url.host == "graph.microsoft.com")
assert str(token_request.url) == "https://identity.example.com/async-handler-9306/oauth2/v2.0/token"
assert token_form == {
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
"client_id": "client-async-handler-9306",
"client_secret": "secret-async-handler-9306",
"assertion": assertion,
"scope": "https://graph.microsoft.com/.default",
"requested_token_use": "on_behalf_of",
}
assert graph_requests[0].headers["Authorization"] == "Bearer graph-access-token"
assert graph_requests[1].url.raw_path == b"/beta/copilot/conversations/conversation%2F42/chat"
assert graph_requests[1].headers["Authorization"] == "Bearer graph-access-token"
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "Copilot response"
def test_token_exchange_cache_hits_and_keys_by_subject_token_and_client_secret() -> None:
client, requests = _sync_client()
exchange_name: Final = "cache-handler-9306"
client_id: Final = "client-cache-handler-9306"
client_secret: Final = "secret-cache-handler-original-9306"
first_assertion: Final = "header-cache-one.payload-cache-one.signature-cache-one"
second_assertion: Final = "header-cache-two.payload-cache-two.signature-cache-two"
def call(assertion: str, secret: str) -> None:
completion(
model=_MODEL,
messages=_messages(),
api_key=None,
litellm_params=_token_exchange_params(exchange_name, client_id, secret),
optional_params={},
secret_fields=_secret_fields(assertion),
client=client,
)
call(first_assertion, client_secret)
call(first_assertion, client_secret)
call(second_assertion, client_secret)
call(second_assertion, "secret-cache-handler-changed-9306")
assert len(_token_requests(requests)) == 3
def test_direct_delegated_graph_token_skips_token_exchange() -> None:
client, requests = _sync_client()
response: Final = completion(
model=_MODEL,
messages=_messages(),
api_key="delegated-graph-token-9306",
litellm_params={},
optional_params={},
client=client,
)
graph_requests: Final = tuple(request for request in requests if request.url.host == "graph.microsoft.com")
assert len(_token_requests(requests)) == 0
assert len(graph_requests) == 2
assert all(request.headers["Authorization"] == "Bearer delegated-graph-token-9306" for request in graph_requests)
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "Copilot response"
def test_token_exchange_credentials_take_precedence_over_direct_api_key() -> None:
assertion: Final = "header-priority.payload-priority.signature-priority"
client, requests = _sync_client()
messages: Final = _messages()
litellm_params: Final = _token_exchange_params(
"priority-9306",
"client-priority-9306",
"secret-priority-9306",
)
optional_params: Final = {"time_zone": "UTC"}
secret_fields: Final = _secret_fields(assertion)
original_messages: Final = deepcopy(messages)
original_litellm_params: Final = dict(litellm_params)
original_optional_params: Final = dict(optional_params)
original_headers: Final = dict(cast(Mapping[str, str], secret_fields["raw_headers"]))
completion(
model=_MODEL,
messages=messages,
api_key="direct-token-must-not-win-9306",
litellm_params=litellm_params,
optional_params=optional_params,
secret_fields=secret_fields,
client=client,
)
graph_requests: Final = tuple(request for request in requests if request.url.host == "graph.microsoft.com")
assert len(_token_requests(requests)) == 1
assert len(graph_requests) == 2
assert all(request.headers["Authorization"] == "Bearer graph-access-token" for request in graph_requests)
assert messages == original_messages
assert litellm_params == original_litellm_params
assert optional_params == original_optional_params
assert secret_fields["raw_headers"] == original_headers
def test_partial_token_exchange_configuration_fails_without_http_calls() -> None:
client, requests = _sync_client()
with pytest.raises(Microsoft365CopilotError) as error:
completion(
model=_MODEL,
messages=_messages(),
api_key="direct-token-must-not-fallback-9306",
litellm_params={
"token_exchange_endpoint": "https://identity.example.com/partial-9306/token",
"client_id": "client-partial-9306",
},
optional_params={},
client=client,
)
assert error.value.status_code == 400
assert len(requests) == 0
def test_api_base_override_clears_exchange_config_and_fails_closed_without_token_post() -> None:
from litellm.router_utils.clientside_credential_handler import get_dynamic_litellm_params
assertion: Final = "header-base-override.payload-base-override.signature-base-override"
oauth_fields: Final = (
"token_exchange_audience",
"token_exchange_endpoint",
"token_exchange_profile",
"token_exchange_scope",
)
litellm_params: Final = get_dynamic_litellm_params(
litellm_params={
"model": _MODEL,
"api_base": "https://graph.microsoft.com/beta",
**_token_exchange_params(
"base-override-9306",
"client-base-override-9306",
"secret-base-override-9306",
),
"token_exchange_audience": "https://graph.microsoft.com",
},
request_kwargs={"api_base": "https://caller-controlled.example"},
)
assert litellm_params["api_base"] == "https://caller-controlled.example"
assert all(field not in litellm_params for field in oauth_fields)
assert "client_id" not in litellm_params
assert "client_secret" not in litellm_params
client, requests = _sync_client()
with pytest.raises(Microsoft365CopilotError) as error:
completion(
model=_MODEL,
messages=_messages(),
api_key=None,
litellm_params=litellm_params,
optional_params={},
secret_fields=_secret_fields(assertion),
client=client,
)
assert error.value.status_code == 401
assert _token_requests(requests) == ()
def test_unknown_token_exchange_profile_fails_without_http_calls() -> None:
client, requests = _sync_client()
litellm_params: Final = {
**_token_exchange_params("unknown-profile-9306", "client-unknown-profile-9306", "secret-unknown-profile-9306"),
"token_exchange_profile": "unknown",
}
with pytest.raises(Microsoft365CopilotError) as error:
completion(
model=_MODEL,
messages=_messages(),
api_key=None,
litellm_params=litellm_params,
optional_params={},
secret_fields=_secret_fields("header.unknown.profile"),
client=client,
)
assert error.value.status_code == 400
assert error.value.message == ("microsoft_365_copilot token_exchange_profile must be 'rfc8693' or 'jwt_bearer_obo'")
assert len(requests) == 0
def test_missing_auth_configuration_fails_without_http_calls() -> None:
client, requests = _sync_client()
with pytest.raises(Microsoft365CopilotError) as error:
completion(
model=_MODEL,
messages=_messages(),
api_key=None,
litellm_params={},
optional_params={},
client=client,
)
assert error.value.status_code == 401
assert len(requests) == 0
@pytest.mark.parametrize(
"secret_fields",
[
None,
SecretFields(raw_headers={"Authorization": "Bearer caller-api-key"}),
],
)
def test_token_exchange_requires_a_caller_bearer_token(secret_fields: SecretFields | None) -> None:
client, requests = _sync_client()
with pytest.raises(Microsoft365CopilotError) as error:
completion(
model=_MODEL,
messages=_messages(),
api_key="direct-token-must-not-fallback-9306",
litellm_params=_token_exchange_params(
"no-assertion-9306",
"client-no-assertion-9306",
"secret-no-assertion-9306",
),
optional_params={},
secret_fields=secret_fields,
client=client,
)
assert error.value.status_code == 401
assert str(error.value) == (
"microsoft_365_copilot with OAuth token exchange requires the caller's IdP-issued access token "
"in the Authorization header"
)
assert len(requests) == 0
def test_token_exchange_invalid_grant_maps_to_401_without_leaking_credentials() -> None:
assertion: Final = "header-invalid-grant.payload-invalid-grant.signature-invalid-grant"
client_secret: Final = "secret-invalid-grant-9306"
client, requests = _sync_client(
oauth_status=400,
oauth_payload={
"error": "invalid_grant",
"error_description": "The assertion is invalid",
},
)
with pytest.raises(Microsoft365CopilotError) as error:
completion(
model=_MODEL,
messages=_messages(),
api_key=None,
litellm_params=_token_exchange_params(
"invalid-grant-9306",
"client-invalid-grant-9306",
client_secret,
),
optional_params={},
secret_fields=_secret_fields(assertion),
client=client,
)
assert error.value.status_code == 401
assert str(error.value) == "OAuth token exchange failed: invalid_grant: The assertion is invalid"
assert assertion not in str(error.value)
assert client_secret not in str(error.value)
assert len(requests) == 1
def test_graph_403_extracts_last_reply_from_stringified_conversation() -> None:
license_message: Final = "It looks like you do not have a valid license"
graph_error: Final = {
"error": {
"code": "UnknownError",
"message": json.dumps({"messages": [{"text": "prompt echo"}, {"text": license_message}]}),
}
}
client, requests = _sync_client(graph_status=403, graph_payload=graph_error)
with pytest.raises(Microsoft365CopilotError) as error:
completion(
model=_MODEL,
messages=_messages(),
api_key="delegated-graph-token-403",
litellm_params={},
optional_params={},
client=client,
)
assert error.value.status_code == 403
assert str(error.value) == license_message
assert len(requests) == 2
def test_graph_error_redacts_caller_token_client_secret_and_access_token() -> None:
assertion: Final = "header-redaction.payload-redaction.signature-redaction"
client_secret: Final = "secret-redaction-9306"
graph_error: Final = {
"error": {
"message": f"token=graph-access-token assertion={assertion} secret={client_secret}",
}
}
client, _ = _sync_client(graph_status=403, graph_payload=graph_error)
with pytest.raises(Microsoft365CopilotError) as error:
completion(
model=_MODEL,
messages=_messages(),
api_key=None,
litellm_params=_token_exchange_params(
"redaction-9306",
"client-redaction-9306",
client_secret,
),
optional_params={},
secret_fields=_secret_fields(assertion),
client=client,
)
assert "graph-access-token" not in str(error.value)
assert assertion not in str(error.value)
assert client_secret not in str(error.value)
def test_streaming_wraps_full_reply_as_fake_stream() -> None:
client, _ = _sync_client()
response: Final = litellm.completion(
model=_MODEL,
messages=_messages(),
api_key="delegated-stream-token-9306",
stream=True,
stream_options={"include_usage": True},
client=client,
)
assert isinstance(response, CustomStreamWrapper)
chunks: Final = tuple(response)
content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
assert content == "Copilot response"
usage_chunk: Final = chunks[-1]
assert usage_chunk.usage.prompt_tokens == token_counter(
model=_MODEL,
messages=_messages(),
)
assert usage_chunk.usage.completion_tokens == token_counter(
model=_MODEL,
text="Copilot response",
count_response_tokens=True,
)
assert usage_chunk.usage.total_tokens == (usage_chunk.usage.prompt_tokens + usage_chunk.usage.completion_tokens)
@pytest.mark.asyncio
async def test_async_streaming_wraps_full_reply_as_fake_stream() -> None:
client, _ = _async_client()
async with client.client:
response: Final = await litellm.acompletion(
model=_MODEL,
messages=_messages(),
api_key="delegated-async-stream-token-9306",
stream=True,
stream_options={"include_usage": True},
client=client,
)
assert isinstance(response, CustomStreamWrapper)
chunks: Final = tuple([chunk async for chunk in response])
contents: Final = tuple(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
assert "".join(contents) == "Copilot response"
usage_chunk: Final = chunks[-1]
assert usage_chunk.usage.prompt_tokens == token_counter(
model=_MODEL,
messages=_messages(),
)
assert usage_chunk.usage.completion_tokens == token_counter(
model=_MODEL,
text="Copilot response",
count_response_tokens=True,
)
assert usage_chunk.usage.total_tokens == (usage_chunk.usage.prompt_tokens + usage_chunk.usage.completion_tokens)
def test_litellm_completion_forwards_token_exchange_settings() -> None:
assertion: Final = "header-dispatch-9306.payload-dispatch-9306.signature-dispatch-9306"
client, requests = _sync_client()
messages: Final[list[AllMessageValues]] = cast(
list[AllMessageValues],
[
{
"role": "user",
"content": [
{"type": "text", "text": "dispatch"},
{"type": "text", "text": "path"},
],
}
],
)
original_messages: Final = deepcopy(messages)
headers: Final[dict[str, str]] = {"X-Caller-Header": "unchanged"}
original_headers: Final = dict(headers)
response: Final = litellm.completion(
model=_MODEL,
messages=messages,
api_key="direct-dispatch-token-must-not-win-9306",
api_base="https://caller-controlled.example",
extra_headers=headers,
token_exchange_endpoint="https://identity.example.com/dispatch-9306/oauth2/v2.0/token",
token_exchange_profile="rfc8693",
token_exchange_scope="scope-dispatch-9306",
token_exchange_audience="audience-dispatch-9306",
client_id="client-dispatch-9306",
client_secret="secret-dispatch-9306",
secret_fields=_secret_fields(assertion),
time_zone="America/New_York",
client=client,
)
graph_requests: Final = tuple(request for request in requests if request.url.host == "graph.microsoft.com")
token_requests: Final = _token_requests(requests)
assert isinstance(response, ModelResponse)
assert len(graph_requests) == 2
assert len(token_requests) == 1
assert str(token_requests[0].url) == "https://identity.example.com/dispatch-9306/oauth2/v2.0/token"
assert {name: values[0] for name, values in parse_qs(token_requests[0].content.decode("utf-8")).items()} == {
"grant_type": "urn:ietf:params:oauth:grant-type:token-exchange",
"client_id": "client-dispatch-9306",
"client_secret": "secret-dispatch-9306",
"subject_token": assertion,
"subject_token_type": "urn:ietf:params:oauth:token-type:access_token",
"scope": "scope-dispatch-9306",
"audience": "audience-dispatch-9306",
}
assert all(request.headers["Authorization"] == "Bearer graph-access-token" for request in graph_requests)
assert assertion not in graph_requests[0].headers["Authorization"]
assert json.loads(graph_requests[1].content) == {
"message": {"text": "dispatch\npath"},
"locationHint": {"timeZone": "America/New_York"},
}
assert messages == original_messages
assert headers == original_headers
def test_litellm_completion_accepts_max_tokens_without_sending_it_to_graph() -> None:
client, requests = _sync_client()
response: Final = litellm.completion(
model=_MODEL,
messages=_messages(),
api_key="delegated-max-token",
max_tokens=16,
max_completion_tokens=32,
client=client,
)
graph_chat_requests: Final = tuple(request for request in requests if request.url.path.endswith("/chat"))
assert isinstance(response, ModelResponse)
assert json.loads(graph_chat_requests[-1].content) == {
"message": {"text": "hello Copilot"},
"locationHint": {"timeZone": "UTC"},
}
def test_litellm_completion_still_rejects_temperature() -> None:
client, _ = _sync_client()
with pytest.raises(litellm.UnsupportedParamsError):
litellm.completion(
model=_MODEL,
messages=_messages(),
api_key="delegated-temperature",
temperature=0.2,
client=client,
)

View file

@ -0,0 +1,335 @@
from copy import deepcopy
from typing import Final, Protocol, cast
import pytest
from pydantic import TypeAdapter
import litellm
from litellm.llms.anthropic.pass_through.adapters.transformation import LiteLLMAnthropicMessagesAdapter
from litellm.llms.microsoft_365_copilot.chat.transformation import (
Microsoft365CopilotChatConfig,
build_chat_request,
extract_graph_error_message,
map_graph_response,
)
from litellm.llms.microsoft_365_copilot.common_utils import Microsoft365CopilotError
from litellm.types.llms.anthropic import AllAnthropicPassThroughMessageValues
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import LlmProviders, Usage
from litellm.utils import ProviderConfigManager, token_counter
class _ModelResponseWithUsage(Protocol):
usage: Usage
def test_build_request_preserves_messages_and_maps_history() -> None:
message_values: Final[list[AllMessageValues]] = cast(
list[AllMessageValues],
[
{"role": "system", "content": "system context"},
{"role": "user", "content": "first question"},
{"role": "assistant", "content": "first answer"},
{
"role": "user",
"content": [
{"type": "text", "text": "second"},
{"type": "text", "text": "question", "cache_control": {"type": "ephemeral"}},
],
},
],
)
original_messages: Final = deepcopy(message_values)
request: Final = build_chat_request(messages=message_values, optional_params={})
assert request == {
"message": {"text": "second\nquestion"},
"additionalContext": [
{"text": "system context", "description": "system message"},
{"text": "first question", "description": "user message"},
{"text": "first answer", "description": "assistant message"},
],
"locationHint": {"timeZone": "UTC"},
}
assert message_values == original_messages
def test_build_request_uses_provider_time_zone_and_omits_empty_history() -> None:
messages: Final[list[AllMessageValues]] = TypeAdapter(list[AllMessageValues]).validate_python(
[{"role": "user", "content": "hello"}]
)
request: Final = build_chat_request(
messages=messages,
optional_params={"time_zone": "America/New_York"},
)
assert request == {
"message": {"text": "hello"},
"locationHint": {"timeZone": "America/New_York"},
}
def test_build_request_uses_last_user_message_when_assistant_is_last() -> None:
messages: Final[list[AllMessageValues]] = TypeAdapter(list[AllMessageValues]).validate_python(
[
{"role": "user", "content": "the prompt"},
{"role": "assistant", "content": "the response"},
]
)
request: Final = build_chat_request(messages=messages, optional_params={})
assert request == {
"message": {"text": "the prompt"},
"additionalContext": [{"text": "the response", "description": "assistant message"}],
"locationHint": {"timeZone": "UTC"},
}
def test_build_request_uses_last_user_message_before_trailing_system() -> None:
environment: Final = "# Environment\nYou have been invoked in the following environment: ..."
messages: Final[list[AllMessageValues]] = TypeAdapter(list[AllMessageValues]).validate_python(
[
{"role": "user", "content": "whats 1+1"},
{"role": "system", "content": environment},
]
)
request: Final = build_chat_request(messages=messages, optional_params={})
assert request == {
"message": {"text": "whats 1+1"},
"additionalContext": [{"text": environment, "description": "system message"}],
"locationHint": {"timeZone": "UTC"},
}
def test_build_request_preserves_context_order_after_last_user_message() -> None:
messages: Final[list[AllMessageValues]] = TypeAdapter(list[AllMessageValues]).validate_python(
[
{"role": "system", "content": "system A"},
{"role": "user", "content": "question one"},
{"role": "assistant", "content": "response one"},
{"role": "user", "content": "question two"},
{"role": "system", "content": "system B"},
]
)
request: Final = build_chat_request(messages=messages, optional_params={})
assert request == {
"message": {"text": "question two"},
"additionalContext": [
{"text": "system A", "description": "system message"},
{"text": "question one", "description": "user message"},
{"text": "response one", "description": "assistant message"},
{"text": "system B", "description": "system message"},
],
"locationHint": {"timeZone": "UTC"},
}
def test_build_request_requires_a_user_message() -> None:
messages: Final[list[AllMessageValues]] = TypeAdapter(list[AllMessageValues]).validate_python(
[{"role": "assistant", "content": "not a prompt"}]
)
with pytest.raises(Microsoft365CopilotError) as error:
build_chat_request(messages=messages, optional_params={})
assert error.value.status_code == 400
assert str(error.value) == "at least one message must have role 'user'"
def test_anthropic_adapter_preserves_trailing_system_context_for_m365() -> None:
environment: Final = "# Environment\nYou have been invoked in the following environment: ..."
desktop_messages: Final[list[AllAnthropicPassThroughMessageValues]] = cast(
list[AllAnthropicPassThroughMessageValues],
[
{"role": "user", "content": "whats 1+1"},
{"role": "system", "content": [{"type": "text", "text": environment}]},
],
)
translated_messages: Final = TypeAdapter(list[AllMessageValues]).validate_python(
LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai(
desktop_messages,
model="microsoft_365_copilot/chat",
custom_llm_provider="microsoft_365_copilot",
)
)
request: Final = build_chat_request(messages=translated_messages, optional_params={})
assert request == {
"message": {"text": "whats 1+1"},
"additionalContext": [{"text": environment, "description": "system message"}],
"locationHint": {"timeZone": "UTC"},
}
def test_build_request_rejects_non_text_content() -> None:
messages: Final[list[AllMessageValues]] = cast(
list[AllMessageValues],
[
{
"role": "user",
"content": [
{"type": "text", "text": "hello"},
{"type": "image_url", "image_url": {"url": "https://example.test/image.png"}},
],
}
],
)
with pytest.raises(Microsoft365CopilotError) as error:
build_chat_request(messages=messages, optional_params={})
assert error.value.status_code == 400
assert str(error.value) == "only text content is supported"
def test_map_graph_response_uses_last_message_and_estimates_usage() -> None:
messages: Final = TypeAdapter(tuple[AllMessageValues, ...]).validate_python(
[{"role": "user", "content": "known Copilot prompt"}]
)
reply: Final = "Copilot reply"
graph_response: Final = {
"messages": [
{"text": "prompt echo"},
{"text": reply},
]
}
response: Final = map_graph_response(
graph_response=graph_response,
model="microsoft_365_copilot/chat",
messages=messages,
)
assert response.model == "microsoft_365_copilot/chat"
assert response.choices[0].message.content == "Copilot reply"
assert response.choices[0].message.role == "assistant"
assert response.choices[0].finish_reason == "stop"
response_with_usage: Final = cast(_ModelResponseWithUsage, response)
expected_prompt_tokens: Final = token_counter(
model="microsoft_365_copilot/chat",
messages=messages,
)
expected_completion_tokens: Final = token_counter(
model="microsoft_365_copilot/chat",
text=reply,
count_response_tokens=True,
)
assert expected_prompt_tokens > 0
assert expected_completion_tokens > 0
assert response_with_usage.usage.prompt_tokens == expected_prompt_tokens
assert response_with_usage.usage.completion_tokens == expected_completion_tokens
assert response_with_usage.usage.total_tokens == expected_prompt_tokens + expected_completion_tokens
single_message_response: Final = map_graph_response(
graph_response={"messages": [{"text": "single reply"}]},
model="microsoft_365_copilot/chat",
messages=messages,
)
assert single_message_response.choices[0].message.content == "single reply"
@pytest.mark.parametrize(
("reply", "expected_reply"),
[
("pongpong", "pong"),
("1 + 1 = **2**.1 + 1 = **2**.", "1 + 1 = **2**."),
("pong", "pong"),
("pongpon", "pongpon"),
("pongPONG", "pongPONG"),
("", ""),
("aaaa", "aa"),
],
)
def test_map_graph_response_collapses_only_exact_duplicate_halves(reply: str, expected_reply: str) -> None:
messages: Final = TypeAdapter(tuple[AllMessageValues, ...]).validate_python(
[{"role": "user", "content": "known Copilot prompt"}]
)
response: Final = map_graph_response(
graph_response={"messages": [{"text": reply}]},
model="microsoft_365_copilot/chat",
messages=messages,
)
assert response.choices[0].message.content == expected_reply
def test_map_graph_response_counts_tokens_for_collapsed_reply() -> None:
messages: Final = TypeAdapter(tuple[AllMessageValues, ...]).validate_python(
[{"role": "user", "content": "known Copilot prompt"}]
)
doubled_reply_response: Final = map_graph_response(
graph_response={"messages": [{"text": "pongpong"}]},
model="microsoft_365_copilot/chat",
messages=messages,
)
single_reply_response: Final = map_graph_response(
graph_response={"messages": [{"text": "pong"}]},
model="microsoft_365_copilot/chat",
messages=messages,
)
doubled_reply_with_usage: Final = cast(_ModelResponseWithUsage, doubled_reply_response)
single_reply_with_usage: Final = cast(_ModelResponseWithUsage, single_reply_response)
assert doubled_reply_with_usage.usage.completion_tokens == single_reply_with_usage.usage.completion_tokens
@pytest.mark.parametrize(
"graph_response",
[
{},
{"messages": []},
{"messages": [{"text": "prompt"}, {}]},
],
)
def test_map_graph_response_rejects_missing_reply_text(graph_response: object) -> None:
with pytest.raises(Microsoft365CopilotError) as error:
map_graph_response(
graph_response=graph_response,
model="microsoft_365_copilot/chat",
messages=(),
)
assert error.value.status_code == 502
def test_extract_graph_error_message_uses_final_stringified_message() -> None:
graph_error: Final = {
"error": {"message": ('{"messages":[{"text":"prompt echo"},{"text":"Copilot requires a valid license"}]}')}
}
message: Final = extract_graph_error_message(graph_error)
empty_reply_error: Final = extract_graph_error_message({"error": {"message": '{"messages":[{"text":""}]}'}})
assert message == "Copilot requires a valid license"
assert empty_reply_error == ""
def test_provider_config_accepts_openai_compatibility_params() -> None:
config: Final = Microsoft365CopilotChatConfig()
assert config.get_supported_openai_params("microsoft_365_copilot/chat") == [
"stream",
"max_tokens",
"max_completion_tokens",
]
def test_provider_is_registered_for_model_resolution_and_chat_config() -> None:
_, provider, _, _ = litellm.get_llm_provider(model="microsoft_365_copilot/chat")
config: Final = ProviderConfigManager.get_provider_chat_config(
model="microsoft_365_copilot/chat",
provider=LlmProviders.MICROSOFT_365_COPILOT,
)
assert provider == LlmProviders.MICROSOFT_365_COPILOT.value
assert LlmProviders.MICROSOFT_365_COPILOT in litellm.provider_list
assert isinstance(config, Microsoft365CopilotChatConfig)

View file

@ -0,0 +1,27 @@
from typing import Final
import pytest
from litellm.llms.microsoft_365_copilot.common_utils import extract_caller_assertion
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
@pytest.mark.parametrize(
("authorization", "expected"),
[
("bEaReR header.payload.signature", "header.payload.signature"),
("Bearer header..signature", None),
("Bearer not-a-jwt", None),
("Basic header.payload.signature", None),
("Bearer header.payload.signature extra", None),
],
)
def test_extract_caller_assertion_requires_a_three_part_bearer_jwt(
authorization: str,
expected: str | None,
) -> None:
secret_fields: Final[SecretFields] = SecretFields(raw_headers={"aUtHoRiZaTiOn": authorization})
assertion: Final = extract_caller_assertion(secret_fields)
assert assertion == expected

View file

@ -35,16 +35,72 @@ from litellm.proxy.auth.auth_utils import (
is_request_body_safe,
)
from litellm.router import Router
from litellm.types.utils import oauth_token_exchange_litellm_params
from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS, OPENAI_WIF_KWARGS_KEYS
@pytest.mark.parametrize("param", sorted(ANTHROPIC_WIF_KWARGS_KEYS | OPENAI_WIF_KWARGS_KEYS))
def test_every_wif_kwarg_key_is_refused_from_a_request_body(param: str):
"""Every key the kwargs funnel carries into litellm_params selects a server-side secret or the
scope a token is minted for, so each one must be refused from a request body even with the
proxy-wide client-credential opt-in; a key added to the funnel without joining the ban shows up
here as a body the proxy accepted."""
with pytest.raises(ValueError, match="server-owned workload identity federation parameter"):
@pytest.mark.parametrize(
"field",
[
"token_exchange_endpoint",
"token_exchange_profile",
"token_exchange_scope",
"token_exchange_audience",
],
)
def test_token_exchange_settings_in_request_body_are_rejected(field: str) -> None:
with pytest.raises(
ValueError,
match="server-owned workload identity federation or OAuth token exchange parameter",
) as error:
is_request_body_safe(
request_body={"model": "microsoft_365_copilot/chat", field: "attacker-chosen"},
general_settings={"allow_client_side_credentials": True},
llm_router=None,
model="microsoft_365_copilot/chat",
route="/v1/chat/completions",
)
assert field in str(error.value)
def test_model_opt_in_cannot_allow_a_token_exchange_endpoint_in_a_request_body() -> None:
router: Final = Router(
model_list=[
{
"model_name": "microsoft_365_copilot/chat",
"litellm_params": {
"model": "microsoft_365_copilot/chat",
"configurable_clientside_auth_params": ["token_exchange_endpoint"],
},
}
]
)
with pytest.raises(
ValueError,
match="server-owned workload identity federation or OAuth token exchange parameter",
):
is_request_body_safe(
request_body={
"model": "microsoft_365_copilot/chat",
"token_exchange_endpoint": "https://identity.example.com/token",
},
general_settings={},
llm_router=router,
model="microsoft_365_copilot/chat",
)
@pytest.mark.parametrize(
"param",
sorted(ANTHROPIC_WIF_KWARGS_KEYS | OPENAI_WIF_KWARGS_KEYS | set(oauth_token_exchange_litellm_params)),
)
def test_every_server_owned_identity_param_is_refused_from_a_request_body(param: str):
with pytest.raises(
ValueError,
match="server-owned workload identity federation or OAuth token exchange parameter",
):
is_request_body_safe(
request_body={"model": "claude-sonnet-5", param: "attacker-chosen"},
general_settings={"allow_client_side_credentials": True},
@ -82,7 +138,10 @@ def test_a_request_body_cannot_pick_a_federated_identity_by_credential_name(monk
],
)
with pytest.raises(ValueError, match="names a credential configured for workload identity federation"):
with pytest.raises(
ValueError,
match="names a credential configured for workload identity federation or OAuth token exchange",
):
is_request_body_safe(
request_body=body,
general_settings={"allow_client_side_credentials": True},
@ -188,7 +247,10 @@ def test_configuring_a_deployment_may_name_a_federated_credential(federated_cred
def test_a_call_still_cannot_pick_a_federated_identity_by_credential_name(federated_credential, route: str | None):
"""The exemption covers the deployment-management routes and nothing that shares their prefix,
so a call still cannot move its token exchange onto a federated credential by naming it."""
with pytest.raises(ValueError, match="names a credential configured for workload identity federation"):
with pytest.raises(
ValueError,
match="names a credential configured for workload identity federation or OAuth token exchange",
):
is_request_body_safe(
request_body={"model": "claude-sonnet-5", "litellm_credential_name": "admin-wif"},
general_settings={"allow_client_side_credentials": True},
@ -202,7 +264,10 @@ def test_a_call_still_cannot_pick_a_federated_identity_by_credential_name(federa
def test_configuring_a_deployment_still_cannot_carry_federation_fields_inline(route: str):
"""Only the credential reference is exempt. Federation fields typed straight into a body stay
refused everywhere, since a stored credential is the surface an admin has to go through."""
with pytest.raises(ValueError, match="server-owned workload identity federation parameter"):
with pytest.raises(
ValueError,
match="server-owned workload identity federation or OAuth token exchange parameter",
):
is_request_body_safe(
request_body={"model": "claude-sonnet-5", "anthropic_federation_rule_id": "fdrl_attacker"},
general_settings={"allow_client_side_credentials": True},
@ -2360,6 +2425,51 @@ class TestGetDynamicLitellmParamsClearsAdminConfigOnBaseOverride:
assert out["api_base"] == "self-hosted.example.com:50051"
assert "use_ssl" not in out
def test_clears_oauth_token_exchange_fields_for_other_provider_base_overrides(self):
from litellm.router_utils.clientside_credential_handler import get_dynamic_litellm_params
oauth_values: Final = {
"token_exchange_audience": "https://graph.microsoft.com",
"token_exchange_endpoint": "https://identity.example.com/token",
"token_exchange_profile": "jwt_bearer_obo",
"token_exchange_scope": "https://graph.microsoft.com/.default",
"client_id": "copilot-client",
"client_secret": "copilot-secret",
}
out = get_dynamic_litellm_params(
litellm_params={
"model": "openai/gpt-4o",
"api_base": "https://admin.example.com/v1",
**oauth_values,
},
request_kwargs={"api_base": "https://caller.example.com/v1"},
)
assert all(field not in out for field in oauth_values)
def test_clears_oauth_token_exchange_fields_for_microsoft_365_copilot_base_overrides(self):
from litellm.router_utils.clientside_credential_handler import get_dynamic_litellm_params
oauth_values: Final = {
"token_exchange_audience": "https://graph.microsoft.com",
"token_exchange_endpoint": "https://identity.example.com/token",
"token_exchange_profile": "jwt_bearer_obo",
"token_exchange_scope": "https://graph.microsoft.com/.default",
"client_id": "copilot-client",
"client_secret": "copilot-secret",
}
out = get_dynamic_litellm_params(
litellm_params={
"model": "microsoft_365_copilot/chat",
"api_base": "https://graph.microsoft.com/beta",
**oauth_values,
},
request_kwargs={"api_base": "https://caller.example.com/v1"},
)
assert out["api_base"] == "https://caller.example.com/v1"
assert all(field not in out for field in oauth_values)
def test_caller_resupplied_value_overrides_admin_value_on_base_override(self):
# When the caller redirects ``api_base`` and *also* supplies their
# own value for one of the admin fields (e.g. ``organization``),

View file

@ -631,6 +631,35 @@ class TestNonAdminCannotPersistWifFieldsOnCredential:
assert response.status_code == 200, response.text
repository.create.assert_awaited_once()
def test_non_admin_cannot_create_a_credential_with_oauth_token_exchange_endpoint(self, restore_credential_list):
with _repository_holding(None) as repository:
response = _post_credential(
{
"credential_name": "attacker-oauth",
"credential_values": {"token_exchange_endpoint": "https://attacker.example/token"},
"credential_info": {"custom_llm_provider": "microsoft_365_copilot"},
},
auth=_as_non_admin,
)
assert response.status_code == 403, response.text
assert "token_exchange_endpoint" in response.json()["error"]["message"]
repository.create.assert_not_awaited()
def test_proxy_admin_can_create_a_credential_with_oauth_token_exchange_endpoint(self, restore_credential_list):
with _repository_holding(None) as repository:
response = _post_credential(
{
"credential_name": "admin-oauth",
"credential_values": {"token_exchange_endpoint": "https://identity.example.com/token"},
"credential_info": {"custom_llm_provider": "microsoft_365_copilot"},
},
auth=_as_admin,
)
assert response.status_code == 200, response.text
repository.create.assert_awaited_once()
def test_non_admin_cannot_create_a_credential_with_an_openai_token_file(self):
with patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.proxy_server.prisma_client", MagicMock()

View file

@ -33,6 +33,7 @@ from litellm.proxy.health_endpoints._health_endpoints import (
from litellm.proxy.health_endpoints._health_endpoints import (
test_model_connection as health_test_model_connection,
)
from litellm.types.utils import oauth_token_exchange_litellm_params
from litellm.types.workload_identity import (
ANTHROPIC_WIF_KWARGS_KEYS,
OPENAI_WIF_KWARGS_KEYS,
@ -1201,7 +1202,7 @@ async def test_test_connection_refuses_a_non_admin_pointing_a_federated_deployme
)
assert exc_info.value.code == "403"
assert "workload identity federation" in exc_info.value.message
assert "workload identity federation or OAuth token exchange" in exc_info.value.message
@pytest.mark.asyncio
@ -2531,6 +2532,10 @@ async def test_health_endpoint_keeps_federation_identity_admin_only():
"anthropic_keycloak_client_id": "litellm-proxy",
"openai_identity_provider_id": "idp_01H",
"openai_service_account_id": "sa_01H",
"token_exchange_audience": "https://graph.microsoft.com",
"token_exchange_endpoint": "https://identity.example.com/token",
"token_exchange_profile": "jwt_bearer_obo",
"token_exchange_scope": "https://graph.microsoft.com/.default",
}
full_model_list = [
{
@ -2585,29 +2590,27 @@ async def test_health_endpoint_keeps_federation_identity_admin_only():
assert non_admin_endpoint["model_id"] == "id-a"
@pytest.mark.parametrize("federation_field", sorted(ANTHROPIC_WIF_KWARGS_KEYS | OPENAI_WIF_KWARGS_KEYS))
def test_no_federation_field_reaches_a_non_admin_health_entry(federation_field: str):
"""Every key that configures workload identity federation either names the identity a
deployment mints as or carries the secret it mints with, and a non-admin who can see the
deployment is healthy must learn neither. Both lists that enforce that are derived from the
same key sets this runs over, so a field added to the funnel without joining either one shows
up here as a value a non-admin could read."""
@pytest.mark.parametrize(
"server_owned_field",
sorted(ANTHROPIC_WIF_KWARGS_KEYS | OPENAI_WIF_KWARGS_KEYS | set(oauth_token_exchange_litellm_params)),
)
def test_no_server_owned_identity_field_reaches_a_non_admin_health_entry(server_owned_field: str):
"""Every server-owned identity field is hidden from non-admin health entries."""
from litellm.proxy.health_check import clean_endpoint_data
from litellm.proxy.health_endpoints._health_endpoints import (
_strip_admin_only_fields_from_health_result,
)
canary = f"CANARY-{federation_field}-VALUE"
canary = f"CANARY-{server_owned_field}-VALUE"
cleaned = clean_endpoint_data(
{"model": "anthropic/claude-sonnet-5", federation_field: canary},
{"model": "anthropic/claude-sonnet-5", server_owned_field: canary},
details=True,
)
stripped = _strip_admin_only_fields_from_health_result(
{"healthy_endpoints": [cleaned], "unhealthy_endpoints": []}
)
stripped = _strip_admin_only_fields_from_health_result({"healthy_endpoints": [cleaned], "unhealthy_endpoints": []})
assert stripped["healthy_endpoints"][0]["model"] == "anthropic/claude-sonnet-5"
assert federation_field not in stripped["healthy_endpoints"][0]
assert server_owned_field not in stripped["healthy_endpoints"][0]
assert canary not in str(stripped)
@ -4081,6 +4084,40 @@ def test_health_test_connection_keeps_error_and_raw_request_through_the_allowlis
assert not {"api_key", "timeout", "exception"} & set(body["result"])
def test_health_test_connection_uses_request_headers_as_secret_fields_and_hides_them() -> None:
app: Final = FastAPI()
app.include_router(_health_endpoints_module.router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
ahealth_check: Final = AsyncMock(return_value={"status": "healthy"})
authorization: Final = "Bearer a.b.c"
body_authorization: Final = "Bearer body-supplied.invalid"
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.ahealth_check", new=ahealth_check),
):
response: Final = TestClient(app).post(
"/health/test_connection",
headers={"Authorization": authorization},
json={
"mode": "chat",
"litellm_params": {
"model": "microsoft_365_copilot/chat",
"secret_fields": {"raw_headers": {"authorization": body_authorization}},
},
},
)
assert response.status_code == 200, response.text
assert response.json()["status"] == "success"
model_params: Final = ahealth_check.call_args.kwargs["model_params"]
assert model_params["secret_fields"]["raw_headers"]["authorization"] == authorization
assert "secret_fields" not in response.text
assert "raw_headers" not in response.text
assert authorization not in response.text
assert body_authorization not in response.text
def test_clean_endpoint_data_keeps_only_json_safe_diagnostics():
"""
LIT-6907: _clean_endpoint_data used to copy every litellm_param not on a
@ -4179,22 +4216,25 @@ class TestConfigBaseForHealthCheck:
"rpm": 100,
}
def _base(self, config, request, allow_client_side_credentials=False):
def _base(self, config, request, allow_client_side_credentials=False, *, selected_by_id: bool):
from litellm.proxy.health_endpoints._health_endpoints import (
_config_base_for_health_check,
)
return _config_base_for_health_check(
config, request, allow_client_side_credentials=allow_client_side_credentials
config,
request,
allow_client_side_credentials=allow_client_side_credentials,
selected_by_id=selected_by_id,
)
def test_request_without_connection_fields_inherits_config(self):
base = self._base(self.CONFIG, {"model": "openai/gpt-4o"})
base = self._base(self.CONFIG, {"model": "openai/gpt-4o"}, selected_by_id=False)
assert base["api_key"] == "sk-configured"
assert base["api_base"] == "https://configured.example/v1"
def test_request_setting_api_base_does_not_inherit_config_credentials(self):
base = self._base(self.CONFIG, {"api_base": "https://caller.example/v1"})
base = self._base(self.CONFIG, {"api_base": "https://caller.example/v1"}, selected_by_id=False)
assert "api_key" not in base
assert "api_base" not in base
assert "vertex_credentials" not in base
@ -4208,7 +4248,7 @@ class TestConfigBaseForHealthCheck:
"api_base": "https://new-deployment.example/v1",
"api_key": "sk-new-deployment",
}
merged = {**self._base(self.CONFIG, request), **request}
merged = {**self._base(self.CONFIG, request, selected_by_id=False), **request}
assert merged["api_base"] == "https://new-deployment.example/v1"
assert merged["api_key"] == "sk-new-deployment"
assert "sk-configured" not in str(merged)
@ -4217,7 +4257,7 @@ class TestConfigBaseForHealthCheck:
"""A request that redirects the destination but supplies no credential
of its own gets none from the configuration."""
request = {"api_base": "https://elsewhere.example"}
merged = {**self._base(self.CONFIG, request), **request}
merged = {**self._base(self.CONFIG, request, selected_by_id=False), **request}
assert "api_key" not in merged
assert "sk-configured" not in str(merged)
@ -4225,6 +4265,7 @@ class TestConfigBaseForHealthCheck:
base = self._base(
{**self.CONFIG, "aws_secret_access_key": "configured-secret"},
{"aws_bedrock_runtime_endpoint": "https://caller.example"},
selected_by_id=False,
)
assert "api_key" not in base
assert "aws_secret_access_key" not in base
@ -4236,6 +4277,7 @@ class TestConfigBaseForHealthCheck:
self.CONFIG,
{"api_base": "https://caller.example/v1"},
allow_client_side_credentials=True,
selected_by_id=False,
)
assert base["api_key"] == "sk-configured"
@ -4243,7 +4285,7 @@ class TestConfigBaseForHealthCheck:
"""A stored-credential name resolves to the same secrets downstream, so a
request that redirects the destination must not keep it either."""
config = {**self.CONFIG, "litellm_credential_name": "OpenAI-prod"}
base = self._base(config, {"api_base": "https://caller.example/v1"})
base = self._base(config, {"api_base": "https://caller.example/v1"}, selected_by_id=False)
assert "litellm_credential_name" not in base
assert "api_key" not in base
@ -4254,19 +4296,28 @@ class TestConfigBaseForHealthCheck:
base = self._base(
config,
{"model": "openai/gpt-4o", "litellm_credential_name": "OpenAI-prod", "custom_llm_provider": "openai"},
selected_by_id=False,
)
assert base["litellm_credential_name"] == "OpenAI-prod"
assert base["api_key"] == "sk-configured"
def test_request_naming_another_credential_does_not_inherit_config_credentials(self):
base = self._base(self.CONFIG, {"model": "openai/gpt-4o", "litellm_credential_name": "Another-cred"})
base = self._base(
self.CONFIG,
{"model": "openai/gpt-4o", "litellm_credential_name": "Another-cred"},
selected_by_id=False,
)
assert "api_key" not in base
assert "api_base" not in base
assert "vertex_credentials" not in base
assert base["rpm"] == 100
def test_blank_credential_name_names_no_credential(self):
base = self._base(self.CONFIG, {"model": "openai/gpt-4o", "litellm_credential_name": ""})
base = self._base(
self.CONFIG,
{"model": "openai/gpt-4o", "litellm_credential_name": ""},
selected_by_id=False,
)
assert base["api_key"] == "sk-configured"
def test_opt_in_does_not_put_config_credentials_over_a_named_credential(self):
@ -4274,9 +4325,76 @@ class TestConfigBaseForHealthCheck:
self.CONFIG,
{"model": "openai/gpt-4o", "litellm_credential_name": "Another-cred"},
allow_client_side_credentials=True,
selected_by_id=False,
)
assert "api_key" not in base
def test_request_with_its_own_api_key_does_not_inherit_config_auth(self):
config: Final = {
**self.CONFIG,
"token_exchange_endpoint": "https://login.example/token",
"client_id": "configured-client-id",
"client_secret": "configured-client-secret",
"litellm_credential_name": "configured-credential",
}
request: Final = {"model": "openai/gpt-4o", "api_key": "sk-request"}
merged: Final = {**self._base(config, request, selected_by_id=False), **request}
assert merged["api_key"] == "sk-request"
assert {
"token_exchange_endpoint",
"client_id",
"client_secret",
"litellm_credential_name",
"api_base",
"vertex_credentials",
}.isdisjoint(merged)
assert merged["rpm"] == 100
assert "sk-configured" not in str(merged)
def test_request_api_key_does_not_inherit_config_auth_with_client_credentials_opt_in(self):
config: Final = {
**self.CONFIG,
"token_exchange_endpoint": "https://login.example/token",
"client_id": "configured-client-id",
"client_secret": "configured-client-secret",
"litellm_credential_name": "configured-credential",
}
request: Final = {"model": "openai/gpt-4o", "api_key": "sk-request"}
base: Final = self._base(config, request, allow_client_side_credentials=True, selected_by_id=False)
merged: Final = {**base, **request}
assert merged["api_key"] == "sk-request"
assert {
"token_exchange_endpoint",
"client_id",
"client_secret",
"litellm_credential_name",
"api_base",
"vertex_credentials",
}.isdisjoint(merged)
assert merged["rpm"] == 100
assert "sk-configured" not in str(merged)
def test_blank_request_api_key_inherits_config_auth(self):
base: Final = self._base(
self.CONFIG,
{"model": "openai/gpt-4o", "api_key": ""},
selected_by_id=False,
)
assert base["api_key"] == "sk-configured"
def test_request_api_key_inherits_saved_endpoint_when_selected_by_id(self):
config: Final = {**self.CONFIG, "api_version": "test-version"}
request: Final = {"model": "openai/gpt-4o", "api_key": "sk-request"}
base: Final = self._base(config, request, selected_by_id=True)
merged: Final = {**base, **request}
assert merged["api_base"] == "https://configured.example/v1"
assert merged["api_version"] == "test-version"
assert merged["api_key"] == "sk-request"
class TestTestConnectionUsesTheNamedCredential:
CREDENTIAL_KEY = "sk-credential-key"
@ -4352,6 +4470,193 @@ class TestTestConnectionUsesTheNamedCredential:
assert response.json()["status"] == "success", response.text
return probe
def test_oauth_connection_test_uses_caller_token_and_keeps_401_without_it(
self, monkeypatch: pytest.MonkeyPatch
) -> None:
caller_token: Final = "a.b.c"
body_token: Final = "body.spoof.token"
endpoint: Final = "https://identity.example.com/health-check/oauth2/v2.0/token"
deployment: Final = {
"model_name": "microsoft_365_copilot/chat",
"litellm_params": {
"model": "microsoft_365_copilot/chat",
"custom_llm_provider": "microsoft_365_copilot",
"token_exchange_endpoint": endpoint,
"client_id": "health-check-client",
"client_secret": "health-check-secret",
"token_exchange_profile": "jwt_bearer_obo",
"token_exchange_scope": "https://graph.microsoft.com/.default",
},
"model_info": {"mode": "chat"},
}
request_body: Final = {
"mode": "chat",
"litellm_params": {
"model": "microsoft_365_copilot/chat",
"secret_fields": {"raw_headers": {"authorization": f"Bearer {body_token}"}},
},
"model_info": {"mode": "chat"},
}
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
app: Final = FastAPI()
app.include_router(_health_endpoints_module.router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
client: Final = TestClient(app)
router: Final = MagicMock()
router.get_model_list.return_value = [deployment]
router.get_deployment.return_value = None
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.llm_router", router),
respx.mock(assert_all_called=True) as respx_mock,
):
token_route: Final = respx_mock.post(endpoint).respond(
json={"access_token": "health-graph-token", "expires_in": 3600}
)
respx_mock.post("https://graph.microsoft.com/beta/copilot/conversations").respond(
status_code=201,
json={"id": "health-check-conversation", "state": "active"},
)
respx_mock.post(
"https://graph.microsoft.com/beta/copilot/conversations/health-check-conversation/chat"
).respond(json={"messages": [{"text": "prompt echo"}, {"text": "Copilot response"}]})
response: Final = client.post(
"/health/test_connection",
headers={"Authorization": f"Bearer {caller_token}"},
json=request_body,
)
assert response.status_code == 200, response.text
assert response.json()["status"] == "success", response.text
token_request: Final = respx_mock.calls[0].request
assert f"assertion={caller_token}".encode() in token_request.content
assert body_token.encode() not in token_request.content
unauthorized_response: Final = client.post("/health/test_connection", json=request_body)
assert unauthorized_response.status_code == 200, unauthorized_response.text
assert unauthorized_response.json()["status"] == "error"
assert (
"requires the caller's IdP-issued access token in the Authorization header"
in unauthorized_response.json()["result"]["error"]
)
assert token_route.call_count == 1
def test_request_api_key_does_not_inherit_microsoft_365_oauth_config(self):
deployment: Final = {
"model_name": "microsoft_365_copilot/chat",
"litellm_params": {
"model": "microsoft_365_copilot/chat",
"custom_llm_provider": "microsoft_365_copilot",
"token_exchange_endpoint": "https://login.example/token",
"client_id": "configured-client-id",
"client_secret": "configured-client-secret",
},
"model_info": {},
}
app: Final = FastAPI()
app.include_router(_health_endpoints_module.router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
router: Final = MagicMock()
router.get_model_list.return_value = [deployment]
ahealth_check: Final = MagicMock(return_value={"status": "healthy"})
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.llm_router", router),
patch("litellm.proxy.proxy_server.premium_user", False),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
AsyncMock(),
),
patch("litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check", ahealth_check),
patch(
"litellm.proxy.health_endpoints._health_endpoints.run_with_timeout",
AsyncMock(return_value={"status": "healthy"}),
),
):
response: Final = TestClient(app).post(
"/health/test_connection",
json={
"litellm_params": {
"model": "microsoft_365_copilot/chat",
"custom_llm_provider": "microsoft_365_copilot",
"api_key": "dummy-static-token",
},
"model_info": {},
},
)
assert response.status_code == 200, response.text
assert response.json()["status"] == "success", response.text
model_params: Final = ahealth_check.call_args.kwargs["model_params"]
assert model_params["api_key"] == "dummy-static-token"
assert "token_exchange_endpoint" not in model_params
assert "client_id" not in model_params
assert "client_secret" not in model_params
def test_request_api_key_keeps_saved_endpoint_when_deployment_selected_by_id(self):
from litellm.types.router import Deployment, LiteLLM_Params
model: Final = "microsoft_365_copilot/chat"
saved_api_base: Final = "https://configured.example/v1"
saved_api_version: Final = "test-version"
deployment: Final = Deployment(
model_name=model,
litellm_params=LiteLLM_Params(
model=model,
custom_llm_provider="microsoft_365_copilot",
api_key="saved-static-token",
api_base=saved_api_base,
api_version=saved_api_version,
),
model_info={"id": "m365-deployment"},
)
app: Final = FastAPI()
app.include_router(_health_endpoints_module.router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
router: Final = MagicMock()
router.get_deployment.return_value = deployment
ahealth_check: Final = MagicMock(return_value={"status": "healthy"})
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.llm_router", router),
patch("litellm.proxy.proxy_server.premium_user", False),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
AsyncMock(),
),
patch("litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check", ahealth_check),
patch(
"litellm.proxy.health_endpoints._health_endpoints.run_with_timeout",
AsyncMock(return_value={"status": "healthy"}),
),
):
response: Final = TestClient(app).post(
"/health/test_connection",
json={
"litellm_params": {
"model": model,
"custom_llm_provider": "microsoft_365_copilot",
"api_key": "replacement-static-token",
},
"model_info": {"id": "m365-deployment"},
},
)
assert response.status_code == 200, response.text
assert response.json()["status"] == "success", response.text
model_params: Final = ahealth_check.call_args.kwargs["model_params"]
assert model_params["api_base"] == saved_api_base
assert model_params["api_version"] == saved_api_version
assert model_params["api_key"] == "replacement-static-token"
def test_named_credentials_key_is_sent_not_the_matched_deployments_key(self, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [self._credential(api_key=self.CREDENTIAL_KEY)])

View file

@ -8598,7 +8598,11 @@ class TestNonAdminCannotPersistWifFieldsOnModel:
),
):
with pytest.raises(
Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity"
Exception,
match=(
"Only proxy admins can change the credentials of a deployment configured for "
"workload identity federation or OAuth token exchange"
),
) as exc_info:
await patch_model(
model_id="m1",
@ -8650,7 +8654,11 @@ class TestNonAdminCannotPersistWifFieldsOnModel:
),
):
with pytest.raises(
Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity"
Exception,
match=(
"Only proxy admins can change the credentials of a deployment configured for "
"workload identity federation or OAuth token exchange"
),
) as exc_info:
await patch_model(
model_id="m1",
@ -8725,15 +8733,26 @@ class TestNonAdminCannotPersistWifFieldsOnModel:
mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once()
@pytest.mark.asyncio
async def test_add_new_model_non_admin_cannot_set_wif_field(self):
async def test_add_new_model_team_admin_cannot_set_oauth_token_exchange_endpoint(self):
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.model_management_endpoints import (
add_new_model,
)
non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER)
non_admin = UserAPIKeyAuth(
user_id="team_admin",
team_id="oauth-exchange-team",
user_role=LitellmUserRoles.INTERNAL_USER,
)
mock_prisma = MagicMock()
mock_prisma.writer_db.litellm_uisettings.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(
return_value=LiteLLM_TeamTable(
team_id="oauth-exchange-team",
team_alias="oauth-exchange-team",
members_with_roles=[Member(user_id="team_admin", role="admin")],
)
)
with (
patch( # test-quality-ok: the proxy wiring under test is what this patches
@ -8754,19 +8773,145 @@ class TestNonAdminCannotPersistWifFieldsOnModel:
model_params=Deployment(
model_name="my-model",
litellm_params=LiteLLM_Params(
model="anthropic/claude-sonnet-4",
anthropic_keycloak_client_secret_ref="os.environ/LITELLM_MASTER_KEY",
model="microsoft_365_copilot/chat",
token_exchange_endpoint="https://identity.example.com/token",
),
model_info={"id": "wif-gate-create-0"},
model_info={"id": "oauth-exchange-create-0", "team_id": "oauth-exchange-team"},
),
user_api_key_dict=non_admin,
)
assert "proxy admin" in str(exc_info.value.message).lower()
assert exc_info.value.param == "anthropic_keycloak_client_secret_ref"
assert exc_info.value.code == "403"
assert exc_info.value.param == "token_exchange_endpoint"
mock_prisma.db.litellm_proxymodeltable.create.assert_not_called()
@staticmethod
def _team_admin_patch_fixtures(
*,
server_owned_oauth: bool,
) -> tuple[UserAPIKeyAuth, MagicMock, MagicMock]:
team_id: Final = "oauth-client-patch-team"
team_admin: Final = UserAPIKeyAuth(
user_id="team_admin",
team_id=team_id,
user_role=LitellmUserRoles.INTERNAL_USER,
)
existing_row: Final = MagicMock()
existing_row.litellm_params = {
"model": "microsoft_365_copilot/chat",
"api_key": "stored-api-key",
"client_id": "stored-client-id",
"client_secret": "stored-client-secret",
**(
{"token_exchange_endpoint": "https://identity.example.com/token"}
if server_owned_oauth
else {}
),
}
existing_row.model_dump.return_value = {
"model_name": "copilot",
"litellm_params": existing_row.litellm_params,
"model_info": {"id": "oauth-client-patch-1", "team_id": team_id},
}
existing_row.model_dump_json.return_value = "{}"
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row)
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=existing_row)
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(
return_value=LiteLLM_TeamTable(
team_id=team_id,
team_alias=team_id,
members_with_roles=[Member(user_id="team_admin", role="admin")],
)
)
return team_admin, mock_prisma, existing_row
@pytest.mark.asyncio
async def test_add_new_model_admin_can_set_wif_field(self):
@pytest.mark.parametrize(
("field", "value"),
(
("client_secret", "replacement-client-secret"),
("client_id", "replacement-client-id"),
),
)
async def test_team_admin_cannot_patch_oauth_client_credentials(
self,
field: str,
value: str,
) -> None:
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model
team_admin, mock_prisma, _ = self._team_admin_patch_fixtures(server_owned_oauth=True)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy wiring under test
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.proxy_server.llm_router",
MagicMock(**{"get_model_ids.return_value": ["oauth-client-patch-1"]}),
),
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
side_effect=lambda value: value,
),
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.management_endpoints.model_management_endpoints.clear_cache",
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
),
):
with pytest.raises(ProxyException) as exc_info:
await patch_model(
model_id="oauth-client-patch-1",
patch_data=updateDeployment(
litellm_params=updateLiteLLMParams.model_validate({field: value})
),
user_api_key_dict=team_admin,
)
assert exc_info.value.code == "403"
mock_prisma.db.litellm_proxymodeltable.update.assert_not_called()
@pytest.mark.asyncio
async def test_team_admin_can_patch_client_secret_without_server_owned_oauth(self) -> None:
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model
team_admin, mock_prisma, existing_row = self._team_admin_patch_fixtures(server_owned_oauth=False)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy wiring under test
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.proxy_server.llm_router",
MagicMock(**{"get_model_ids.return_value": ["oauth-client-patch-1"]}),
),
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
side_effect=lambda value: value,
),
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.management_endpoints.model_management_endpoints.clear_cache",
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
),
):
result = await patch_model(
model_id="oauth-client-patch-1",
patch_data=updateDeployment(
litellm_params=updateLiteLLMParams(client_secret="replacement-client-secret")
),
user_api_key_dict=team_admin,
)
assert result is existing_row
saved_params = json.loads(
mock_prisma.db.litellm_proxymodeltable.update.await_args.kwargs["data"]["litellm_params"]
)
assert saved_params["client_secret"] == "replacement-client-secret"
@pytest.mark.asyncio
async def test_add_new_model_proxy_admin_can_set_oauth_token_exchange_endpoint(self):
from litellm.proxy.management_endpoints.model_management_endpoints import (
add_new_model,
)
@ -8808,10 +8953,12 @@ class TestNonAdminCannotPersistWifFieldsOnModel:
model_params=Deployment(
model_name="my-model",
litellm_params=LiteLLM_Params(
model="anthropic/claude-sonnet-4",
anthropic_keycloak_client_secret_ref="os.environ/ANTHROPIC_WIF_CLIENT_SECRET",
model="microsoft_365_copilot/chat",
token_exchange_endpoint="https://identity.example.com/token",
client_id="copilot-client",
client_secret="copilot-secret",
),
model_info={"id": "wif-gate-create-1"},
model_info={"id": "oauth-exchange-create-1"},
),
user_api_key_dict=admin,
)
@ -8855,7 +9002,11 @@ class TestNonAdminCannotPersistWifFieldsOnModel:
),
):
with pytest.raises(
Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity"
Exception,
match=(
"Only proxy admins can change the credentials of a deployment configured for "
"workload identity federation or OAuth token exchange"
),
) as exc_info:
await update_model(
model_params=updateDeployment(
@ -9022,7 +9173,11 @@ class TestWifBoundaryReadsTheResultingDeployment:
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test
):
with pytest.raises(
Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity"
Exception,
match=(
"Only proxy admins can change the credentials of a deployment configured for "
"workload identity federation or OAuth token exchange"
),
):
await patch_model(
model_id="m1",
@ -9074,7 +9229,11 @@ class TestWifBoundaryReadsTheResultingDeployment:
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test
):
with pytest.raises(
Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity"
Exception,
match=(
"Only proxy admins can change the credentials of a deployment configured for "
"workload identity federation or OAuth token exchange"
),
):
await patch_model(
model_id="m1",
@ -9092,28 +9251,38 @@ class TestWifBoundaryReadsTheResultingDeployment:
imports whatever the credential holds, so the resulting deployment federates."""
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model
non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER)
non_admin = UserAPIKeyAuth(
user_id="team_admin",
team_id="oauth-exchange-team",
user_role=LitellmUserRoles.INTERNAL_USER,
)
plain_row = MagicMock()
plain_row.litellm_params = {"model": "anthropic/claude-sonnet-4"}
plain_row.litellm_params = {"model": "microsoft_365_copilot/chat"}
plain_row.model_dump.return_value = {
"model_name": "claude",
"model_name": "copilot",
"litellm_params": plain_row.litellm_params,
"model_info": {"id": "m1"},
"model_info": {"id": "m1", "team_id": "oauth-exchange-team"},
}
# The credential is served from the row rather than this pod's memory, which is both the
# multi-pod case and the one the gate must not miss.
admin_credential_row = {
"credential_name": "admin-wif",
"credential_values": {
"anthropic_federation_rule_id": "fdrl_admin",
"anthropic_organization_id": "org-admin",
"token_exchange_endpoint": "https://identity.example.com/token",
},
"credential_info": {"custom_llm_provider": "anthropic"},
"credential_info": {"custom_llm_provider": "microsoft_365_copilot"},
}
mock_prisma = MagicMock()
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=plain_row)
mock_prisma.db.litellm_credentialstable.find_unique = AsyncMock(return_value=admin_credential_row)
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(
return_value=LiteLLM_TeamTable(
team_id="oauth-exchange-team",
team_alias="oauth-exchange-team",
members_with_roles=[Member(user_id="team_admin", role="admin")],
)
)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy wiring under test
@ -9124,8 +9293,12 @@ class TestWifBoundaryReadsTheResultingDeployment:
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test
):
with pytest.raises(
Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity"
):
Exception,
match=(
"Only proxy admins can change the credentials of a deployment configured for "
"workload identity federation or OAuth token exchange"
),
) as exc_info:
await patch_model(
model_id="m1",
patch_data=updateDeployment(
@ -9133,6 +9306,64 @@ class TestWifBoundaryReadsTheResultingDeployment:
),
user_api_key_dict=non_admin,
)
assert getattr(exc_info.value, "param", "") == "token_exchange_endpoint"
@pytest.mark.asyncio
async def test_proxy_admin_can_attach_an_existing_oauth_exchange_credential(self, monkeypatch):
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model
admin = UserAPIKeyAuth(user_id="proxy-admin", user_role=LitellmUserRoles.PROXY_ADMIN)
plain_row = MagicMock()
plain_row.litellm_params = {"model": "microsoft_365_copilot/chat"}
plain_row.model_dump.return_value = {
"model_name": "copilot",
"litellm_params": plain_row.litellm_params,
"model_info": {"id": "m1"},
}
plain_row.model_dump_json.return_value = "{}"
updated_row = MagicMock()
updated_row.model_dump_json.return_value = "{}"
mock_prisma = MagicMock()
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=plain_row)
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row)
monkeypatch.setattr(
"litellm.credential_list",
[
CredentialItem(
credential_name="admin-oauth-exchange",
credential_values={"token_exchange_endpoint": "https://identity.example.com/token"},
credential_info={"custom_llm_provider": "microsoft_365_copilot"},
)
],
)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]})),
patch("litellm.proxy.proxy_server.store_model_in_db", True),
patch("litellm.proxy.proxy_server.premium_user", True),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
side_effect=lambda value: value,
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.clear_cache",
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
),
):
await patch_model(
model_id="m1",
patch_data=updateDeployment(
litellm_params=updateLiteLLMParams(litellm_credential_name="admin-oauth-exchange")
),
user_api_key_dict=admin,
)
mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once()
saved_params = json.loads(
mock_prisma.db.litellm_proxymodeltable.update.await_args.kwargs["data"]["litellm_params"]
)
assert saved_params["litellm_credential_name"] == "admin-oauth-exchange"
@pytest.mark.asyncio
async def test_non_admin_cannot_modify_a_deployment_whose_stored_credential_name_is_encrypted(self, monkeypatch):
@ -9179,7 +9410,11 @@ class TestWifBoundaryReadsTheResultingDeployment:
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test
):
with pytest.raises(
Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity"
Exception,
match=(
"Only proxy admins can change the credentials of a deployment configured for "
"workload identity federation or OAuth token exchange"
),
):
await patch_model(
model_id="m1",
@ -9210,12 +9445,11 @@ class TestFederationGateScopesToWhatTheWriteTouches:
def _federated_row(cls):
row = MagicMock()
row.litellm_params = {
"model": "anthropic/claude-sonnet-4",
"anthropic_federation_rule_id": "fdrl_admin",
"anthropic_organization_id": "org-admin",
"model": "microsoft_365_copilot/chat",
"token_exchange_endpoint": "https://identity.example.com/token",
}
row.model_dump.return_value = {
"model_name": "claude",
"model_name": "copilot",
"litellm_params": row.litellm_params,
"model_info": {"id": "m1", "team_id": cls._TEAM_ID},
}
@ -9236,7 +9470,7 @@ class TestFederationGateScopesToWhatTheWriteTouches:
return mock_prisma
@pytest.mark.asyncio
async def test_team_admin_can_still_set_rpm_on_a_federated_deployment(self):
async def test_team_admin_can_still_set_rpm_on_an_oauth_exchange_deployment(self):
"""rpm cannot move or re-scope the token the deployment mints, so it stays a team edit."""
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model
@ -9287,7 +9521,10 @@ class TestFederationGateScopesToWhatTheWriteTouches:
):
with pytest.raises(
Exception,
match="Only proxy admins can change the credentials of a deployment configured for workload identity",
match=(
"Only proxy admins can change the credentials of a deployment configured for "
"workload identity federation or OAuth token exchange"
),
):
await patch_model(
model_id="m1",

View file

@ -113,6 +113,10 @@ DEPLOYMENT_PARAM_CLASSIFICATION: Final[Mapping[str, Classification]] = MappingPr
"azure_ad_token": Unplanted(),
"client_secret": Unplanted(),
"azure_password": Unplanted(),
"token_exchange_endpoint": NotSecret("OAuth token endpoint URL"),
"token_exchange_profile": NotSecret("OAuth grant profile"),
"token_exchange_scope": NotSecret("OAuth scope string"),
"token_exchange_audience": NotSecret("OAuth audience identifier"),
"vertex_credentials": Secret("B4v"),
"aws_access_key_id": Unplanted(),
"aws_secret_access_key": Secret("B4"),

View file

@ -5,6 +5,7 @@ import time
from typing import Final
from unittest.mock import MagicMock, patch
import httpx
import pytest
import litellm
@ -25,6 +26,7 @@ from litellm.router_utils.cooldown_handlers import (
_should_run_cooldown_logic,
async_get_cooldown_deployments,
cast_exception_status_to_int,
get_cooldown_deployments,
mark_advisor_orchestration_failure,
should_cooldown_based_on_allowed_fails_policy,
)
@ -32,6 +34,7 @@ from litellm.router_utils.fallback_event_handlers import (
_trigger_cooldown_for_failed_deployment,
)
from litellm.router_utils.router_callbacks.track_deployment_metrics import (
get_deployment_failures_for_current_minute,
increment_deployment_failures_for_current_minute,
increment_deployment_successes_for_current_minute,
)
@ -40,6 +43,7 @@ from litellm.types.router import (
DeploymentTypedDict,
LiteLLMParamsTypedDict,
)
from litellm.types.utils import CredentialItem
from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome
@ -1260,6 +1264,183 @@ class TestDeploymentCallbackOnFailureCooldownTimePrecedence:
call_kwargs = mock_set_cooldown.call_args[1]
assert call_kwargs["time_to_cooldown"] == 20.0, "litellm_params.cooldown_time must still be honored"
@pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown")
class TestCallerScopedOAuthAuthFailureCooldown:
@staticmethod
def _router(
model_id: str,
litellm_params: dict[str, object],
allowed_fails_policy: dict[str, int] | None = None,
) -> Router:
return _make_router(
model_list=[
{
"model_name": "copilot",
"litellm_params": {
"model": "microsoft_365_copilot/chat",
**litellm_params,
},
"model_info": {
"id": model_id,
**(
{"allowed_fails_policy": allowed_fails_policy}
if allowed_fails_policy is not None
else {}
),
},
}
]
)
@staticmethod
def _auth_exception(status: int) -> Exception:
if status == 401:
return litellm.AuthenticationError("caller assertion rejected", "microsoft_365_copilot", "copilot")
return litellm.PermissionDeniedError(
"caller assertion rejected",
"microsoft_365_copilot",
"copilot",
response=httpx.Response(
status_code=403,
request=httpx.Request("GET", "https://litellm.ai"),
),
)
@staticmethod
def _callback(router: Router, model_id: str, exception: Exception) -> bool:
return router.deployment_callback_on_failure(
kwargs={
"exception": exception,
"litellm_params": {"model_info": {"id": model_id}},
},
completion_response=None,
start_time=0,
end_time=1,
)
@pytest.mark.asyncio
@pytest.mark.parametrize("status", (401, 403))
async def test_inline_oauth_caller_auth_failure_does_not_cooldown(self, status: int) -> None:
model_id: Final = "inline-oauth"
router: Final = self._router(
model_id,
{
"token_exchange_endpoint": "https://identity.example.com/token",
"client_id": "copilot-client",
"client_secret": "copilot-secret",
},
)
result: Final = self._callback(router, model_id, self._auth_exception(status))
assert result is False
assert get_deployment_failures_for_current_minute(router, model_id) == 0
assert get_cooldown_deployments(router, parent_otel_span=None) == []
@pytest.mark.asyncio
@pytest.mark.parametrize("status", (401, 403))
async def test_named_oauth_credential_caller_auth_failure_does_not_cooldown(
self,
status: int,
monkeypatch: pytest.MonkeyPatch,
) -> None:
model_id: Final = "named-oauth"
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="copilot-oauth",
credential_values={
"token_exchange_endpoint": "https://identity.example.com/token",
"client_id": "copilot-client",
"client_secret": "copilot-secret",
},
credential_info={"custom_llm_provider": "microsoft_365_copilot"},
)
],
)
router: Final = self._router(model_id, {"litellm_credential_name": "copilot-oauth"})
result: Final = self._callback(router, model_id, self._auth_exception(status))
assert result is False
assert get_deployment_failures_for_current_minute(router, model_id) == 0
assert get_cooldown_deployments(router, parent_otel_span=None) == []
@pytest.mark.asyncio
async def test_api_key_deployment_still_cools_down_on_401(self) -> None:
model_id: Final = "api-key-only"
router: Final = self._router(model_id, {"api_key": "sk-test"})
self._callback(
router,
model_id,
litellm.AuthenticationError("upstream rejected the API key", "openai", "gpt-4o-mini"),
)
assert get_deployment_failures_for_current_minute(router, model_id) == 1
assert get_cooldown_deployments(router, parent_otel_span=None) == [model_id]
@pytest.mark.asyncio
@pytest.mark.parametrize("status", (429, 500))
async def test_oauth_deployment_still_cools_down_on_provider_failures(self, status: int) -> None:
model_id: Final = f"oauth-provider-{status}"
router: Final = self._router(
model_id,
{
"token_exchange_endpoint": "https://identity.example.com/token",
"client_id": "copilot-client",
"client_secret": "copilot-secret",
},
allowed_fails_policy={
"RateLimitErrorAllowedFails": 0,
"InternalServerErrorAllowedFails": 0,
},
)
exception: Final[Exception] = (
litellm.RateLimitError("rate limited", "microsoft_365_copilot", "copilot")
if status == 429
else litellm.InternalServerError("provider failed", "microsoft_365_copilot", "copilot")
)
self._callback(router, model_id, exception)
assert get_deployment_failures_for_current_minute(router, model_id) == 1
assert get_cooldown_deployments(router, parent_otel_span=None) == [model_id]
@pytest.mark.parametrize("status", (401, 403))
def test_fallback_oauth_caller_auth_failure_does_not_cooldown(self, status: int) -> None:
model_id: Final = f"fallback-oauth-{status}"
router: Final = self._router(
model_id,
{
"token_exchange_endpoint": "https://identity.example.com/token",
"client_id": "copilot-client",
"client_secret": "copilot-secret",
},
)
exception: Final = self._auth_exception(status)
exception.failed_deployment_id = model_id
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exception)
assert get_deployment_failures_for_current_minute(router, model_id) == 0
assert get_cooldown_deployments(router, parent_otel_span=None) == []
@pytest.mark.asyncio
async def test_fallback_api_key_auth_failure_still_cools_down(self) -> None:
model_id: Final = "fallback-api-key"
router: Final = self._router(model_id, {"api_key": "sk-test"})
exception: Final = litellm.AuthenticationError("API key rejected", "openai", "gpt-4o-mini")
exception.failed_deployment_id = model_id
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exception)
assert get_deployment_failures_for_current_minute(router, model_id) == 1
assert get_cooldown_deployments(router, parent_otel_span=None) == [model_id]
@pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown")
class TestNewAllowedFailsPolicyFields:
def test_service_unavailable_error_matched_by_policy(self):

View file

@ -1165,7 +1165,10 @@ async def test_a_stored_fallback_target_cannot_carry_a_federation_field():
so a stored key/team/global fallback could otherwise set the workspace a federation token is
minted for. The request itself is already forbidden to carry these, and a stored setting is
not a more trusted source than the request."""
with pytest.raises(ValueError, match="server-owned workload identity federation parameter"):
with pytest.raises(
ValueError,
match="server-owned workload identity federation or OAuth token exchange parameter",
):
await run_async_fallback(
litellm_router=FakeRouter(),
fallback_model_group=[{"model": "anthropic-backup", "anthropic_federation_workspace_id": "wrkspc_other"}],

View file

@ -45,6 +45,7 @@ _IN_MEMORY_ONLY_CALLERS: Final = frozenset(
"litellm/integrations/newrelic/newrelic_team_handler.py",
"litellm/integrations/shadow_eval_logger.py",
"litellm/litellm_core_utils/litellm_logging.py",
"litellm/litellm_core_utils/oauth_token_exchange.py",
"litellm/litellm_core_utils/prompt_templates/factory.py",
"litellm/litellm_core_utils/prompt_templates/image_handling.py",
"litellm/litellm_core_utils/specialty_caches/dynamic_logging_cache.py",

View file

@ -78,6 +78,10 @@ CONNECTION_NAMES: Final = (
"tenant_id",
"client_id",
"client_secret",
"token_exchange_endpoint",
"token_exchange_profile",
"token_exchange_scope",
"token_exchange_audience",
"azure_username",
"azure_password",
"azure_scope",
@ -678,6 +682,10 @@ NAMES_SHARED_WITH_TYPED_MODELS: Final[Mapping[str, tuple[str, ...]]] = MappingPr
"bedrock_tags",
"client_id",
"client_secret",
"token_exchange_audience",
"token_exchange_endpoint",
"token_exchange_profile",
"token_exchange_scope",
"openai_identity_provider_id",
"openai_identity_token_file",
"openai_service_account_id",

View file

@ -19,6 +19,7 @@ from litellm.types.utils import (
CustomPricingLiteLLMParams,
MirroredPricingParams,
anthropic_wif_litellm_params,
oauth_token_exchange_litellm_params,
openai_wif_litellm_params,
server_owned_wif_litellm_params,
)
@ -325,8 +326,16 @@ def test_openai_wif_fields_round_trip_through_model_dump():
assert dumped[field] == value, field
def test_server_owned_registry_is_anthropic_plus_openai():
assert server_owned_wif_litellm_params == anthropic_wif_litellm_params + openai_wif_litellm_params
def test_server_owned_registry_includes_anthropic_openai_and_oauth_token_exchange():
assert oauth_token_exchange_litellm_params == (
"token_exchange_audience",
"token_exchange_endpoint",
"token_exchange_profile",
"token_exchange_scope",
)
assert server_owned_wif_litellm_params == (
anthropic_wif_litellm_params + openai_wif_litellm_params + oauth_token_exchange_litellm_params
)
assert set(openai_wif_litellm_params) == {
"openai_identity_provider_id",
"openai_service_account_id",
@ -334,6 +343,11 @@ def test_server_owned_registry_is_anthropic_plus_openai():
}
def test_credential_litellm_params_declares_each_oauth_token_exchange_field():
for field in oauth_token_exchange_litellm_params:
assert field in CredentialLiteLLMParams.model_fields, field
def test_server_owned_wif_fields_present_reports_openai_fields():
assert server_owned_wif_fields_present(
{"openai_identity_token_file": "/var/run/secrets/tokens/openai", "model": "gpt-4o"}

View file

@ -1,12 +1,22 @@
import { fireEvent, renderHook, screen, waitFor, within, renderWithProviders } from "../../../tests/test-utils";
import {
chooseSelectOption,
fireEvent,
renderHook,
screen,
waitFor,
within,
renderWithProviders,
} from "../../../tests/test-utils";
import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event";
import { describe, expect, it, vi } from "vitest";
import type { Team } from "../key_team_helpers/key_list";
import { credentialCreateCall, type CredentialItem } from "../networking";
import { credentialCreateCall, modelCreateCall, type CredentialItem } from "../networking";
import { Providers } from "../provider_info_helpers";
import { projectMountedValues, useMountRegistry, type MountedFormValues } from "../common_components/MountedFormField";
import { useForm } from "react-hook-form";
import AddModelForm from "./AddModelForm";
import { handleAddModelSubmit } from "./handle_add_model_submit";
import { toast } from "@/lib/toast";
vi.mock("../molecules/models/ProviderLogo", () => ({
ProviderLogo: ({ provider, className }: { provider: string; className?: string }) => (
@ -35,6 +45,7 @@ vi.mock("../networking", async () => {
}),
testConnectionRequest: vi.fn().mockResolvedValue({ status: "success" }),
credentialCreateCall: vi.fn().mockResolvedValue({ success: true }),
modelCreateCall: vi.fn().mockResolvedValue({ success: true }),
getProviderCreateMetadata: vi.fn().mockResolvedValue([
{
provider: "OpenAI",
@ -43,6 +54,27 @@ vi.mock("../networking", async () => {
default_model_placeholder: "gpt-3.5-turbo",
credential_fields: [],
},
{
provider: "MICROSOFT_365_COPILOT",
provider_display_name: "Microsoft 365 Copilot",
litellm_provider: "microsoft_365_copilot",
default_model_placeholder: "microsoft_365_copilot/chat",
credential_fields: [
{ key: "token_exchange_endpoint", label: "Token Endpoint URL", field_type: "text" },
{
key: "token_exchange_profile",
label: "Exchange Grant",
field_type: "select",
options: ["jwt_bearer_obo", "rfc8693"],
default_value: "jwt_bearer_obo",
},
{ key: "client_id", label: "Client ID", field_type: "text" },
{ key: "client_secret", label: "Client Secret", field_type: "password" },
{ key: "token_exchange_scope", label: "Scope", field_type: "text" },
{ key: "token_exchange_audience", label: "Audience", field_type: "text" },
{ key: "api_key", label: "Delegated Access Token", field_type: "password" },
],
},
]),
};
});
@ -57,6 +89,27 @@ vi.mock("@/app/(dashboard)/hooks/providers/useProviderFields", () => ({
default_model_placeholder: "gpt-3.5-turbo",
credential_fields: [],
},
{
provider: "MICROSOFT_365_COPILOT",
provider_display_name: "Microsoft 365 Copilot",
litellm_provider: "microsoft_365_copilot",
default_model_placeholder: "microsoft_365_copilot/chat",
credential_fields: [
{ key: "token_exchange_endpoint", label: "Token Endpoint URL", field_type: "text" },
{
key: "token_exchange_profile",
label: "Exchange Grant",
field_type: "select",
options: ["jwt_bearer_obo", "rfc8693"],
default_value: "jwt_bearer_obo",
},
{ key: "client_id", label: "Client ID", field_type: "text" },
{ key: "client_secret", label: "Client Secret", field_type: "password" },
{ key: "token_exchange_scope", label: "Scope", field_type: "text" },
{ key: "token_exchange_audience", label: "Audience", field_type: "text" },
{ key: "api_key", label: "Delegated Access Token", field_type: "password" },
],
},
],
isLoading: false,
error: null,
@ -501,6 +554,128 @@ describe("AddModelForm", () => {
});
});
describe("credential-only provider auth types", () => {
const renderAsAdmin = async () => {
const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized"));
mockUseAuthorized.default.mockReturnValue(mockAuthorizedUser("proxy_admin", "user-1", true));
const props = createTestProps();
props.selectedProvider = Providers.MICROSOFT_365_COPILOT;
renderWithProviders(<AddModelForm {...props} />);
await screen.findByText("Existing Credentials");
return props;
};
const submitModel = async (props: ReturnType<typeof createTestProps>) => {
const credentialName = props.form.getValues("litellm_credential_name");
const apiKey = props.form.getValues("api_key");
const modelValues = {
...(credentialName ? { litellm_credential_name: credentialName } : {}),
...(apiKey ? { api_key: apiKey } : {}),
custom_llm_provider: "MICROSOFT_365_COPILOT",
model_mappings: [
{
public_name: "copilot-model",
litellm_model: "microsoft_365_copilot/chat",
},
],
};
await handleAddModelSubmit(modelValues, "test-access-token", { resetFields: vi.fn() });
};
it("creates and attaches an OAuth exchange credential without mounting its fields on the model", async () => {
const user = userEvent.setup();
vi.mocked(credentialCreateCall).mockClear();
vi.mocked(modelCreateCall).mockClear();
const props = await renderAsAdmin();
await user.click(screen.getByRole("button", { name: "Create credential" }));
const dialog = await screen.findByRole("dialog");
expect(within(dialog).getByPlaceholderText("Select a provider")).toHaveValue(Providers.MICROSOFT_365_COPILOT);
expect(within(dialog).getByRole("combobox", { name: "Auth Type:" })).toHaveTextContent(
"OAuth token exchange (on-behalf-of)",
);
await chooseSelectOption(
user,
within(dialog).getByRole("combobox", { name: "Exchange Grant" }),
"jwt_bearer_obo",
);
fireEvent.change(within(dialog).getByLabelText("Credential Name:"), { target: { value: "m365-obo" } });
fireEvent.change(within(dialog).getByLabelText("Token Endpoint URL"), {
target: { value: "https://login.example.test/token" },
});
fireEvent.change(within(dialog).getByLabelText("Client ID"), { target: { value: "client-id" } });
fireEvent.change(within(dialog).getByLabelText("Client Secret"), { target: { value: "client-secret" } });
fireEvent.change(within(dialog).getByLabelText("Scope"), { target: { value: "scope" } });
fireEvent.change(within(dialog).getByLabelText("Audience"), { target: { value: "audience" } });
await user.click(within(dialog).getByRole("button", { name: "Add Credential" }));
await waitFor(() =>
expect(credentialCreateCall).toHaveBeenCalledWith("test-access-token", {
credential_name: "m365-obo",
credential_values: {
token_exchange_endpoint: "https://login.example.test/token",
token_exchange_profile: "jwt_bearer_obo",
client_id: "client-id",
client_secret: "client-secret",
token_exchange_scope: "scope",
token_exchange_audience: "audience",
},
credential_info: { custom_llm_provider: Providers.MICROSOFT_365_COPILOT },
}),
);
expect(props.form.getValues("litellm_credential_name")).toBe("m365-obo");
expect(props.mountedValues()).not.toHaveProperty("token_exchange_endpoint");
props.handleOk.mockImplementation(async () => {
await submitModel(props);
return true;
});
await user.click(screen.getByRole("button", { name: "Add Model" }));
expect(props.handleOk).toHaveBeenCalledOnce();
await waitFor(() => expect(modelCreateCall).toHaveBeenCalledOnce());
const modelPayload = vi.mocked(modelCreateCall).mock.calls[0][1];
expect(modelPayload.litellm_params).toMatchObject({ litellm_credential_name: "m365-obo" });
expect(Object.keys(modelPayload.litellm_params).some((key) => key.startsWith("token_exchange_"))).toBe(false);
});
it("blocks model creation when OAuth exchange has no credential", async () => {
const props = await renderAsAdmin();
const errorToast = vi.spyOn(toast, "error");
await userEvent.click(screen.getByRole("button", { name: "Add Model" }));
expect(props.handleOk).not.toHaveBeenCalled();
expect(errorToast).toHaveBeenCalledWith("Create or select a credential for this auth type");
errorToast.mockRestore();
});
it("keeps static delegated tokens inline in the model request", async () => {
const user = userEvent.setup();
vi.mocked(modelCreateCall).mockClear();
const props = await renderAsAdmin();
await chooseSelectOption(
user,
screen.getByRole("combobox", { name: "Auth Type:" }),
"Static delegated access token",
);
await waitFor(() =>
expect(screen.getByRole("combobox", { name: "Auth Type:" })).toHaveTextContent("Static delegated access token"),
);
fireEvent.change(screen.getByLabelText("Delegated Access Token"), { target: { value: "delegated-token" } });
props.handleOk.mockImplementation(async () => {
await submitModel(props);
return true;
});
await user.click(screen.getByRole("button", { name: "Add Model" }));
expect(props.handleOk).toHaveBeenCalledOnce();
await waitFor(() => expect(modelCreateCall).toHaveBeenCalledOnce());
expect(vi.mocked(modelCreateCall).mock.calls[0][1].litellm_params).toMatchObject({
api_key: "delegated-token",
});
});
});
describe("cache control bindings reach the parent form store", () => {
const renderWithForm = async () => {
const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized"));

View file

@ -36,6 +36,7 @@ import ConditionalPublicModelName from "./conditional_public_model_name";
import LiteLLMModelNameField from "./litellm_model_name";
import ConnectionErrorDisplay from "./model_connection_test";
import ProviderSpecificFields from "./provider_specific_fields";
import { authTypesFor } from "./provider_auth_types";
import { TEST_MODES } from "./add_model_modes";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { credentialsKeys } from "@/app/(dashboard)/hooks/credentials/useCredentials";
@ -99,11 +100,16 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
const { data: tagsList } = useTags();
const selectedCredentialName = useWatch({ control: form.control, name: "litellm_credential_name" });
const queryClient = useQueryClient();
const [isCredentialModalOpen, setIsCredentialModalOpen] = useState(false);
const [credentialModalAuthTypeId, setCredentialModalAuthTypeId] = useState<string | undefined>();
const [selectedAuthTypeId, setSelectedAuthTypeId] = useState("");
const canCreateCredential = isProxyAdminRole(userRole ?? "");
const [isFederatedCredentialModalOpen, setIsFederatedCredentialModalOpen] = useState(false);
const canCreateFederatedCredential =
isProxyAdminRole(userRole ?? "") && federatedProviderOf(selectedProvider) !== null;
const selectedAuthType = authTypesFor(selectedProvider).find(({ id }) => id === selectedAuthTypeId);
const handleCreateFederatedCredential = async (values: Record<string, unknown>) => {
const handleCreateCredential = async (values: Record<string, unknown>) => {
const credential = buildCredential(values, withoutRestrictedFields(values));
try {
await credentialCreateCall(accessToken, credential);
@ -112,11 +118,25 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
return;
}
toast.success("Credential added successfully");
setIsCredentialModalOpen(false);
setIsFederatedCredentialModalOpen(false);
await queryClient.invalidateQueries({ queryKey: credentialsKeys.all });
form.setValue("litellm_credential_name", credential.credential_name, { shouldDirty: true });
};
const openCredentialModal = (authTypeId?: string) => {
setCredentialModalAuthTypeId(authTypeId);
setIsCredentialModalOpen(true);
};
const handleSubmit = async () => {
if (selectedAuthType?.credentialOnly && !selectedCredentialName) {
toast.error("Create or select a credential for this auth type");
return false;
}
return handleOk();
};
const handleTestConnection = async () => {
setIsTestingConnection(true);
setConnectionTestId(`test-${Date.now()}`);
@ -188,7 +208,7 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
<form
onSubmit={(event) => {
event.preventDefault();
void handleOk().then((submitted) => {
void handleSubmit().then((submitted) => {
if (submitted) {
setTeamAdminSelectedTeam(null);
}
@ -331,7 +351,12 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
<span className="px-4 text-muted-foreground text-sm">OR</span>
<div className="grow border-t border-border"></div>
</div>
<ProviderSpecificFields selectedProvider={selectedProvider} />
<ProviderSpecificFields
selectedProvider={selectedProvider}
context="model"
onAuthTypeChange={setSelectedAuthTypeId}
onCreateCredential={canCreateCredential ? openCredentialModal : undefined}
/>
{canCreateFederatedCredential && (
<div className="mb-4 flex flex-col items-start gap-2">
<span className="text-sm text-muted-foreground">
@ -476,6 +501,17 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
</CardContent>
</Card>
{isCredentialModalOpen && (
<CredentialModal
open
mode="add"
initialProvider={selectedProvider}
initialAuthTypeId={credentialModalAuthTypeId}
providerLocked
onCancel={() => setIsCredentialModalOpen(false)}
onSubmit={handleCreateCredential}
/>
)}
{isFederatedCredentialModalOpen && (
<CredentialModal
open
@ -484,7 +520,7 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
initialAuthMethod="federation"
providerLocked
onCancel={() => setIsFederatedCredentialModalOpen(false)}
onSubmit={handleCreateFederatedCredential}
onSubmit={handleCreateCredential}
/>
)}
{/* Test Connection Results Modal */}

View file

@ -0,0 +1,81 @@
import { describe, expect, it } from "vitest";
import {
PROVIDER_AUTH_TYPES,
authTypeFieldKeys,
authTypesFor,
hiddenAuthFieldKeys,
inferAuthTypeId,
type ProviderAuthType,
} from "./provider_auth_types";
import { Providers } from "../provider_info_helpers";
describe("provider auth types", () => {
it("resolves provider enum keys and display names", () => {
expect(authTypesFor("MICROSOFT_365_COPILOT")).toEqual(PROVIDER_AUTH_TYPES.MICROSOFT_365_COPILOT);
expect(authTypesFor(Providers.MICROSOFT_365_COPILOT)).toEqual(PROVIDER_AUTH_TYPES.MICROSOFT_365_COPILOT);
expect(authTypesFor(null)).toEqual([]);
expect(authTypesFor("unknown")).toEqual([]);
});
it("returns the unique field keys for every auth type", () => {
expect(authTypeFieldKeys("MICROSOFT_365_COPILOT")).toEqual([
"token_exchange_endpoint",
"token_exchange_profile",
"client_id",
"client_secret",
"token_exchange_scope",
"token_exchange_audience",
"api_key",
]);
});
it("infers auth types from required fields and falls back when only optional defaults are present", () => {
const authTypes = authTypesFor("MICROSOFT_365_COPILOT");
const delegatedValues: Record<string, unknown> = {
token_exchange_profile: "jwt_bearer_obo",
token_exchange_scope: "https://graph.microsoft.com/.default",
api_key: "delegated-token",
};
const emptyValues: Record<string, unknown> = {
token_exchange_endpoint: "",
client_id: null,
client_secret: undefined,
token_exchange_profile: "jwt_bearer_obo",
token_exchange_scope: "https://graph.microsoft.com/.default",
api_key: "",
};
expect(inferAuthTypeId(authTypes, delegatedValues)).toBe("static_token");
expect(inferAuthTypeId(authTypes, emptyValues)).toBe("oauth_token_exchange");
expect(inferAuthTypeId([], {})).toBe("");
});
it("hides only non-selected keys not shared with the selected type", () => {
const authTypes: readonly ProviderAuthType[] = [
{
id: "first",
label: "First",
description: "First auth type",
fieldKeys: ["shared", "first_key"],
requiredFieldKeys: ["first_key"],
},
{
id: "second",
label: "Second",
description: "Second auth type",
fieldKeys: ["shared", "second_key"],
requiredFieldKeys: ["second_key"],
},
{
id: "third",
label: "Third",
description: "Third auth type",
fieldKeys: ["third_key"],
requiredFieldKeys: ["third_key"],
},
];
expect(hiddenAuthFieldKeys(authTypes, "first")).toEqual(["second_key", "third_key"]);
});
});

View file

@ -0,0 +1,76 @@
import { Providers } from "../provider_info_helpers";
export interface ProviderAuthType {
readonly id: string;
readonly label: string;
readonly description: string;
readonly fieldKeys: readonly string[];
readonly requiredFieldKeys: readonly string[];
readonly fixedValues?: Readonly<Record<string, string>>;
readonly credentialOnly?: boolean;
}
const EMPTY_PROVIDER_AUTH_TYPES: readonly ProviderAuthType[] = [];
export const PROVIDER_AUTH_TYPES: Partial<Record<keyof typeof Providers, readonly ProviderAuthType[]>> = {
MICROSOFT_365_COPILOT: [
{
id: "oauth_token_exchange",
label: "OAuth token exchange (on-behalf-of)",
credentialOnly: true,
description:
"LiteLLM exchanges each caller's IdP-issued JWT at your IdP's token endpoint for a delegated token. Microsoft Graph only accepts Microsoft Entra tokens, so for Microsoft 365 Copilot the token endpoint must be Entra and callers must send an Entra-issued JWT for this app. Users can still sign in through any IdP federated with Entra.",
fieldKeys: [
"token_exchange_endpoint",
"token_exchange_profile",
"client_id",
"client_secret",
"token_exchange_scope",
"token_exchange_audience",
],
requiredFieldKeys: ["token_exchange_endpoint", "client_id", "client_secret"],
},
{
id: "static_token",
label: "Static delegated access token",
description: "Sends one pre-acquired delegated Microsoft Graph token for every caller.",
fieldKeys: ["api_key"],
requiredFieldKeys: ["api_key"],
},
],
};
export const authTypesFor = (provider: string | null): readonly ProviderAuthType[] => {
const providerKeys: readonly (keyof typeof Providers)[] = Object.keys(Providers) as (keyof typeof Providers)[];
const providerKey =
provider === null ? undefined : providerKeys.find((key) => key === provider || Providers[key] === provider);
return providerKey === undefined
? EMPTY_PROVIDER_AUTH_TYPES
: PROVIDER_AUTH_TYPES[providerKey] ?? EMPTY_PROVIDER_AUTH_TYPES;
};
export const authTypeFieldKeys = (provider: string | null): readonly string[] =>
authTypesFor(provider)
.flatMap(({ fieldKeys }) => fieldKeys)
.filter((fieldKey, index, fieldKeys) => fieldKeys.indexOf(fieldKey) === index);
export const inferAuthTypeId = (authTypes: readonly ProviderAuthType[], values: Record<string, unknown>): string =>
authTypes.find(({ requiredFieldKeys }) =>
requiredFieldKeys.some(
(fieldKey) => values[fieldKey] !== undefined && values[fieldKey] !== null && values[fieldKey] !== "",
),
)?.id ??
authTypes[0]?.id ??
"";
export const hiddenAuthFieldKeys = (authTypes: readonly ProviderAuthType[], selectedId: string): readonly string[] => {
const selectedFieldKeys = authTypes.find(({ id }) => id === selectedId)?.fieldKeys ?? [];
return authTypes
.filter(({ id }) => id !== selectedId)
.flatMap(({ fieldKeys }) => fieldKeys)
.filter(
(fieldKey, index, fieldKeys) => !selectedFieldKeys.includes(fieldKey) && fieldKeys.indexOf(fieldKey) === index,
);
};

View file

@ -1,7 +1,10 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import type { ReactNode } from "react";
import { beforeAll, describe, expect, it, vi } from "vitest";
import { useFormContext } from "react-hook-form";
import { chooseSelectOption } from "@/../tests/test-utils";
import { Providers } from "../provider_info_helpers";
import type { MountedFormValues } from "../common_components/MountedFormField";
import { MountedFormHost } from "../../../tests/mounted-form-host";
@ -109,6 +112,41 @@ vi.mock("../networking", async () => {
},
],
},
{
provider: "MICROSOFT_365_COPILOT",
provider_display_name: Providers.MICROSOFT_365_COPILOT,
litellm_provider: "microsoft_365_copilot",
default_model_placeholder: "microsoft_365_copilot/chat",
credential_fields: [
{
key: "token_exchange_endpoint",
label: "Token Endpoint URL",
placeholder: "https://login.microsoftonline.com/<tenant-id>/oauth2/v2.0/token",
field_type: "text",
},
{
key: "token_exchange_profile",
label: "Exchange Grant",
field_type: "select",
options: ["jwt_bearer_obo", "rfc8693"],
default_value: "jwt_bearer_obo",
},
{ key: "client_id", label: "Client ID", field_type: "text" },
{ key: "client_secret", label: "Client Secret", field_type: "password" },
{
key: "token_exchange_scope",
label: "Scope",
field_type: "text",
default_value: "https://graph.microsoft.com/.default",
},
{
key: "token_exchange_audience",
label: "Audience",
field_type: "text",
},
{ key: "api_key", label: "Delegated Access Token", field_type: "password" },
],
},
]),
};
});
@ -144,6 +182,16 @@ const VertexCredentialsProbe = () => {
return <output data-testid="vertex-credentials">{String(watch("vertex_credentials") ?? "")}</output>;
};
const ValidationForm = ({ children }: { readonly children: ReactNode }) => {
const form = useFormContext<MountedFormValues>();
return (
<form onSubmit={form.handleSubmit(() => {})}>
{children}
<button type="submit">Submit</button>
</form>
);
};
describe("ProviderSpecificFields", () => {
it("reads a picked service-account file into the vertex credentials field", async () => {
const queryClient = createQueryClient();
@ -212,6 +260,7 @@ describe("ProviderSpecificFields", () => {
const apiKeyLabel = await screen.findByLabelText("OpenAI API Key");
expect(apiKeyLabel).toBeInTheDocument();
expect(screen.queryByLabelText("Auth Type:")).not.toBeInTheDocument();
const apiBaseInput = screen.getByPlaceholderText("https://api.openai.com/v1");
expect(apiBaseInput).toBeInTheDocument();
@ -419,4 +468,192 @@ describe("ProviderSpecificFields", () => {
expect(apiVersionInput).toHaveValue("2025-01-01-preview");
});
});
it("defaults Microsoft 365 Copilot to OAuth token exchange fields", async () => {
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
</MountedFormHost>
</QueryClientProvider>,
);
expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent(
"OAuth token exchange (on-behalf-of)",
);
expect(screen.getByLabelText("Token Endpoint URL")).toBeInTheDocument();
expect(screen.getByLabelText("Client ID")).toBeInTheDocument();
expect(screen.getByLabelText("Client Secret")).toBeInTheDocument();
expect(screen.queryByLabelText("Delegated Access Token")).not.toBeInTheDocument();
});
it("offers a credential-only OAuth mode in model context", async () => {
const onCreateCredential = vi.fn();
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ProviderSpecificFields
selectedProvider="MICROSOFT_365_COPILOT"
context="model"
onCreateCredential={onCreateCredential}
/>
</MountedFormHost>
</QueryClientProvider>,
);
expect(
await screen.findByText("This auth type is saved as an LLM credential and attached to this model."),
).toBeInTheDocument();
expect(screen.queryByLabelText("Token Endpoint URL")).not.toBeInTheDocument();
await userEvent.click(screen.getByRole("button", { name: "Create credential" }));
expect(onCreateCredential).toHaveBeenCalledWith("oauth_token_exchange");
});
it("hides credential-only auth modes from non-admin model forms", async () => {
const user = userEvent.setup();
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" context="model" />
</MountedFormHost>
</QueryClientProvider>,
);
const authTypeSelect = await screen.findByRole("combobox", { name: "Auth Type:" });
expect(authTypeSelect).toHaveTextContent("Static delegated access token");
await user.click(authTypeSelect);
expect((await screen.findAllByRole("option")).map((option) => option.textContent)).toEqual([
"Static delegated access token",
]);
expect(screen.queryByLabelText("Token Endpoint URL")).not.toBeInTheDocument();
expect(screen.queryByLabelText("Client ID")).not.toBeInTheDocument();
expect(screen.queryByLabelText("Client Secret")).not.toBeInTheDocument();
expect(screen.getByLabelText("Delegated Access Token")).toBeInTheDocument();
});
it("switches Microsoft 365 Copilot fields to a static delegated access token", async () => {
const user = userEvent.setup();
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
</MountedFormHost>
</QueryClientProvider>,
);
await screen.findByLabelText("Token Endpoint URL");
await chooseSelectOption(
user,
screen.getByRole("combobox", { name: "Auth Type:" }),
"Static delegated access token",
);
expect(await screen.findByLabelText("Delegated Access Token")).toBeInTheDocument();
expect(screen.queryByLabelText("Token Endpoint URL")).not.toBeInTheDocument();
expect(screen.queryByLabelText("Client ID")).not.toBeInTheDocument();
expect(screen.queryByLabelText("Client Secret")).not.toBeInTheDocument();
expect(screen.queryByLabelText("Scope")).not.toBeInTheDocument();
expect(screen.queryByLabelText("Audience")).not.toBeInTheDocument();
});
it("infers a static delegated token from a preloaded api_key", async () => {
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<MountedFormHost defaultValues={{ api_key: "delegated-token" }}>
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
</MountedFormHost>
</QueryClientProvider>,
);
expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent(
"Static delegated access token",
);
expect(screen.getByLabelText("Delegated Access Token")).toHaveValue("delegated-token");
expect(screen.queryByLabelText("Token Endpoint URL")).not.toBeInTheDocument();
});
it("re-infers the auth type when the selected provider changes", async () => {
const user = userEvent.setup();
const queryClient = createQueryClient();
const { rerender } = render(
<QueryClientProvider client={queryClient}>
<MountedFormHost defaultValues={{ api_key: "delegated-token" }}>
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
</MountedFormHost>
</QueryClientProvider>,
);
await chooseSelectOption(
user,
await screen.findByRole("combobox", { name: "Auth Type:" }),
"OAuth token exchange (on-behalf-of)",
);
rerender(
<QueryClientProvider client={queryClient}>
<MountedFormHost defaultValues={{ api_key: "delegated-token" }}>
<ProviderSpecificFields selectedProvider="OpenAI" />
</MountedFormHost>
</QueryClientProvider>,
);
await screen.findByLabelText("OpenAI API Key");
rerender(
<QueryClientProvider client={queryClient}>
<MountedFormHost defaultValues={{ api_key: "delegated-token" }}>
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
</MountedFormHost>
</QueryClientProvider>,
);
expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent(
"Static delegated access token",
);
});
it("requires only the required OAuth token exchange fields", async () => {
const user = userEvent.setup();
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ValidationForm>
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
</ValidationForm>
</MountedFormHost>
</QueryClientProvider>,
);
const tokenEndpoint = await screen.findByLabelText("Token Endpoint URL");
const clientId = screen.getByLabelText("Client ID");
const clientSecret = screen.getByLabelText("Client Secret");
expect(tokenEndpoint).toBeInTheDocument();
expect(clientId).toBeInTheDocument();
expect(clientSecret).toBeInTheDocument();
await user.click(screen.getByRole("button", { name: "Submit" }));
expect(await screen.findAllByText("Required")).toHaveLength(3);
});
it("does not render an Auth Type selector for OpenAI", async () => {
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ProviderSpecificFields selectedProvider="OpenAI" />
</MountedFormHost>
</QueryClientProvider>,
);
await screen.findByLabelText("OpenAI API Key");
expect(screen.queryByRole("combobox", { name: "Auth Type:" })).not.toBeInTheDocument();
});
});

View file

@ -13,6 +13,7 @@ import {
type MountedFieldControlProps,
type MountedFormValues,
} from "../common_components/MountedFormField";
import { authTypesFor, hiddenAuthFieldKeys, inferAuthTypeId } from "./provider_auth_types";
import { ProviderCredentialFieldMetadata } from "../networking";
import { Providers } from "../provider_info_helpers";
import { labelWithHint } from "@/components/shared/form/LabelWithHint";
@ -21,6 +22,10 @@ import type { ProviderFieldValidators } from "../model_add/credential_federation
interface ProviderSpecificFieldsProps {
selectedProvider: string | null;
hiddenFieldKeys?: readonly string[];
context?: "model" | "credential";
onCreateCredential?: (authTypeId: string) => void;
initialAuthTypeId?: string;
onAuthTypeChange?: (authTypeId: string) => void;
fieldValidators?: ProviderFieldValidators;
}
@ -83,14 +88,40 @@ const mapFieldMetadataToUiField = (field: ProviderCredentialFieldMetadata): Prov
const providerFieldsByDisplayName: Record<string, ProviderCredentialField[]> = {};
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
const ProviderSpecificFieldsContent: React.FC<ProviderSpecificFieldsProps> = ({
selectedProvider,
hiddenFieldKeys,
context = "credential",
onCreateCredential,
initialAuthTypeId,
onAuthTypeChange,
fieldValidators,
}) => {
const selectedProviderEnum = Providers[selectedProvider as keyof typeof Providers] as Providers;
const form = useFormContext<MountedFormValues>();
const authTypes = authTypesFor(selectedProvider);
const availableAuthTypes = React.useMemo(
() =>
authTypes.filter(
({ credentialOnly }) => context !== "model" || onCreateCredential !== undefined || !credentialOnly,
),
[authTypes, context, onCreateCredential],
);
const [selectedAuthTypeId, setSelectedAuthTypeId] = React.useState(() =>
initialAuthTypeId && availableAuthTypes.some(({ id }) => id === initialAuthTypeId)
? initialAuthTypeId
: inferAuthTypeId(availableAuthTypes, form.getValues()),
);
const selectedAuthType = availableAuthTypes.find(({ id }) => id === selectedAuthTypeId) ?? availableAuthTypes[0];
const credentialOnlySelected = context === "model" && selectedAuthType?.credentialOnly === true;
const authTypeSelectId = React.useId();
const credentialsFileRef = React.useRef<HTMLInputElement>(null);
React.useEffect(() => {
if (selectedAuthType) {
onAuthTypeChange?.(selectedAuthType.id);
}
}, [onAuthTypeChange, selectedAuthType]);
const pickCredentialsFile =
(onLoaded: (contents: string) => void) => (event: React.ChangeEvent<HTMLInputElement>) => {
const file = event.target.files?.[0];
@ -172,11 +203,34 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
return mapped;
}, [selectedProviderEnum, selectedProvider, providerMetadata]);
const allFields = React.useMemo(
() => (hiddenFieldKeys ? providerFields.filter((field) => !hiddenFieldKeys.includes(field.key)) : providerFields),
[providerFields, hiddenFieldKeys],
const authTypeHiddenFieldKeys = React.useMemo(
() => hiddenAuthFieldKeys(authTypes, selectedAuthType?.id ?? ""),
[authTypes, selectedAuthType?.id],
);
const allFields = React.useMemo(() => {
if (credentialOnlySelected) {
return [];
}
if (availableAuthTypes.length === 0) {
return hiddenFieldKeys ? providerFields.filter((field) => !hiddenFieldKeys.includes(field.key)) : providerFields;
}
const fieldsToHide = [...(hiddenFieldKeys ?? []), ...authTypeHiddenFieldKeys];
return providerFields
.filter((field) => !fieldsToHide.includes(field.key))
.map((field) => (selectedAuthType?.requiredFieldKeys.includes(field.key) ? { ...field, required: true } : field));
}, [
providerFields,
hiddenFieldKeys,
availableAuthTypes,
authTypeHiddenFieldKeys,
selectedAuthType,
credentialOnlySelected,
]);
const hasApiVersionField = React.useMemo(() => allFields.some((field) => field.key === "api_version"), [allFields]);
const lastInferredApiVersionRef = React.useRef<string | null>(null);
@ -291,6 +345,44 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
return (
<>
{availableAuthTypes.length > 0 && selectedAuthType && (
<div className="mb-4 flex flex-col gap-2">
<label htmlFor={authTypeSelectId} className="text-sm font-medium">
{labelWithHint("Auth Type:", "Select how LiteLLM authenticates to this provider.")}
</label>
<Select
items={availableAuthTypes.map(({ id, label }) => ({ value: id, label }))}
value={selectedAuthType.id}
onValueChange={(value) => setSelectedAuthTypeId(value ?? "")}
>
<SelectTrigger id={authTypeSelectId} className="w-full">
<SelectValue />
</SelectTrigger>
<SelectContent>
{availableAuthTypes.map((authType) => (
<SelectItem key={authType.id} value={authType.id}>
{authType.label}
</SelectItem>
))}
</SelectContent>
</Select>
<p className="text-sm text-muted-foreground">{selectedAuthType.description}</p>
</div>
)}
{credentialOnlySelected ? (
<>
<p className="mb-4 text-sm">This auth type is saved as an LLM credential and attached to this model.</p>
<Button type="button" onClick={() => onCreateCredential?.(selectedAuthType.id)}>
Create credential
</Button>
</>
) : (
Object.entries(selectedAuthType?.fixedValues ?? {}).map(([fieldKey, fixedValue]) => (
<MountedFormField key={`${selectedProvider}:${fieldKey}`} name={fieldKey} defaultValue={fixedValue} bare>
{() => null}
</MountedFormField>
))
)}
{isLoading && allFields.length === 0 && <p className="text-sm mb-2">Loading provider fields...</p>}
{loadError && allFields.length === 0 && (
<p className="text-sm mb-2 text-destructive">
@ -341,4 +433,8 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
);
};
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = (props) => (
<ProviderSpecificFieldsContent key={props.selectedProvider ?? "no-provider"} {...props} />
);
export default ProviderSpecificFields;

View file

@ -40,6 +40,26 @@ vi.mock("../networking", async () => {
{ key: "api_key", label: "Azure API Key", field_type: "password" },
],
},
{
provider: "MICROSOFT_365_COPILOT",
provider_display_name: Providers.MICROSOFT_365_COPILOT,
litellm_provider: "microsoft_365_copilot",
credential_fields: [
{ key: "token_exchange_endpoint", label: "Token Endpoint URL", field_type: "text" },
{
key: "token_exchange_profile",
label: "Exchange Grant",
field_type: "select",
options: ["jwt_bearer_obo", "rfc8693"],
default_value: "jwt_bearer_obo",
},
{ key: "client_id", label: "Client ID", field_type: "text" },
{ key: "client_secret", label: "Client Secret", field_type: "password" },
{ key: "token_exchange_scope", label: "Scope", field_type: "text" },
{ key: "token_exchange_audience", label: "Audience", field_type: "text" },
{ key: "api_key", label: "Delegated Access Token", field_type: "password" },
],
},
]),
};
});
@ -84,6 +104,19 @@ const unknownSourceCredential: CredentialItem = {
credential_info: { custom_llm_provider: "anthropic" },
};
const microsoftCopilotCredential: CredentialItem = {
credential_name: "microsoft-copilot",
credential_values: {
token_exchange_endpoint: "https://identity.example.com/stored-token",
token_exchange_profile: "jwt_bearer_obo",
client_id: "stored-client",
client_secret: "stored-secret",
token_exchange_scope: "https://graph.microsoft.com/.default",
token_exchange_audience: "https://graph.microsoft.com",
},
credential_info: { custom_llm_provider: "MICROSOFT_365_COPILOT" },
};
const renderModal = (props: Partial<React.ComponentProps<typeof CredentialModal>> = {}) => {
const onSubmit = vi.fn();
render(
@ -109,6 +142,34 @@ const chooseProvider = async (user: ReturnType<typeof userEvent.setup>, provider
};
describe("CredentialModal with Anthropic workload identity federation", () => {
it("deletes stored exchange fields when switching to a static token", async () => {
const user = userEvent.setup();
const onSubmit = renderModal({ mode: "edit", existingCredential: microsoftCopilotCredential });
expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent(
"OAuth token exchange (on-behalf-of)",
);
await screen.findByLabelText("Token Endpoint URL");
await chooseOption(user, /^Auth Type:/, "Static delegated access token");
fill("Delegated Access Token", "delegated-graph-token");
await user.click(screen.getByRole("button", { name: "Update Credential" }));
const [values, valuesToDelete] = onSubmit.mock.calls[0];
expect(values).toEqual({
credential_name: "microsoft-copilot",
custom_llm_provider: "MICROSOFT_365_COPILOT",
api_key: "delegated-graph-token",
});
expect([...valuesToDelete].sort()).toEqual([
"client_id",
"client_secret",
"token_exchange_audience",
"token_exchange_endpoint",
"token_exchange_profile",
"token_exchange_scope",
]);
});
it("does not carry the previous provider's base URL into the Anthropic form or its payload", async () => {
const user = userEvent.setup();
const onSubmit = renderModal();

View file

@ -4,6 +4,7 @@ import { SimpleTooltip } from "@/components/ui/tooltip";
import { Button } from "@/components/ui/button";
import { useState } from "react";
import { FormProvider, useForm } from "react-hook-form";
import { authTypeFieldKeys } from "../add_model/provider_auth_types";
import ProviderSpecificFields from "../add_model/provider_specific_fields";
import { requiredRule } from "../common_components/formRules";
import { labelWithHint } from "@/components/shared/form/LabelWithHint";
@ -59,6 +60,7 @@ interface CredentialModalProps {
existingCredential?: CredentialItem | null;
initialProvider?: string | null;
initialAuthMethod?: AuthMethod;
initialAuthTypeId?: string;
providerLocked?: boolean;
}
@ -106,6 +108,7 @@ export default function CredentialModal({
existingCredential = null,
initialProvider = null,
initialAuthMethod,
initialAuthTypeId,
providerLocked = false,
}: CredentialModalProps) {
const isEdit = mode === "edit";
@ -163,10 +166,16 @@ export default function CredentialModal({
onSubmit({ ...meta, ...buildCreateCredentialValues(withoutRestrictedFields(values), selection) }, []);
return;
}
const patch = sameProvider(selectedProvider, storedProvider)
? buildCredentialPatch(storedValues, withoutRestrictedFields(values), storedSelection, selection)
: buildProviderChangePatch(storedValues, withoutRestrictedFields(values), selection);
onSubmit({ ...meta, ...patch.credential_values }, patch.credential_values_to_delete);
const sameStoredProvider = sameProvider(selectedProvider, storedProvider);
const projectedValues = withoutRestrictedFields(values);
const patch = sameStoredProvider
? buildCredentialPatch(storedValues, projectedValues, storedSelection, selection)
: buildProviderChangePatch(storedValues, projectedValues, selection);
const authTypeValuesToDelete = sameStoredProvider
? authTypeFieldKeys(selectedProvider).filter((key) => key in storedValues && !(key in projectedValues))
: [];
const valuesToDelete = Array.from(new Set([...patch.credential_values_to_delete, ...authTypeValuesToDelete]));
onSubmit({ ...meta, ...patch.credential_values }, valuesToDelete);
};
const closeAndReset = () => {
@ -283,6 +292,7 @@ export default function CredentialModal({
<ProviderSpecificFields
selectedProvider={selectedProvider}
hiddenFieldKeys={selection.authMethod === "federation" ? API_KEY_FIELDS : NO_HIDDEN_FIELDS}
initialAuthTypeId={initialAuthTypeId}
fieldValidators={providerFieldValidators(selection)}
/>

View file

@ -9,6 +9,11 @@ import {
} from "./provider_info_helpers";
describe("provider_info_helpers", () => {
it("maps Microsoft 365 Copilot to its chat model provider and placeholder", () => {
expect(provider_map.MICROSOFT_365_COPILOT).toBe("microsoft_365_copilot");
expect(getPlaceholder(Providers.MICROSOFT_365_COPILOT)).toBe("microsoft_365_copilot/chat");
});
describe("getProviderLogoAndName", () => {
it("should return empty logo and dash display name when providerValue is empty", () => {
const result = getProviderLogoAndName("");

View file

@ -132,6 +132,7 @@ export enum Providers {
LM_STUDIO = "Lm Studio",
LLAMA = "Meta Llama",
MARITALK = "Maritalk",
MICROSOFT_365_COPILOT = "Microsoft 365 Copilot",
MiniMax = "MiniMax",
MistralAI = "Mistral AI",
MOONSHOT = "Moonshot",
@ -252,6 +253,7 @@ export const provider_map: Record<string, string> = {
LLAMAFILE: "llamafile",
LLAMA: "meta_llama",
LM_STUDIO: "lm_studio",
MICROSOFT_365_COPILOT: "microsoft_365_copilot",
MARITALK: "maritalk",
MiniMax: "minimax",
MistralAI: "mistral",
@ -364,6 +366,7 @@ export const providerLogoMap: Partial<Record<Providers, string>> = {
[Providers.LM_STUDIO]: lmstudioLogo.src,
[Providers.LLAMA]: metaLlamaLogo.src,
[Providers.MiniMax]: minimaxLogo.src,
[Providers.MICROSOFT_365_COPILOT]: microsoftAzureLogo.src,
[Providers.MistralAI]: mistralLogo.src,
[Providers.MOONSHOT]: moonshotLogo.src,
[Providers.MORPH]: morphLogo.src,
@ -458,6 +461,7 @@ const providerPlaceholderMap: Partial<Record<Providers, string>> = {
[Providers.FalAI]: "fal_ai/fal-ai/flux-pro/v1.1-ultra",
[Providers.Google_AI_Studio]: "gemini-pro",
[Providers.JinaAI]: "jina_ai/",
[Providers.MICROSOFT_365_COPILOT]: "microsoft_365_copilot/chat",
[Providers.NVIDIA_RIVA]: "nvidia_riva/nvidia/parakeet-ctc-1_1b-asr",
[Providers.Oracle]: "oci/xai.grok-4",
[Providers.RunwayML]: "runwayml/gen4_turbo",

View file

@ -36336,6 +36336,14 @@ export interface components {
}[] | null;
/** Timeout */
timeout?: number | string | null;
/** Token Exchange Audience */
token_exchange_audience?: string | null;
/** Token Exchange Endpoint */
token_exchange_endpoint?: string | null;
/** Token Exchange Profile */
token_exchange_profile?: string | null;
/** Token Exchange Scope */
token_exchange_scope?: string | null;
/** Tpm */
tpm?: number | null;
/** Use Chat Completions Api */
@ -51664,6 +51672,14 @@ export interface components {
}[] | null;
/** Timeout */
timeout?: number | string | null;
/** Token Exchange Audience */
token_exchange_audience?: string | null;
/** Token Exchange Endpoint */
token_exchange_endpoint?: string | null;
/** Token Exchange Profile */
token_exchange_profile?: string | null;
/** Token Exchange Scope */
token_exchange_scope?: string | null;
/** Tpm */
tpm?: number | null;
/** Use Chat Completions Api */