From fbf4af301cfdc6bf0d861eb243682e3e9f1164c4 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Fri, 9 Oct 2026 15:19:07 -0700 Subject: [PATCH] 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> --- .github/workflows/test-unit.yml | 1 + litellm/__init__.py | 4 + litellm/caching/caching_handler.py | 52 +- litellm/constants.py | 8 +- .../litellm_core_utils/get_litellm_params.py | 9 + .../oauth_token_exchange.py | 307 ++++++++ .../llms/microsoft_365_copilot/__init__.py | 3 + .../microsoft_365_copilot/chat/__init__.py | 3 + .../microsoft_365_copilot/chat/handler.py | 580 +++++++++++++++ .../chat/transformation.py | 332 +++++++++ .../microsoft_365_copilot/common_utils.py | 47 ++ litellm/main.py | 56 ++ ...odel_prices_and_context_window_backup.json | 7 + litellm/proxy/auth/auth_utils.py | 3 +- .../common_utils/credential_hydration.py | 16 +- .../proxy/credential_endpoints/endpoints.py | 9 +- litellm/proxy/health_check.py | 6 +- .../health_endpoints/_health_endpoints.py | 60 +- .../model_management_endpoints.py | 2 +- .../provider_create_fields.json | 78 ++ litellm/router.py | 19 + .../clientside_credential_handler.py | 13 +- litellm/router_utils/cooldown_handlers.py | 24 + .../router_utils/fallback_event_handlers.py | 9 + litellm/types/litellm_params.py | 4 + litellm/types/router.py | 42 +- litellm/types/utils.py | 11 +- litellm/utils.py | 10 + model_prices_and_context_window.json | 7 + provider_endpoints_support.json | 18 + tests/unit/caching/test_caching_handler.py | 143 +++- .../test_oauth_token_exchange.py | 289 ++++++++ .../anthropic/test_anthropic_common_utils.py | 5 +- .../llms/microsoft_365_copilot/__init__.py | 0 .../microsoft_365_copilot/chat/__init__.py | 0 .../chat/test_handler.py | 677 ++++++++++++++++++ .../chat/test_transformation.py | 335 +++++++++ .../test_common_utils.py | 27 + tests/unit/proxy/auth/test_auth_utils.py | 130 +++- .../credential_endpoints/test_endpoints.py | 29 + .../health_endpoints/test_health_endpoints.py | 351 ++++++++- .../test_model_management_endpoints.py | 301 +++++++- .../proxy/test_credential_slot_registry.py | 4 + .../router_utils/test_cooldown_handlers.py | 181 +++++ .../test_fallback_event_handlers.py | 5 +- tests/unit/test_internal_context.py | 1 + tests/unit/types/test_litellm_params.py | 8 + tests/unit/types/test_router.py | 18 +- .../AddModelForm.integration.test.tsx | 179 ++++- .../src/components/add_model/AddModelForm.tsx | 44 +- .../add_model/provider_auth_types.test.ts | 81 +++ .../add_model/provider_auth_types.ts | 76 ++ .../provider_specific_fields.test.tsx | 237 ++++++ .../add_model/provider_specific_fields.tsx | 104 ++- .../CredentialModal.integration.test.tsx | 61 ++ .../components/model_add/CredentialModal.tsx | 18 +- .../components/provider_info_helpers.test.tsx | 5 + .../src/components/provider_info_helpers.tsx | 4 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 16 + 59 files changed, 4925 insertions(+), 144 deletions(-) create mode 100644 litellm/litellm_core_utils/oauth_token_exchange.py create mode 100644 litellm/llms/microsoft_365_copilot/__init__.py create mode 100644 litellm/llms/microsoft_365_copilot/chat/__init__.py create mode 100644 litellm/llms/microsoft_365_copilot/chat/handler.py create mode 100644 litellm/llms/microsoft_365_copilot/chat/transformation.py create mode 100644 litellm/llms/microsoft_365_copilot/common_utils.py create mode 100644 tests/unit/litellm_core_utils/test_oauth_token_exchange.py create mode 100644 tests/unit/llms/microsoft_365_copilot/__init__.py create mode 100644 tests/unit/llms/microsoft_365_copilot/chat/__init__.py create mode 100644 tests/unit/llms/microsoft_365_copilot/chat/test_handler.py create mode 100644 tests/unit/llms/microsoft_365_copilot/chat/test_transformation.py create mode 100644 tests/unit/llms/microsoft_365_copilot/test_common_utils.py create mode 100644 ui/litellm-dashboard/src/components/add_model/provider_auth_types.test.ts create mode 100644 ui/litellm-dashboard/src/components/add_model/provider_auth_types.ts diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 4c3c3aa1cfa..ab2fb4f2754 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -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 diff --git a/litellm/__init__.py b/litellm/__init__.py index 7d3fb19d662..d4ab0d648a7 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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, diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 9ffc95c6260..c25d523a8c3 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -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, diff --git a/litellm/constants.py b/litellm/constants.py index aee101dabb5..5b202f9281c 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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")) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index bab516f54cc..0aeedafb956 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -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) ) diff --git a/litellm/litellm_core_utils/oauth_token_exchange.py b/litellm/litellm_core_utils/oauth_token_exchange.py new file mode 100644 index 00000000000..d6783c6920e --- /dev/null +++ b/litellm/litellm_core_utils/oauth_token_exchange.py @@ -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) diff --git a/litellm/llms/microsoft_365_copilot/__init__.py b/litellm/llms/microsoft_365_copilot/__init__.py new file mode 100644 index 00000000000..c79aec23b7c --- /dev/null +++ b/litellm/llms/microsoft_365_copilot/__init__.py @@ -0,0 +1,3 @@ +from .chat.transformation import Microsoft365CopilotChatConfig + +__all__ = ["Microsoft365CopilotChatConfig"] diff --git a/litellm/llms/microsoft_365_copilot/chat/__init__.py b/litellm/llms/microsoft_365_copilot/chat/__init__.py new file mode 100644 index 00000000000..311d9ce87ef --- /dev/null +++ b/litellm/llms/microsoft_365_copilot/chat/__init__.py @@ -0,0 +1,3 @@ +from .transformation import Microsoft365CopilotChatConfig + +__all__ = ["Microsoft365CopilotChatConfig"] diff --git a/litellm/llms/microsoft_365_copilot/chat/handler.py b/litellm/llms/microsoft_365_copilot/chat/handler.py new file mode 100644 index 00000000000..5554c337c0d --- /dev/null +++ b/litellm/llms/microsoft_365_copilot/chat/handler.py @@ -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 diff --git a/litellm/llms/microsoft_365_copilot/chat/transformation.py b/litellm/llms/microsoft_365_copilot/chat/transformation.py new file mode 100644 index 00000000000..6b7366685fd --- /dev/null +++ b/litellm/llms/microsoft_365_copilot/chat/transformation.py @@ -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) diff --git a/litellm/llms/microsoft_365_copilot/common_utils.py b/litellm/llms/microsoft_365_copilot/common_utils.py new file mode 100644 index 00000000000..7f20b796f48 --- /dev/null +++ b/litellm/llms/microsoft_365_copilot/common_utils.py @@ -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 diff --git a/litellm/main.py b/litellm/main.py index 5c2e4d2b2a1..4c72cf26359 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index f3ec5db965b..032c262cbfa 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index d31ce4dae03..54705a8e171 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -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." ) diff --git a/litellm/proxy/common_utils/credential_hydration.py b/litellm/proxy/common_utils/credential_hydration.py index 2aabe8cad5c..2a66a3ec2c4 100644 --- a/litellm/proxy/common_utils/credential_hydration.py +++ b/litellm/proxy/common_utils/credential_hydration.py @@ -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()) diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 8df22f458cf..877fd058eb6 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -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, diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index ce30334b8e9..a5133dd384c 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -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", diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 557ff2680b4..261bf99a6f2 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -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", diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index deefea6259b..5a313b14c4a 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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, diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index 965d8dae4d2..f20369e279a 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -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//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", diff --git a/litellm/router.py b/litellm/router.py index dfdcaf2c1ca..15b8282b54d 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 ) diff --git a/litellm/router_utils/clientside_credential_handler.py b/litellm/router_utils/clientside_credential_handler.py index 186772925a9..e65bd777eca 100644 --- a/litellm/router_utils/clientside_credential_handler.py +++ b/litellm/router_utils/clientside_credential_handler.py @@ -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 diff --git a/litellm/router_utils/cooldown_handlers.py b/litellm/router_utils/cooldown_handlers.py index a6af6073a4c..84693589873 100644 --- a/litellm/router_utils/cooldown_handlers.py +++ b/litellm/router_utils/cooldown_handlers.py @@ -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) diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 239ee400503..d95774d4d85 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -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) diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index 8b254cb4b67..23162c8c156 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -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 diff --git a/litellm/types/router.py b/litellm/types/router.py index 67683e400ed..53df3ce6773 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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." ) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index a99c22e9479..01ebf139a88 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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" diff --git a/litellm/utils.py b/litellm/utils.py index b51f349c4e6..b67a7294c6d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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.""" diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index f3ec5db965b..032c262cbfa 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index ef1fa270b06..33ac4e4e60b 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -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", diff --git a/tests/unit/caching/test_caching_handler.py b/tests/unit/caching/test_caching_handler.py index d98c02667ee..2520f2eb2da 100644 --- a/tests/unit/caching/test_caching_handler.py +++ b/tests/unit/caching/test_caching_handler.py @@ -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=[ diff --git a/tests/unit/litellm_core_utils/test_oauth_token_exchange.py b/tests/unit/litellm_core_utils/test_oauth_token_exchange.py new file mode 100644 index 00000000000..2a313539765 --- /dev/null +++ b/tests/unit/litellm_core_utils/test_oauth_token_exchange.py @@ -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 diff --git a/tests/unit/llms/anthropic/test_anthropic_common_utils.py b/tests/unit/llms/anthropic/test_anthropic_common_utils.py index 265ea082683..93ec7b7b7ca 100644 --- a/tests/unit/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/unit/llms/anthropic/test_anthropic_common_utils.py @@ -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={}, diff --git a/tests/unit/llms/microsoft_365_copilot/__init__.py b/tests/unit/llms/microsoft_365_copilot/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/microsoft_365_copilot/chat/__init__.py b/tests/unit/llms/microsoft_365_copilot/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/microsoft_365_copilot/chat/test_handler.py b/tests/unit/llms/microsoft_365_copilot/chat/test_handler.py new file mode 100644 index 00000000000..7be5db74af5 --- /dev/null +++ b/tests/unit/llms/microsoft_365_copilot/chat/test_handler.py @@ -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, + ) diff --git a/tests/unit/llms/microsoft_365_copilot/chat/test_transformation.py b/tests/unit/llms/microsoft_365_copilot/chat/test_transformation.py new file mode 100644 index 00000000000..22cb8da3aed --- /dev/null +++ b/tests/unit/llms/microsoft_365_copilot/chat/test_transformation.py @@ -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) diff --git a/tests/unit/llms/microsoft_365_copilot/test_common_utils.py b/tests/unit/llms/microsoft_365_copilot/test_common_utils.py new file mode 100644 index 00000000000..fec367418ad --- /dev/null +++ b/tests/unit/llms/microsoft_365_copilot/test_common_utils.py @@ -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 diff --git a/tests/unit/proxy/auth/test_auth_utils.py b/tests/unit/proxy/auth/test_auth_utils.py index 52c5fa2a41e..412a104f252 100644 --- a/tests/unit/proxy/auth/test_auth_utils.py +++ b/tests/unit/proxy/auth/test_auth_utils.py @@ -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``), diff --git a/tests/unit/proxy/credential_endpoints/test_endpoints.py b/tests/unit/proxy/credential_endpoints/test_endpoints.py index 83d6167256a..b69150a95ec 100644 --- a/tests/unit/proxy/credential_endpoints/test_endpoints.py +++ b/tests/unit/proxy/credential_endpoints/test_endpoints.py @@ -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() diff --git a/tests/unit/proxy/health_endpoints/test_health_endpoints.py b/tests/unit/proxy/health_endpoints/test_health_endpoints.py index 5aa763fcc9c..bf5deb4d362 100644 --- a/tests/unit/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/unit/proxy/health_endpoints/test_health_endpoints.py @@ -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)]) diff --git a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py index f44b2b57ccd..dbb2f16d056 100644 --- a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py @@ -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", diff --git a/tests/unit/proxy/test_credential_slot_registry.py b/tests/unit/proxy/test_credential_slot_registry.py index 044b00b33bb..6332038feb8 100644 --- a/tests/unit/proxy/test_credential_slot_registry.py +++ b/tests/unit/proxy/test_credential_slot_registry.py @@ -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"), diff --git a/tests/unit/router_utils/test_cooldown_handlers.py b/tests/unit/router_utils/test_cooldown_handlers.py index d2e21f172db..d219abb7d3d 100644 --- a/tests/unit/router_utils/test_cooldown_handlers.py +++ b/tests/unit/router_utils/test_cooldown_handlers.py @@ -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): diff --git a/tests/unit/router_utils/test_fallback_event_handlers.py b/tests/unit/router_utils/test_fallback_event_handlers.py index c23a8c95906..febb33aab7d 100644 --- a/tests/unit/router_utils/test_fallback_event_handlers.py +++ b/tests/unit/router_utils/test_fallback_event_handlers.py @@ -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"}], diff --git a/tests/unit/test_internal_context.py b/tests/unit/test_internal_context.py index 665ed4f4a8f..1f43f9c6312 100644 --- a/tests/unit/test_internal_context.py +++ b/tests/unit/test_internal_context.py @@ -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", diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index f8e2befe237..f9bb957e777 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -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", diff --git a/tests/unit/types/test_router.py b/tests/unit/types/test_router.py index d2817ae6a90..dce4225f70a 100644 --- a/tests/unit/types/test_router.py +++ b/tests/unit/types/test_router.py @@ -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"} diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx index e80d0e554e6..8d0880360cd 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx @@ -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(); + await screen.findByText("Existing Credentials"); + return props; + }; + + const submitModel = async (props: ReturnType) => { + 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")); diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx index e50750318c7..fbdab07a545 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx @@ -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 = ({ 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(); + 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) => { + const handleCreateCredential = async (values: Record) => { const credential = buildCredential(values, withoutRestrictedFields(values)); try { await credentialCreateCall(accessToken, credential); @@ -112,11 +118,25 @@ const AddModelForm: React.FC = ({ 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 = ({
{ event.preventDefault(); - void handleOk().then((submitted) => { + void handleSubmit().then((submitted) => { if (submitted) { setTeamAdminSelectedTeam(null); } @@ -331,7 +351,12 @@ const AddModelForm: React.FC = ({ OR
- + {canCreateFederatedCredential && (
@@ -476,6 +501,17 @@ const AddModelForm: React.FC = ({ + {isCredentialModalOpen && ( + setIsCredentialModalOpen(false)} + onSubmit={handleCreateCredential} + /> + )} {isFederatedCredentialModalOpen && ( = ({ initialAuthMethod="federation" providerLocked onCancel={() => setIsFederatedCredentialModalOpen(false)} - onSubmit={handleCreateFederatedCredential} + onSubmit={handleCreateCredential} /> )} {/* Test Connection Results Modal */} diff --git a/ui/litellm-dashboard/src/components/add_model/provider_auth_types.test.ts b/ui/litellm-dashboard/src/components/add_model/provider_auth_types.test.ts new file mode 100644 index 00000000000..9670c6c6e27 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/provider_auth_types.test.ts @@ -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 = { + token_exchange_profile: "jwt_bearer_obo", + token_exchange_scope: "https://graph.microsoft.com/.default", + api_key: "delegated-token", + }; + const emptyValues: Record = { + 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"]); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/provider_auth_types.ts b/ui/litellm-dashboard/src/components/add_model/provider_auth_types.ts new file mode 100644 index 00000000000..37fbcb73e47 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/provider_auth_types.ts @@ -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>; + readonly credentialOnly?: boolean; +} + +const EMPTY_PROVIDER_AUTH_TYPES: readonly ProviderAuthType[] = []; + +export const PROVIDER_AUTH_TYPES: Partial> = { + 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 => + 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, + ); +}; diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx index 1ea9505201c..09d48f828e1 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx @@ -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//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 {String(watch("vertex_credentials") ?? "")}; }; +const ValidationForm = ({ children }: { readonly children: ReactNode }) => { + const form = useFormContext(); + return ( + {})}> + {children} + + + ); +}; + 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( + + + + + , + ); + + 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( + + + + + , + ); + + 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( + + + + + , + ); + + 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( + + + + + , + ); + + 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( + + + + + , + ); + + 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( + + + + + , + ); + + await chooseSelectOption( + user, + await screen.findByRole("combobox", { name: "Auth Type:" }), + "OAuth token exchange (on-behalf-of)", + ); + + rerender( + + + + + , + ); + await screen.findByLabelText("OpenAI API Key"); + + rerender( + + + + + , + ); + + 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( + + + + + + + , + ); + + 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( + + + + + , + ); + + await screen.findByLabelText("OpenAI API Key"); + + expect(screen.queryByRole("combobox", { name: "Auth Type:" })).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx index 69f11897009..a38b077270f 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx +++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx @@ -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 = {}; -const ProviderSpecificFields: React.FC = ({ +const ProviderSpecificFieldsContent: React.FC = ({ selectedProvider, hiddenFieldKeys, + context = "credential", + onCreateCredential, + initialAuthTypeId, + onAuthTypeChange, fieldValidators, }) => { const selectedProviderEnum = Providers[selectedProvider as keyof typeof Providers] as Providers; const form = useFormContext(); + 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(null); + + React.useEffect(() => { + if (selectedAuthType) { + onAuthTypeChange?.(selectedAuthType.id); + } + }, [onAuthTypeChange, selectedAuthType]); const pickCredentialsFile = (onLoaded: (contents: string) => void) => (event: React.ChangeEvent) => { const file = event.target.files?.[0]; @@ -172,11 +203,34 @@ const ProviderSpecificFields: React.FC = ({ 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(null); @@ -291,6 +345,44 @@ const ProviderSpecificFields: React.FC = ({ return ( <> + {availableAuthTypes.length > 0 && selectedAuthType && ( +
+ + +

{selectedAuthType.description}

+
+ )} + {credentialOnlySelected ? ( + <> +

This auth type is saved as an LLM credential and attached to this model.

+ + + ) : ( + Object.entries(selectedAuthType?.fixedValues ?? {}).map(([fieldKey, fixedValue]) => ( + + {() => null} + + )) + )} {isLoading && allFields.length === 0 &&

Loading provider fields...

} {loadError && allFields.length === 0 && (

@@ -341,4 +433,8 @@ const ProviderSpecificFields: React.FC = ({ ); }; +const ProviderSpecificFields: React.FC = (props) => ( + +); + export default ProviderSpecificFields; diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.integration.test.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.integration.test.tsx index 4d39e0365cd..4b9f4e61a8f 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialModal.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.integration.test.tsx @@ -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> = {}) => { const onSubmit = vi.fn(); render( @@ -109,6 +142,34 @@ const chooseProvider = async (user: ReturnType, 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(); diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx index 81b4fef9a6d..c97e4658817 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx @@ -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({ diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx index cecd6ce01bf..e1f379b11bd 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx @@ -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(""); diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx index 177d4be3247..c543de7e58c 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx @@ -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 = { 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> = { [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> = { [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", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index a3fae7a2f4f..4c072cfb7ee 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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 */