mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
feat(microsoft_365_copilot): add Microsoft 365 Copilot chat provider with OAuth token exchange (#45158)
* feat(microsoft_365_copilot): add Microsoft 365 Copilot chat provider with OAuth token exchange * fix(ui): create a credential for credential-only auth types in Add Model * fix(microsoft_365_copilot): accept max_tokens and list the chat model * fix(microsoft_365_copilot): register provider model set * fix(ui): hide other auth types' fields when credential-only types are filtered * fix(health): do not inherit stored auth settings when a connection test brings its own api_key * docs(health): restore connection-test inheritance docstrings * fix(lint): suppress justified provider boundary casts * fix(lint): satisfy M365 type-discipline gate * fix(types): eliminate type-check gate regressions * style: shorten suppression reasons so ruff format leaves them on one line * refactor(types): drop shared-file type widening unrelated to the Copilot provider * fix(health): pass caller headers as secret_fields to connection-test probes * fix(types): type connection-test probe and token counter locally * fix(lint): drop cast import from connection-test header getter * chore: remove explanatory comments flagged by review * fix(proxy): gate OAuth client credentials behind proxy admin * fix(router): skip cooldown on caller-scoped OAuth auth failures * fix(router): skip fallback cooldown on caller-scoped OAuth auth failures * fix(router): type caller-scoped cooldown lookups * fix(ci): include Microsoft 365 Copilot tests in unit shard * fix(proxy): keep saved endpoint when test_connection retests a deployment by id * fix(m365): collapse doubled Graph replies * fix(m365): preserve trailing system messages --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
2248fcefda
commit
fbf4af301c
59 changed files with 4925 additions and 144 deletions
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
||||
|
|
|
|||
307
litellm/litellm_core_utils/oauth_token_exchange.py
Normal file
307
litellm/litellm_core_utils/oauth_token_exchange.py
Normal file
|
|
@ -0,0 +1,307 @@
|
|||
import hashlib
|
||||
import json
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import (
|
||||
Final,
|
||||
Literal,
|
||||
Protocol,
|
||||
TypeAlias,
|
||||
cast, # noqa: TID251 # adapter protocols cover untyped cache and HTTP handler methods
|
||||
)
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import OAUTH_TOKEN_EXCHANGE_CACHE_SAFETY_MARGIN_SECONDS
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
|
||||
TokenExchangeProfile: TypeAlias = Literal["rfc8693", "jwt_bearer_obo"]
|
||||
|
||||
SUBJECT_TOKEN_TYPE_ACCESS_TOKEN: Final = "urn:ietf:params:oauth:token-type:access_token"
|
||||
_RFC8693_GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:token-exchange"
|
||||
_JWT_BEARER_GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:jwt-bearer"
|
||||
_TOKEN_CACHE: Final = InMemoryCache(max_size_in_memory=1000, default_ttl=600)
|
||||
|
||||
|
||||
class _TokenCache(Protocol):
|
||||
def set_cache(self, *, key: str, value: str, ttl: int) -> None: ...
|
||||
|
||||
def get_cache(self, *, key: str) -> object: ...
|
||||
|
||||
|
||||
class _SyncTokenExchangeClient(Protocol):
|
||||
def post(
|
||||
self,
|
||||
url: str,
|
||||
*,
|
||||
data: Mapping[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> httpx.Response: ...
|
||||
|
||||
|
||||
class _AsyncTokenExchangeClient(Protocol):
|
||||
async def post(
|
||||
self,
|
||||
url: str,
|
||||
*,
|
||||
data: Mapping[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> httpx.Response: ...
|
||||
|
||||
|
||||
class OAuthTokenExchangeError(Exception):
|
||||
def __init__(self, status_code: int, message: str) -> None:
|
||||
self.status_code: Final = status_code
|
||||
self.message: Final = message
|
||||
super().__init__(message)
|
||||
|
||||
|
||||
class _OAuthTokenResponse(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
|
||||
|
||||
access_token: str
|
||||
expires_in: int = Field(gt=0)
|
||||
|
||||
|
||||
class _OAuthErrorResponse(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
|
||||
|
||||
error: str | None = None
|
||||
error_description: str | None = None
|
||||
|
||||
|
||||
_TOKEN_RESPONSE_ADAPTER: Final = TypeAdapter(_OAuthTokenResponse)
|
||||
_ERROR_RESPONSE_ADAPTER: Final = TypeAdapter(_OAuthErrorResponse)
|
||||
|
||||
|
||||
def _is_valid_token_endpoint(value: str) -> bool:
|
||||
try:
|
||||
endpoint: Final = httpx.URL(value)
|
||||
except httpx.InvalidURL:
|
||||
return False
|
||||
return endpoint.scheme.casefold() == "https" and bool(endpoint.host)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, repr=False)
|
||||
class OAuthTokenExchangeConfig:
|
||||
token_endpoint: str
|
||||
client_id: str
|
||||
client_secret: str
|
||||
profile: TokenExchangeProfile
|
||||
scope: str | None
|
||||
audience: str | None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not _is_valid_token_endpoint(self.token_endpoint):
|
||||
raise OAuthTokenExchangeError(
|
||||
status_code=400,
|
||||
message="token_exchange_endpoint must be an HTTPS URL with a host",
|
||||
)
|
||||
if self.profile not in ("rfc8693", "jwt_bearer_obo"):
|
||||
raise OAuthTokenExchangeError(status_code=400, message="Unsupported OAuth token exchange profile")
|
||||
if self.profile == "jwt_bearer_obo" and (self.scope is None or not self.scope.strip()):
|
||||
raise OAuthTokenExchangeError(
|
||||
status_code=400,
|
||||
message="jwt_bearer_obo requires a non-empty token exchange scope",
|
||||
)
|
||||
|
||||
|
||||
def redact_sensitive_values(message: str, sensitive_values: tuple[str, ...]) -> str:
|
||||
redacted_values: Final = tuple(sorted((value for value in sensitive_values if value), key=len, reverse=True))
|
||||
if not redacted_values:
|
||||
return message
|
||||
pattern: Final = re.compile("|".join(re.escape(value) for value in redacted_values))
|
||||
return pattern.sub("[REDACTED]", message)
|
||||
|
||||
|
||||
def _cache_key(config: OAuthTokenExchangeConfig, subject_token: str) -> str:
|
||||
key_material: Final = json.dumps(
|
||||
{
|
||||
"token_endpoint": config.token_endpoint,
|
||||
"client_id": config.client_id,
|
||||
"client_secret": config.client_secret,
|
||||
"profile": config.profile,
|
||||
"scope": config.scope,
|
||||
"audience": config.audience,
|
||||
"subject_token": subject_token,
|
||||
},
|
||||
ensure_ascii=True,
|
||||
separators=(",", ":"),
|
||||
sort_keys=True,
|
||||
)
|
||||
return hashlib.sha256(key_material.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def build_token_exchange_form(
|
||||
config: OAuthTokenExchangeConfig,
|
||||
subject_token: str,
|
||||
) -> dict[str, str]: # mutable-ok: the form builder contract requires an HTTP form dictionary
|
||||
match config.profile:
|
||||
case "rfc8693":
|
||||
return {
|
||||
"grant_type": _RFC8693_GRANT_TYPE,
|
||||
"client_id": config.client_id,
|
||||
"client_secret": config.client_secret,
|
||||
"subject_token": subject_token,
|
||||
"subject_token_type": SUBJECT_TOKEN_TYPE_ACCESS_TOKEN,
|
||||
**({"scope": config.scope} if config.scope else {}),
|
||||
**({"audience": config.audience} if config.audience else {}),
|
||||
}
|
||||
case "jwt_bearer_obo":
|
||||
if config.scope is None or not config.scope.strip():
|
||||
raise OAuthTokenExchangeError(
|
||||
status_code=400,
|
||||
message="jwt_bearer_obo requires a non-empty token exchange scope",
|
||||
)
|
||||
return {
|
||||
"grant_type": _JWT_BEARER_GRANT_TYPE,
|
||||
"client_id": config.client_id,
|
||||
"client_secret": config.client_secret,
|
||||
"assertion": subject_token,
|
||||
"scope": config.scope,
|
||||
"requested_token_use": "on_behalf_of",
|
||||
}
|
||||
|
||||
|
||||
def _error_message(response: httpx.Response, sensitive_values: tuple[str, ...]) -> str:
|
||||
try:
|
||||
error_response: Final = _ERROR_RESPONSE_ADAPTER.validate_python(response.json())
|
||||
except (ValidationError, ValueError):
|
||||
return redact_sensitive_values(
|
||||
"OAuth token exchange failed: unknown_error: invalid response",
|
||||
sensitive_values,
|
||||
)
|
||||
error: Final = error_response.error or "unknown_error"
|
||||
description: Final = error_response.error_description or "no description provided"
|
||||
return redact_sensitive_values(
|
||||
f"OAuth token exchange failed: {error}: {description}",
|
||||
sensitive_values,
|
||||
)
|
||||
|
||||
|
||||
def _response_token(response: httpx.Response, config: OAuthTokenExchangeConfig, subject_token: str) -> str:
|
||||
try:
|
||||
token_response: Final = _TOKEN_RESPONSE_ADAPTER.validate_python(response.json())
|
||||
except (ValidationError, ValueError):
|
||||
raise OAuthTokenExchangeError(
|
||||
status_code=502,
|
||||
message="OAuth token exchange returned an invalid token response",
|
||||
) from None
|
||||
effective_ttl: Final = token_response.expires_in - OAUTH_TOKEN_EXCHANGE_CACHE_SAFETY_MARGIN_SECONDS
|
||||
if effective_ttl > 0:
|
||||
cache: Final = cast(_TokenCache, _TOKEN_CACHE) # cast-ok: InMemoryCache exposes untyped cache methods
|
||||
cache.set_cache(
|
||||
key=_cache_key(config, subject_token),
|
||||
value=token_response.access_token,
|
||||
ttl=effective_ttl,
|
||||
)
|
||||
return token_response.access_token
|
||||
|
||||
|
||||
def _error_status_code(status_code: int) -> int:
|
||||
if status_code == 429:
|
||||
return status_code
|
||||
if 400 <= status_code < 500:
|
||||
return 401
|
||||
return status_code
|
||||
|
||||
|
||||
def _post(
|
||||
client: HTTPHandler,
|
||||
config: OAuthTokenExchangeConfig,
|
||||
request_data: Mapping[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> httpx.Response:
|
||||
try:
|
||||
http_client: Final = cast( # cast-ok: HTTPHandler.post has an untyped response contract
|
||||
_SyncTokenExchangeClient,
|
||||
client,
|
||||
)
|
||||
return http_client.post(config.token_endpoint, data=dict(request_data), timeout=timeout)
|
||||
except httpx.HTTPStatusError as error:
|
||||
return error.response
|
||||
except httpx.HTTPError:
|
||||
raise OAuthTokenExchangeError(
|
||||
status_code=502,
|
||||
message="OAuth token exchange request failed",
|
||||
) from None
|
||||
|
||||
|
||||
async def _apost(
|
||||
client: AsyncHTTPHandler,
|
||||
config: OAuthTokenExchangeConfig,
|
||||
request_data: Mapping[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> httpx.Response:
|
||||
try:
|
||||
http_client: Final = cast( # cast-ok: AsyncHTTPHandler.post has an untyped response contract
|
||||
_AsyncTokenExchangeClient,
|
||||
client,
|
||||
)
|
||||
return await http_client.post(config.token_endpoint, data=dict(request_data), timeout=timeout)
|
||||
except httpx.HTTPStatusError as error:
|
||||
return error.response
|
||||
except httpx.HTTPError:
|
||||
raise OAuthTokenExchangeError(
|
||||
status_code=502,
|
||||
message="OAuth token exchange request failed",
|
||||
) from None
|
||||
|
||||
|
||||
def exchange_token(
|
||||
client: HTTPHandler,
|
||||
config: OAuthTokenExchangeConfig,
|
||||
subject_token: str,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> str:
|
||||
cache_key: Final = _cache_key(config, subject_token)
|
||||
cache: Final = cast(_TokenCache, _TOKEN_CACHE) # cast-ok: InMemoryCache exposes untyped cache methods
|
||||
cached_token: Final = cache.get_cache(key=cache_key)
|
||||
if isinstance(cached_token, str):
|
||||
return cached_token
|
||||
response: Final = _post(
|
||||
client=client,
|
||||
config=config,
|
||||
request_data=build_token_exchange_form(config, subject_token),
|
||||
timeout=timeout,
|
||||
)
|
||||
if response.status_code != 200:
|
||||
raise OAuthTokenExchangeError(
|
||||
status_code=_error_status_code(response.status_code),
|
||||
message=_error_message(
|
||||
response=response,
|
||||
sensitive_values=(config.client_secret, subject_token),
|
||||
),
|
||||
)
|
||||
return _response_token(response=response, config=config, subject_token=subject_token)
|
||||
|
||||
|
||||
async def aexchange_token(
|
||||
client: AsyncHTTPHandler,
|
||||
config: OAuthTokenExchangeConfig,
|
||||
subject_token: str,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> str:
|
||||
cache_key: Final = _cache_key(config, subject_token)
|
||||
cache: Final = cast(_TokenCache, _TOKEN_CACHE) # cast-ok: InMemoryCache exposes untyped cache methods
|
||||
cached_token: Final = cache.get_cache(key=cache_key)
|
||||
if isinstance(cached_token, str):
|
||||
return cached_token
|
||||
response: Final = await _apost(
|
||||
client=client,
|
||||
config=config,
|
||||
request_data=build_token_exchange_form(config, subject_token),
|
||||
timeout=timeout,
|
||||
)
|
||||
if response.status_code != 200:
|
||||
raise OAuthTokenExchangeError(
|
||||
status_code=_error_status_code(response.status_code),
|
||||
message=_error_message(
|
||||
response=response,
|
||||
sensitive_values=(config.client_secret, subject_token),
|
||||
),
|
||||
)
|
||||
return _response_token(response=response, config=config, subject_token=subject_token)
|
||||
3
litellm/llms/microsoft_365_copilot/__init__.py
Normal file
3
litellm/llms/microsoft_365_copilot/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .chat.transformation import Microsoft365CopilotChatConfig
|
||||
|
||||
__all__ = ["Microsoft365CopilotChatConfig"]
|
||||
3
litellm/llms/microsoft_365_copilot/chat/__init__.py
Normal file
3
litellm/llms/microsoft_365_copilot/chat/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .transformation import Microsoft365CopilotChatConfig
|
||||
|
||||
__all__ = ["Microsoft365CopilotChatConfig"]
|
||||
580
litellm/llms/microsoft_365_copilot/chat/handler.py
Normal file
580
litellm/llms/microsoft_365_copilot/chat/handler.py
Normal file
|
|
@ -0,0 +1,580 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import (
|
||||
Final,
|
||||
Protocol,
|
||||
cast, # noqa: TID251 # adapter protocols cover pluggable logging and untyped HTTP methods
|
||||
)
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
from aiohttp import ClientSession
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm import LlmProviders
|
||||
from litellm.constants import (
|
||||
MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_PROFILE,
|
||||
MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_SCOPE,
|
||||
MICROSOFT_GRAPH_BETA_BASE,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.oauth_token_exchange import (
|
||||
OAuthTokenExchangeConfig,
|
||||
OAuthTokenExchangeError,
|
||||
TokenExchangeProfile,
|
||||
aexchange_token,
|
||||
exchange_token,
|
||||
)
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.llms.custom_httpx import http_handler
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
)
|
||||
from litellm.llms.custom_httpx.llm_http_handler import MockResponseIterator
|
||||
from litellm.llms.microsoft_365_copilot.chat.transformation import (
|
||||
GraphChatRequest,
|
||||
build_chat_request,
|
||||
extract_graph_error_message,
|
||||
map_graph_response,
|
||||
parse_graph_conversation_id,
|
||||
)
|
||||
from litellm.llms.microsoft_365_copilot.common_utils import (
|
||||
Microsoft365CopilotError,
|
||||
extract_caller_assertion,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, repr=False)
|
||||
class _TokenExchangeCredentials:
|
||||
config: OAuthTokenExchangeConfig
|
||||
subject_token: str
|
||||
|
||||
|
||||
class _CopilotLogging(Protocol):
|
||||
def pre_call(
|
||||
self,
|
||||
*,
|
||||
input: Sequence[AllMessageValues],
|
||||
api_key: str,
|
||||
model: str,
|
||||
additional_args: Mapping[str, object],
|
||||
) -> object: ...
|
||||
|
||||
def post_call(
|
||||
self,
|
||||
*,
|
||||
original_response: object,
|
||||
input: Sequence[AllMessageValues],
|
||||
api_key: str,
|
||||
additional_args: Mapping[str, object],
|
||||
) -> object: ...
|
||||
|
||||
|
||||
class _SyncGraphClient(Protocol):
|
||||
def post(
|
||||
self,
|
||||
url: str,
|
||||
*,
|
||||
headers: Mapping[str, str],
|
||||
json: Mapping[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> httpx.Response: ...
|
||||
|
||||
|
||||
class _AsyncGraphClient(Protocol):
|
||||
async def post(
|
||||
self,
|
||||
url: str,
|
||||
*,
|
||||
headers: Mapping[str, str],
|
||||
json: Mapping[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> httpx.Response: ...
|
||||
|
||||
|
||||
class _HTTPClientFactory(Protocol):
|
||||
def get_httpx_client(self, params: Mapping[str, object] | None = None) -> HTTPHandler: ...
|
||||
|
||||
def get_async_httpx_client(
|
||||
self,
|
||||
llm_provider: LlmProviders | str,
|
||||
params: Mapping[str, object] | None = None,
|
||||
shared_session: ClientSession | None = None,
|
||||
) -> AsyncHTTPHandler: ...
|
||||
|
||||
|
||||
_JSON_VALUE_ADAPTER: Final = TypeAdapter(object)
|
||||
_HTTP_CLIENT_FACTORY: Final = cast( # cast-ok: shared cached-client methods have partially typed signatures
|
||||
_HTTPClientFactory, http_handler
|
||||
)
|
||||
|
||||
|
||||
def _get_sync_http_client(params: Mapping[str, object] | None) -> HTTPHandler:
|
||||
return _HTTP_CLIENT_FACTORY.get_httpx_client(dict(params) if params is not None else None)
|
||||
|
||||
|
||||
def _get_async_http_client(
|
||||
provider: LlmProviders | str,
|
||||
params: Mapping[str, object] | None,
|
||||
shared_session: ClientSession | None,
|
||||
) -> AsyncHTTPHandler:
|
||||
return _HTTP_CLIENT_FACTORY.get_async_httpx_client(
|
||||
provider,
|
||||
dict(params) if params is not None else None,
|
||||
shared_session,
|
||||
)
|
||||
|
||||
|
||||
def _token_exchange_profile(value: object) -> TokenExchangeProfile:
|
||||
if value is None:
|
||||
return MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_PROFILE
|
||||
if value == "rfc8693":
|
||||
return "rfc8693"
|
||||
if value == "jwt_bearer_obo":
|
||||
return "jwt_bearer_obo"
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=400,
|
||||
message="microsoft_365_copilot token_exchange_profile must be 'rfc8693' or 'jwt_bearer_obo'",
|
||||
)
|
||||
|
||||
|
||||
def _get_token_exchange_credentials(
|
||||
litellm_params: Mapping[str, object],
|
||||
secret_fields: SecretFields | None,
|
||||
) -> _TokenExchangeCredentials | None:
|
||||
token_endpoint: Final[object] = litellm_params.get("token_exchange_endpoint")
|
||||
client_id: Final[object] = litellm_params.get("client_id")
|
||||
client_secret: Final[object] = litellm_params.get("client_secret")
|
||||
credential_values: Final = (token_endpoint, client_id, client_secret)
|
||||
configured_values: Final = tuple(value is not None and value != "" for value in credential_values)
|
||||
if not any(configured_values):
|
||||
return None
|
||||
if not all(configured_values):
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=400,
|
||||
message="microsoft_365_copilot requires token_exchange_endpoint, client_id, and client_secret together",
|
||||
)
|
||||
if not isinstance(token_endpoint, str) or not isinstance(client_id, str) or not isinstance(client_secret, str):
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=400,
|
||||
message=(
|
||||
"microsoft_365_copilot requires token_exchange_endpoint, client_id, and client_secret "
|
||||
"as non-empty strings"
|
||||
),
|
||||
)
|
||||
if not token_endpoint.strip() or not client_id.strip() or not client_secret.strip():
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=400,
|
||||
message="microsoft_365_copilot requires token_exchange_endpoint, client_id, and client_secret together",
|
||||
)
|
||||
profile: Final = _token_exchange_profile(litellm_params.get("token_exchange_profile"))
|
||||
scope_value: Final[object] = litellm_params.get("token_exchange_scope")
|
||||
if scope_value is not None and not isinstance(scope_value, str):
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=400,
|
||||
message="microsoft_365_copilot token_exchange_scope must be a string",
|
||||
)
|
||||
audience_value: Final[object] = litellm_params.get("token_exchange_audience")
|
||||
if audience_value is not None and not isinstance(audience_value, str):
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=400,
|
||||
message="microsoft_365_copilot token_exchange_audience must be a string",
|
||||
)
|
||||
scope: Final = MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_SCOPE if scope_value is None else scope_value
|
||||
try:
|
||||
config: Final = OAuthTokenExchangeConfig(
|
||||
token_endpoint=token_endpoint,
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
profile=profile,
|
||||
scope=scope,
|
||||
audience=audience_value,
|
||||
)
|
||||
except OAuthTokenExchangeError as error:
|
||||
raise Microsoft365CopilotError(status_code=error.status_code, message=error.message) from error
|
||||
subject_token: Final = extract_caller_assertion(secret_fields)
|
||||
if subject_token is None:
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=401,
|
||||
message=(
|
||||
"microsoft_365_copilot with OAuth token exchange requires the caller's "
|
||||
"IdP-issued access token in the Authorization header"
|
||||
),
|
||||
)
|
||||
return _TokenExchangeCredentials(config=config, subject_token=subject_token)
|
||||
|
||||
|
||||
def _resolve_access_token(
|
||||
client: HTTPHandler,
|
||||
api_key: str | None,
|
||||
litellm_params: Mapping[str, object],
|
||||
secret_fields: SecretFields | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> str:
|
||||
exchange_credentials: Final = _get_token_exchange_credentials(
|
||||
litellm_params=litellm_params,
|
||||
secret_fields=secret_fields,
|
||||
)
|
||||
if exchange_credentials is not None:
|
||||
try:
|
||||
return exchange_token(
|
||||
client=client,
|
||||
config=exchange_credentials.config,
|
||||
subject_token=exchange_credentials.subject_token,
|
||||
timeout=timeout,
|
||||
)
|
||||
except OAuthTokenExchangeError as error:
|
||||
raise Microsoft365CopilotError(status_code=error.status_code, message=error.message) from error
|
||||
if isinstance(api_key, str) and api_key:
|
||||
return api_key
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=401,
|
||||
message="microsoft_365_copilot requires api_key or a complete token exchange configuration",
|
||||
)
|
||||
|
||||
|
||||
async def _aresolve_access_token(
|
||||
client: AsyncHTTPHandler,
|
||||
api_key: str | None,
|
||||
litellm_params: Mapping[str, object],
|
||||
secret_fields: SecretFields | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> str:
|
||||
exchange_credentials: Final = _get_token_exchange_credentials(
|
||||
litellm_params=litellm_params,
|
||||
secret_fields=secret_fields,
|
||||
)
|
||||
if exchange_credentials is not None:
|
||||
try:
|
||||
return await aexchange_token(
|
||||
client=client,
|
||||
config=exchange_credentials.config,
|
||||
subject_token=exchange_credentials.subject_token,
|
||||
timeout=timeout,
|
||||
)
|
||||
except OAuthTokenExchangeError as error:
|
||||
raise Microsoft365CopilotError(status_code=error.status_code, message=error.message) from error
|
||||
if isinstance(api_key, str) and api_key:
|
||||
return api_key
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=401,
|
||||
message="microsoft_365_copilot requires api_key or a complete token exchange configuration",
|
||||
)
|
||||
|
||||
|
||||
def _json_body(response: httpx.Response) -> object:
|
||||
try:
|
||||
return _JSON_VALUE_ADAPTER.validate_json(response.content)
|
||||
except ValidationError:
|
||||
return {}
|
||||
|
||||
|
||||
def _graph_error(
|
||||
response: httpx.Response,
|
||||
sensitive_values: tuple[str, ...],
|
||||
) -> Microsoft365CopilotError:
|
||||
return Microsoft365CopilotError(
|
||||
status_code=response.status_code,
|
||||
message=extract_graph_error_message(
|
||||
response_body=_json_body(response),
|
||||
sensitive_values=sensitive_values,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _pre_call(
|
||||
logging_obj: Logging | None,
|
||||
messages: Sequence[AllMessageValues],
|
||||
model: str,
|
||||
request_data: GraphChatRequest,
|
||||
) -> None:
|
||||
if logging_obj is None:
|
||||
return
|
||||
typed_logging: Final = cast( # cast-ok: Logging callback arguments are provider-pluggable
|
||||
_CopilotLogging,
|
||||
logging_obj,
|
||||
)
|
||||
typed_logging.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
model=model,
|
||||
additional_args={
|
||||
"api_base": MICROSOFT_GRAPH_BETA_BASE,
|
||||
"complete_input_dict": request_data,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _post_call(
|
||||
logging_obj: Logging | None,
|
||||
messages: Sequence[AllMessageValues],
|
||||
request_data: GraphChatRequest,
|
||||
status_code: int,
|
||||
) -> None:
|
||||
if logging_obj is None:
|
||||
return
|
||||
typed_logging: Final = cast( # cast-ok: Logging callback arguments are provider-pluggable
|
||||
_CopilotLogging,
|
||||
logging_obj,
|
||||
)
|
||||
typed_logging.post_call(
|
||||
original_response={"status_code": status_code},
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"api_base": MICROSOFT_GRAPH_BETA_BASE,
|
||||
"complete_input_dict": request_data,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _post_graph_request(
|
||||
client: HTTPHandler,
|
||||
url: str,
|
||||
access_token: str,
|
||||
request_data: GraphChatRequest | Mapping[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> httpx.Response:
|
||||
try:
|
||||
http_client: Final = cast( # cast-ok: HTTPHandler.post has an untyped response contract
|
||||
_SyncGraphClient,
|
||||
client,
|
||||
)
|
||||
request_mapping: Final[Mapping[str, object]] = request_data
|
||||
return http_client.post(
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {access_token}"},
|
||||
json=dict(request_mapping),
|
||||
timeout=timeout,
|
||||
)
|
||||
except httpx.HTTPStatusError as error:
|
||||
return error.response
|
||||
except httpx.HTTPError:
|
||||
raise Microsoft365CopilotError(status_code=502, message="Microsoft Graph request failed") from None
|
||||
|
||||
|
||||
async def _apost_graph_request(
|
||||
client: AsyncHTTPHandler,
|
||||
url: str,
|
||||
access_token: str,
|
||||
request_data: GraphChatRequest | Mapping[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> httpx.Response:
|
||||
try:
|
||||
http_client: Final = cast( # cast-ok: AsyncHTTPHandler.post has an untyped response contract
|
||||
_AsyncGraphClient,
|
||||
client,
|
||||
)
|
||||
request_mapping: Final[Mapping[str, object]] = request_data
|
||||
return await http_client.post(
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {access_token}"},
|
||||
json=dict(request_mapping),
|
||||
timeout=timeout,
|
||||
)
|
||||
except httpx.HTTPStatusError as error:
|
||||
return error.response
|
||||
except httpx.HTTPError:
|
||||
raise Microsoft365CopilotError(status_code=502, message="Microsoft Graph request failed") from None
|
||||
|
||||
|
||||
def _stream_response(
|
||||
response: ModelResponse,
|
||||
model: str,
|
||||
logging_obj: Logging | None,
|
||||
) -> CustomStreamWrapper:
|
||||
if logging_obj is None:
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=500,
|
||||
message="Microsoft 365 Copilot streaming requires a logging context",
|
||||
)
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=MockResponseIterator(model_response=response),
|
||||
model=model,
|
||||
custom_llm_provider=LlmProviders.MICROSOFT_365_COPILOT.value,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
|
||||
def _raise_graph_error_if_needed(
|
||||
response: httpx.Response,
|
||||
sensitive_values: tuple[str, ...],
|
||||
logging_obj: Logging | None,
|
||||
messages: Sequence[AllMessageValues],
|
||||
request_data: GraphChatRequest,
|
||||
) -> None:
|
||||
if response.is_success:
|
||||
return
|
||||
_post_call(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
request_data=request_data,
|
||||
status_code=response.status_code,
|
||||
)
|
||||
raise _graph_error(response=response, sensitive_values=sensitive_values)
|
||||
|
||||
|
||||
def _timeout_value(timeout: float | str | httpx.Timeout | None) -> float | httpx.Timeout | None:
|
||||
if not isinstance(timeout, str):
|
||||
return timeout
|
||||
try:
|
||||
return float(timeout)
|
||||
except ValueError:
|
||||
raise Microsoft365CopilotError(status_code=400, message="timeout must be a number") from None
|
||||
|
||||
|
||||
def completion(
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
api_key: str | None,
|
||||
litellm_params: Mapping[str, object],
|
||||
optional_params: Mapping[str, object],
|
||||
secret_fields: SecretFields | None = None,
|
||||
stream: bool = False,
|
||||
timeout: float | str | httpx.Timeout | None = None,
|
||||
logging_obj: Logging | None = None,
|
||||
client: HTTPHandler | None = None,
|
||||
) -> ModelResponse | CustomStreamWrapper:
|
||||
request_data: Final = build_chat_request(
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
http_client: Final = client if client is not None else _get_sync_http_client({})
|
||||
timeout_value: Final = _timeout_value(timeout)
|
||||
_pre_call(logging_obj=logging_obj, messages=messages, model=model, request_data=request_data)
|
||||
access_token: Final = _resolve_access_token(
|
||||
client=http_client,
|
||||
api_key=api_key,
|
||||
litellm_params=litellm_params,
|
||||
secret_fields=secret_fields,
|
||||
timeout=timeout_value,
|
||||
)
|
||||
assertion: Final = extract_caller_assertion(secret_fields)
|
||||
configured_client_secret: Final[object] = litellm_params.get("client_secret")
|
||||
sensitive_client_secret: Final = configured_client_secret if isinstance(configured_client_secret, str) else ""
|
||||
sensitive_values: Final = (access_token, api_key or "", assertion or "", sensitive_client_secret)
|
||||
conversation_url: Final = f"{MICROSOFT_GRAPH_BETA_BASE}/copilot/conversations"
|
||||
conversation_response: Final = _post_graph_request(
|
||||
client=http_client,
|
||||
url=conversation_url,
|
||||
access_token=access_token,
|
||||
request_data={},
|
||||
timeout=timeout_value,
|
||||
)
|
||||
_raise_graph_error_if_needed(
|
||||
response=conversation_response,
|
||||
sensitive_values=sensitive_values,
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
request_data=request_data,
|
||||
)
|
||||
conversation_id: Final = parse_graph_conversation_id(_json_body(conversation_response))
|
||||
chat_url: Final = f"{conversation_url}/{quote(conversation_id, safe='')}/chat"
|
||||
chat_response: Final = _post_graph_request(
|
||||
client=http_client,
|
||||
url=chat_url,
|
||||
access_token=access_token,
|
||||
request_data=request_data,
|
||||
timeout=timeout_value,
|
||||
)
|
||||
_raise_graph_error_if_needed(
|
||||
response=chat_response,
|
||||
sensitive_values=sensitive_values,
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
request_data=request_data,
|
||||
)
|
||||
_post_call(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
request_data=request_data,
|
||||
status_code=chat_response.status_code,
|
||||
)
|
||||
response: Final = map_graph_response(
|
||||
graph_response=_json_body(chat_response),
|
||||
model=model,
|
||||
messages=messages,
|
||||
)
|
||||
return _stream_response(response, model, logging_obj) if stream else response
|
||||
|
||||
|
||||
async def acompletion(
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
api_key: str | None,
|
||||
litellm_params: Mapping[str, object],
|
||||
optional_params: Mapping[str, object],
|
||||
secret_fields: SecretFields | None = None,
|
||||
stream: bool = False,
|
||||
timeout: float | str | httpx.Timeout | None = None,
|
||||
logging_obj: Logging | None = None,
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
shared_session: ClientSession | None = None,
|
||||
) -> ModelResponse | CustomStreamWrapper:
|
||||
request_data: Final = build_chat_request(
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
http_client: Final = (
|
||||
client if client is not None else _get_async_http_client(LlmProviders.MICROSOFT_365_COPILOT, {}, shared_session)
|
||||
)
|
||||
timeout_value: Final = _timeout_value(timeout)
|
||||
_pre_call(logging_obj=logging_obj, messages=messages, model=model, request_data=request_data)
|
||||
access_token: Final = await _aresolve_access_token(
|
||||
client=http_client,
|
||||
api_key=api_key,
|
||||
litellm_params=litellm_params,
|
||||
secret_fields=secret_fields,
|
||||
timeout=timeout_value,
|
||||
)
|
||||
assertion: Final = extract_caller_assertion(secret_fields)
|
||||
configured_client_secret: Final[object] = litellm_params.get("client_secret")
|
||||
sensitive_client_secret: Final = configured_client_secret if isinstance(configured_client_secret, str) else ""
|
||||
sensitive_values: Final = (access_token, api_key or "", assertion or "", sensitive_client_secret)
|
||||
conversation_url: Final = f"{MICROSOFT_GRAPH_BETA_BASE}/copilot/conversations"
|
||||
conversation_response: Final = await _apost_graph_request(
|
||||
client=http_client,
|
||||
url=conversation_url,
|
||||
access_token=access_token,
|
||||
request_data={},
|
||||
timeout=timeout_value,
|
||||
)
|
||||
_raise_graph_error_if_needed(
|
||||
response=conversation_response,
|
||||
sensitive_values=sensitive_values,
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
request_data=request_data,
|
||||
)
|
||||
conversation_id: Final = parse_graph_conversation_id(_json_body(conversation_response))
|
||||
chat_url: Final = f"{conversation_url}/{quote(conversation_id, safe='')}/chat"
|
||||
chat_response: Final = await _apost_graph_request(
|
||||
client=http_client,
|
||||
url=chat_url,
|
||||
access_token=access_token,
|
||||
request_data=request_data,
|
||||
timeout=timeout_value,
|
||||
)
|
||||
_raise_graph_error_if_needed(
|
||||
response=chat_response,
|
||||
sensitive_values=sensitive_values,
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
request_data=request_data,
|
||||
)
|
||||
_post_call(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
request_data=request_data,
|
||||
status_code=chat_response.status_code,
|
||||
)
|
||||
response: Final = map_graph_response(
|
||||
graph_response=_json_body(chat_response),
|
||||
model=model,
|
||||
messages=messages,
|
||||
)
|
||||
return _stream_response(response, model, logging_obj) if stream else response
|
||||
332
litellm/llms/microsoft_365_copilot/chat/transformation.py
Normal file
332
litellm/llms/microsoft_365_copilot/chat/transformation.py
Normal file
|
|
@ -0,0 +1,332 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Final,
|
||||
Literal,
|
||||
Protocol,
|
||||
cast, # noqa: TID251 # AllMessageValues variants are read-only TypedDict mappings
|
||||
)
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
import litellm.utils as litellm_utils
|
||||
from litellm.constants import MICROSOFT_365_COPILOT_DEFAULT_TIME_ZONE
|
||||
from litellm.litellm_core_utils.oauth_token_exchange import redact_sensitive_values
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.llms.microsoft_365_copilot.common_utils import Microsoft365CopilotError
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import Choices, Message, ModelResponse, Usage
|
||||
|
||||
|
||||
class _TokenCounter(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
model: str = "",
|
||||
*,
|
||||
messages: Sequence[AllMessageValues | Message] | None = None,
|
||||
text: str | None = None,
|
||||
count_response_tokens: bool | None = False,
|
||||
) -> int: ...
|
||||
|
||||
|
||||
_TOKEN_COUNTER: Final = cast( # cast-ok: only the typed token-counter arguments used here are narrowed locally
|
||||
_TokenCounter, litellm_utils.token_counter
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer
|
||||
|
||||
|
||||
class _TextPart(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
|
||||
|
||||
type: Literal["text"]
|
||||
text: str
|
||||
|
||||
|
||||
class _GraphConversation(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
|
||||
|
||||
id: str
|
||||
|
||||
|
||||
class _GraphMessage(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
|
||||
|
||||
text: str | None = None
|
||||
|
||||
|
||||
class _GraphConversationResponse(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
|
||||
|
||||
messages: Sequence[_GraphMessage] | None = None
|
||||
|
||||
|
||||
class _GraphError(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
|
||||
|
||||
message: str | None = None
|
||||
|
||||
|
||||
class _GraphErrorResponse(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True, strict=True)
|
||||
|
||||
error: _GraphError | None = None
|
||||
|
||||
|
||||
class _GraphRequestMessage(TypedDict):
|
||||
text: ReadOnly[str]
|
||||
|
||||
|
||||
class _GraphAdditionalContext(TypedDict):
|
||||
text: ReadOnly[str]
|
||||
description: ReadOnly[str]
|
||||
|
||||
|
||||
class _GraphLocationHint(TypedDict):
|
||||
timeZone: ReadOnly[str]
|
||||
|
||||
|
||||
class GraphChatRequest(TypedDict):
|
||||
message: ReadOnly[_GraphRequestMessage]
|
||||
locationHint: ReadOnly[_GraphLocationHint]
|
||||
additionalContext: NotRequired[ReadOnly[Sequence[_GraphAdditionalContext]]]
|
||||
|
||||
|
||||
_TEXT_PART_ADAPTER: Final = TypeAdapter(_TextPart)
|
||||
_GRAPH_CONVERSATION_ADAPTER: Final = TypeAdapter(_GraphConversation)
|
||||
_GRAPH_RESPONSE_ADAPTER: Final = TypeAdapter(_GraphConversationResponse)
|
||||
_GRAPH_ERROR_ADAPTER: Final = TypeAdapter(_GraphErrorResponse)
|
||||
_GRAPH_CHAT_REQUEST_ADAPTER: Final = TypeAdapter(GraphChatRequest)
|
||||
_CONTENT_PARTS_ADAPTER: Final = TypeAdapter(list[object])
|
||||
_JSON_VALUE_ADAPTER: Final = TypeAdapter(object)
|
||||
|
||||
|
||||
class Microsoft365CopilotChatConfig(BaseConfig):
|
||||
def get_supported_openai_params(self, model: str) -> list[str]: # mutable-ok: BaseConfig requires a list return
|
||||
return ["stream", "max_tokens", "max_completion_tokens"]
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
headers: Mapping[str, object],
|
||||
) -> dict[str, object]: # mutable-ok: BaseConfig requires a mutable JSON dictionary
|
||||
return {**build_chat_request(messages=messages, optional_params=optional_params)}
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: "Logging",
|
||||
request_data: Mapping[str, object],
|
||||
messages: Sequence[AllMessageValues],
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
encoding: "Tokenizer | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
return map_graph_response(
|
||||
graph_response=_response_json(raw_response.content),
|
||||
model=model,
|
||||
messages=messages,
|
||||
)
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: Mapping[str, object] | httpx.Headers,
|
||||
) -> BaseLLMException:
|
||||
return Microsoft365CopilotError(status_code=status_code, message=error_message)
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
optional_params: Mapping[str, object],
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict[str, object]: # mutable-ok: BaseConfig requires a mutable parameter dictionary
|
||||
return dict(optional_params)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, object],
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict[str, object]: # mutable-ok: BaseConfig requires a mutable header dictionary
|
||||
return dict(headers)
|
||||
|
||||
|
||||
def _content_to_text(content: object) -> str:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if not isinstance(content, list):
|
||||
raise Microsoft365CopilotError(status_code=400, message="only text content is supported")
|
||||
try:
|
||||
content_parts: Final = _CONTENT_PARTS_ADAPTER.validate_python(content)
|
||||
text_parts: Final = tuple(_TEXT_PART_ADAPTER.validate_python(part) for part in content_parts)
|
||||
except ValidationError:
|
||||
raise Microsoft365CopilotError(status_code=400, message="only text content is supported") from None
|
||||
return "\n".join(part.text for part in text_parts)
|
||||
|
||||
|
||||
def _message_content(message: AllMessageValues) -> object:
|
||||
message_fields: Final = cast( # cast-ok: AllMessageValues variants are read-only TypedDict mappings
|
||||
Mapping[str, object],
|
||||
message,
|
||||
)
|
||||
return message_fields.get("content")
|
||||
|
||||
|
||||
def _response_json(response_content: bytes) -> object:
|
||||
try:
|
||||
return _JSON_VALUE_ADAPTER.validate_json(response_content)
|
||||
except ValidationError:
|
||||
return {}
|
||||
|
||||
|
||||
def build_chat_request(
|
||||
messages: Sequence[AllMessageValues],
|
||||
optional_params: Mapping[str, object],
|
||||
) -> GraphChatRequest:
|
||||
last_user_index: Final = next(
|
||||
(index for index in reversed(range(len(messages))) if messages[index]["role"] == "user"),
|
||||
None,
|
||||
)
|
||||
if last_user_index is None:
|
||||
raise Microsoft365CopilotError(status_code=400, message="at least one message must have role 'user'")
|
||||
last_user_message: Final = messages[last_user_index]
|
||||
history: Final = tuple(
|
||||
{
|
||||
"text": _content_to_text(_message_content(message)),
|
||||
"description": f"{message['role']} message",
|
||||
}
|
||||
for index, message in enumerate(messages)
|
||||
if index != last_user_index
|
||||
)
|
||||
requested_time_zone: Final = optional_params.get("time_zone")
|
||||
time_zone: Final = (
|
||||
requested_time_zone
|
||||
if isinstance(requested_time_zone, str) and requested_time_zone
|
||||
else MICROSOFT_365_COPILOT_DEFAULT_TIME_ZONE
|
||||
)
|
||||
request_data: Final = {
|
||||
"message": {"text": _content_to_text(_message_content(last_user_message))},
|
||||
"locationHint": {"timeZone": time_zone},
|
||||
**({"additionalContext": list(history)} if history else {}),
|
||||
}
|
||||
try:
|
||||
return _GRAPH_CHAT_REQUEST_ADAPTER.validate_python(request_data)
|
||||
except ValidationError:
|
||||
raise Microsoft365CopilotError(status_code=400, message="only text content is supported") from None
|
||||
|
||||
|
||||
def _collapse_doubled_reply(text: str) -> str:
|
||||
if not text or len(text) % 2 != 0:
|
||||
return text
|
||||
half: Final = len(text) // 2
|
||||
return text[:half] if text[:half] == text[half:] else text
|
||||
|
||||
|
||||
def map_graph_response(
|
||||
graph_response: object,
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> ModelResponse:
|
||||
try:
|
||||
parsed_response: Final = _GRAPH_RESPONSE_ADAPTER.validate_python(graph_response)
|
||||
except ValidationError:
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=502,
|
||||
message="Microsoft Graph returned an invalid Copilot conversation response",
|
||||
) from None
|
||||
if not parsed_response.messages:
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=502,
|
||||
message="Microsoft Graph returned a Copilot response without messages",
|
||||
)
|
||||
reply: Final = parsed_response.messages[-1].text
|
||||
if reply is None:
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=502,
|
||||
message="Microsoft Graph returned a Copilot response without reply text",
|
||||
)
|
||||
normalized_reply: Final = _collapse_doubled_reply(reply)
|
||||
prompt_tokens: Final = _TOKEN_COUNTER(model=model, messages=messages)
|
||||
completion_tokens: Final = _TOKEN_COUNTER(
|
||||
model=model,
|
||||
text=normalized_reply,
|
||||
count_response_tokens=True,
|
||||
)
|
||||
return ModelResponse(
|
||||
model=model,
|
||||
choices=[
|
||||
Choices(
|
||||
index=0,
|
||||
message=Message(content=normalized_reply, role="assistant"),
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
usage=Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def parse_graph_conversation_id(graph_conversation: object) -> str:
|
||||
try:
|
||||
conversation: Final = _GRAPH_CONVERSATION_ADAPTER.validate_python(graph_conversation)
|
||||
except ValidationError:
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=502,
|
||||
message="Microsoft Graph returned an invalid Copilot conversation response",
|
||||
) from None
|
||||
if not conversation.id:
|
||||
raise Microsoft365CopilotError(
|
||||
status_code=502,
|
||||
message="Microsoft Graph returned a Copilot conversation without an id",
|
||||
)
|
||||
return conversation.id
|
||||
|
||||
|
||||
def _last_message_text(graph_response: object) -> str | None:
|
||||
try:
|
||||
parsed_response: Final = _GRAPH_RESPONSE_ADAPTER.validate_python(graph_response)
|
||||
except ValidationError:
|
||||
return None
|
||||
if not parsed_response.messages:
|
||||
return None
|
||||
return parsed_response.messages[-1].text
|
||||
|
||||
|
||||
def extract_graph_error_message(
|
||||
response_body: object,
|
||||
sensitive_values: tuple[str, ...] = (),
|
||||
) -> str:
|
||||
try:
|
||||
parsed_error: Final = _GRAPH_ERROR_ADAPTER.validate_python(response_body)
|
||||
except ValidationError:
|
||||
return "Microsoft Graph request failed"
|
||||
message: Final = parsed_error.error.message if parsed_error.error is not None else None
|
||||
if message is None:
|
||||
return "Microsoft Graph request failed"
|
||||
try:
|
||||
nested_response: Final = _JSON_VALUE_ADAPTER.validate_json(message)
|
||||
except ValidationError:
|
||||
return redact_sensitive_values(message, sensitive_values)
|
||||
nested_reply: Final = _last_message_text(nested_response)
|
||||
return redact_sensitive_values(nested_reply if nested_reply is not None else message, sensitive_values)
|
||||
47
litellm/llms/microsoft_365_copilot/common_utils.py
Normal file
47
litellm/llms/microsoft_365_copilot/common_utils.py
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
|
||||
|
||||
|
||||
class Microsoft365CopilotError(BaseLLMException):
|
||||
pass
|
||||
|
||||
|
||||
_RAW_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str])
|
||||
_SECRET_FIELDS_ADAPTER: Final = TypeAdapter(SecretFields)
|
||||
|
||||
|
||||
def extract_caller_assertion(secret_fields: Mapping[str, object] | None) -> str | None:
|
||||
if secret_fields is None:
|
||||
return None
|
||||
raw_headers: Final = secret_fields.get("raw_headers")
|
||||
try:
|
||||
headers: Final = _RAW_HEADERS_ADAPTER.validate_python(raw_headers)
|
||||
except ValidationError:
|
||||
return None
|
||||
authorization_values: Final = tuple(value for name, value in headers.items() if name.casefold() == "authorization")
|
||||
if len(authorization_values) != 1:
|
||||
return None
|
||||
authorization_parts: Final = authorization_values[0].split(maxsplit=1)
|
||||
if len(authorization_parts) != 2 or authorization_parts[0].casefold() != "bearer":
|
||||
return None
|
||||
assertion: Final = authorization_parts[1].strip()
|
||||
assertion_parts: Final = tuple(assertion.split("."))
|
||||
if (
|
||||
len(assertion_parts) != 3
|
||||
or any(not part for part in assertion_parts)
|
||||
or any(character.isspace() for character in assertion)
|
||||
):
|
||||
return None
|
||||
return assertion
|
||||
|
||||
|
||||
def as_secret_fields(secret_fields: object) -> SecretFields | None:
|
||||
try:
|
||||
return _SECRET_FIELDS_ADAPTER.validate_python(secret_fields)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -2135,6 +2135,84 @@
|
|||
],
|
||||
"default_model_placeholder": "gpt-3.5-turbo"
|
||||
},
|
||||
{
|
||||
"provider": "MICROSOFT_365_COPILOT",
|
||||
"provider_display_name": "Microsoft 365 Copilot",
|
||||
"litellm_provider": "microsoft_365_copilot",
|
||||
"credential_fields": [
|
||||
{
|
||||
"key": "token_exchange_endpoint",
|
||||
"label": "Token Endpoint URL",
|
||||
"placeholder": "https://login.microsoftonline.com/<tenant-id>/oauth2/v2.0/token",
|
||||
"tooltip": "Your IdP's OAuth 2.0 token endpoint. For Microsoft 365 Copilot this must be Microsoft Entra.",
|
||||
"required": false,
|
||||
"field_type": "text",
|
||||
"options": null,
|
||||
"default_value": null
|
||||
},
|
||||
{
|
||||
"key": "token_exchange_profile",
|
||||
"label": "Exchange Grant",
|
||||
"placeholder": null,
|
||||
"tooltip": "jwt_bearer_obo: on-behalf-of with the JWT bearer grant (Microsoft Entra). rfc8693: standard OAuth 2.0 token exchange.",
|
||||
"required": false,
|
||||
"field_type": "select",
|
||||
"options": ["jwt_bearer_obo", "rfc8693"],
|
||||
"default_value": "jwt_bearer_obo"
|
||||
},
|
||||
{
|
||||
"key": "client_id",
|
||||
"label": "Client ID",
|
||||
"placeholder": null,
|
||||
"tooltip": null,
|
||||
"required": false,
|
||||
"field_type": "text",
|
||||
"options": null,
|
||||
"default_value": null
|
||||
},
|
||||
{
|
||||
"key": "client_secret",
|
||||
"label": "Client Secret",
|
||||
"placeholder": null,
|
||||
"tooltip": null,
|
||||
"required": false,
|
||||
"field_type": "password",
|
||||
"options": null,
|
||||
"default_value": null
|
||||
},
|
||||
{
|
||||
"key": "token_exchange_scope",
|
||||
"label": "Scope",
|
||||
"placeholder": null,
|
||||
"tooltip": null,
|
||||
"required": false,
|
||||
"field_type": "text",
|
||||
"options": null,
|
||||
"default_value": "https://graph.microsoft.com/.default"
|
||||
},
|
||||
{
|
||||
"key": "token_exchange_audience",
|
||||
"label": "Audience",
|
||||
"placeholder": null,
|
||||
"tooltip": "Optional. Sent only with the rfc8693 grant.",
|
||||
"required": false,
|
||||
"field_type": "text",
|
||||
"options": null,
|
||||
"default_value": null
|
||||
},
|
||||
{
|
||||
"key": "api_key",
|
||||
"label": "Delegated Access Token",
|
||||
"placeholder": null,
|
||||
"tooltip": null,
|
||||
"required": false,
|
||||
"field_type": "password",
|
||||
"options": null,
|
||||
"default_value": null
|
||||
}
|
||||
],
|
||||
"default_model_placeholder": "microsoft_365_copilot/chat"
|
||||
},
|
||||
{
|
||||
"provider": "MistralAI",
|
||||
"provider_display_name": "Mistral AI",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
|
|
|
|||
289
tests/unit/litellm_core_utils/test_oauth_token_exchange.py
Normal file
289
tests/unit/litellm_core_utils/test_oauth_token_exchange.py
Normal file
|
|
@ -0,0 +1,289 @@
|
|||
from collections.abc import Callable
|
||||
from dataclasses import replace
|
||||
from typing import Final
|
||||
from urllib.parse import parse_qs
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils import oauth_token_exchange
|
||||
from litellm.litellm_core_utils.oauth_token_exchange import (
|
||||
OAuthTokenExchangeConfig,
|
||||
OAuthTokenExchangeError,
|
||||
TokenExchangeProfile,
|
||||
aexchange_token,
|
||||
build_token_exchange_form,
|
||||
exchange_token,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
|
||||
|
||||
class _AsyncCaptureTransport(httpx.AsyncBaseTransport):
|
||||
def __init__(self, responder: Callable[[httpx.Request], httpx.Response]) -> None:
|
||||
self._responder = responder
|
||||
|
||||
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return self._responder(request)
|
||||
|
||||
|
||||
def _config(
|
||||
*,
|
||||
token_endpoint: str = "https://identity.example.com/oauth2/token",
|
||||
client_id: str = "client-9306",
|
||||
client_secret: str = "secret-9306",
|
||||
profile: TokenExchangeProfile = "rfc8693",
|
||||
scope: str | None = "graph.scope",
|
||||
audience: str | None = "graph-api",
|
||||
) -> OAuthTokenExchangeConfig:
|
||||
return OAuthTokenExchangeConfig(
|
||||
token_endpoint=token_endpoint,
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
profile=profile,
|
||||
scope=scope,
|
||||
audience=audience,
|
||||
)
|
||||
|
||||
|
||||
def _sync_client(
|
||||
status_code: int,
|
||||
payload: object,
|
||||
) -> tuple[HTTPHandler, list[httpx.Request]]:
|
||||
requests: Final[list[httpx.Request]] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
requests.append(request)
|
||||
return httpx.Response(status_code=status_code, json=payload, request=request)
|
||||
|
||||
return HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond))), requests
|
||||
|
||||
|
||||
def _decode_form(request: httpx.Request) -> dict[str, str]:
|
||||
return {name: values[0] for name, values in parse_qs(request.content.decode("utf-8")).items()}
|
||||
|
||||
|
||||
def test_rfc8693_form_contains_exact_fields_and_optional_scope_and_audience() -> None:
|
||||
config: Final = _config()
|
||||
|
||||
form: Final = build_token_exchange_form(config, "subject-token-9306")
|
||||
|
||||
assert form == {
|
||||
"grant_type": "urn:ietf:params:oauth:grant-type:token-exchange",
|
||||
"client_id": "client-9306",
|
||||
"client_secret": "secret-9306",
|
||||
"subject_token": "subject-token-9306",
|
||||
"subject_token_type": "urn:ietf:params:oauth:token-type:access_token",
|
||||
"scope": "graph.scope",
|
||||
"audience": "graph-api",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("scope", "audience", "expected_optional_fields"),
|
||||
[
|
||||
(None, None, {}),
|
||||
("", "", {}),
|
||||
("graph.scope", None, {"scope": "graph.scope"}),
|
||||
(None, "graph-api", {"audience": "graph-api"}),
|
||||
],
|
||||
)
|
||||
def test_rfc8693_form_omits_empty_or_unset_optional_fields(
|
||||
scope: str | None,
|
||||
audience: str | None,
|
||||
expected_optional_fields: dict[str, str],
|
||||
) -> None:
|
||||
config: Final = _config(scope=scope, audience=audience)
|
||||
|
||||
form: Final = build_token_exchange_form(config, "subject-token-9306")
|
||||
|
||||
assert form == {
|
||||
"grant_type": "urn:ietf:params:oauth:grant-type:token-exchange",
|
||||
"client_id": "client-9306",
|
||||
"client_secret": "secret-9306",
|
||||
"subject_token": "subject-token-9306",
|
||||
"subject_token_type": "urn:ietf:params:oauth:token-type:access_token",
|
||||
**expected_optional_fields,
|
||||
}
|
||||
|
||||
|
||||
def test_jwt_bearer_obo_form_requires_scope_and_ignores_audience() -> None:
|
||||
config: Final = _config(profile="jwt_bearer_obo", audience="ignored-audience")
|
||||
|
||||
form: Final = build_token_exchange_form(config, "subject-token-9306")
|
||||
|
||||
assert form == {
|
||||
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
|
||||
"client_id": "client-9306",
|
||||
"client_secret": "secret-9306",
|
||||
"assertion": "subject-token-9306",
|
||||
"scope": "graph.scope",
|
||||
"requested_token_use": "on_behalf_of",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("scope", [None, "", " "])
|
||||
def test_jwt_bearer_obo_rejects_empty_scope(scope: str | None) -> None:
|
||||
with pytest.raises(OAuthTokenExchangeError) as error:
|
||||
_config(profile="jwt_bearer_obo", scope=scope)
|
||||
|
||||
assert error.value.status_code == 400
|
||||
assert error.value.message == "jwt_bearer_obo requires a non-empty token exchange scope"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint",
|
||||
[
|
||||
"http://identity.example.com/token",
|
||||
"https:///token",
|
||||
"https://identity.example.com:invalid/token",
|
||||
"not-a-url",
|
||||
],
|
||||
)
|
||||
def test_config_rejects_invalid_token_endpoint(endpoint: str) -> None:
|
||||
with pytest.raises(OAuthTokenExchangeError) as error:
|
||||
_config(token_endpoint=endpoint)
|
||||
|
||||
assert error.value.status_code == 400
|
||||
assert error.value.message == "token_exchange_endpoint must be an HTTPS URL with a host"
|
||||
|
||||
|
||||
def test_cache_separates_each_config_field_and_subject_token() -> None:
|
||||
base_config: Final = _config(
|
||||
token_endpoint="https://identity.example.com/cache-separation-9306",
|
||||
client_id="cache-client-9306",
|
||||
client_secret="cache-secret-9306",
|
||||
profile="rfc8693",
|
||||
scope="cache-scope-9306",
|
||||
audience="cache-audience-9306",
|
||||
)
|
||||
config_and_subjects: Final = (
|
||||
(base_config, "cache-subject-9306"),
|
||||
(replace(base_config, token_endpoint="https://identity.example.com/other"), "cache-subject-9306"),
|
||||
(replace(base_config, client_id="other-client"), "cache-subject-9306"),
|
||||
(replace(base_config, client_secret="other-secret"), "cache-subject-9306"),
|
||||
(replace(base_config, profile="jwt_bearer_obo"), "cache-subject-9306"),
|
||||
(replace(base_config, scope="other-scope"), "cache-subject-9306"),
|
||||
(replace(base_config, audience="other-audience"), "cache-subject-9306"),
|
||||
(base_config, "other-subject"),
|
||||
)
|
||||
response_count: Final[list[int]] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
response_count.append(len(response_count) + 1)
|
||||
return httpx.Response(
|
||||
status_code=200,
|
||||
json={"access_token": f"cached-token-{len(response_count)}", "expires_in": 3600},
|
||||
request=request,
|
||||
)
|
||||
|
||||
client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond)))
|
||||
returned_tokens: Final = tuple(
|
||||
exchange_token(client, config, subject_token) for config, subject_token in config_and_subjects
|
||||
)
|
||||
repeated_tokens: Final = tuple(
|
||||
exchange_token(client, config, subject_token) for config, subject_token in config_and_subjects
|
||||
)
|
||||
|
||||
assert len(response_count) == len(config_and_subjects)
|
||||
assert returned_tokens == repeated_tokens
|
||||
|
||||
|
||||
def test_exchange_applies_cache_safety_margin() -> None:
|
||||
cache_margin: Final = oauth_token_exchange.OAUTH_TOKEN_EXCHANGE_CACHE_SAFETY_MARGIN_SECONDS
|
||||
no_cache_client, no_cache_requests = _sync_client(
|
||||
200,
|
||||
{"access_token": "ttl-boundary-token-9306", "expires_in": cache_margin},
|
||||
)
|
||||
no_cache_config: Final = _config(token_endpoint="https://identity.example.com/ttl-boundary-9306")
|
||||
cache_client, cache_requests = _sync_client(
|
||||
200,
|
||||
{"access_token": "ttl-cached-token-9306", "expires_in": cache_margin + 60},
|
||||
)
|
||||
cache_config: Final = _config(token_endpoint="https://identity.example.com/ttl-cached-9306")
|
||||
|
||||
exchange_token(no_cache_client, no_cache_config, "ttl-boundary-subject-9306")
|
||||
exchange_token(no_cache_client, no_cache_config, "ttl-boundary-subject-9306")
|
||||
exchange_token(cache_client, cache_config, "ttl-cached-subject-9306")
|
||||
exchange_token(cache_client, cache_config, "ttl-cached-subject-9306")
|
||||
|
||||
assert len(no_cache_requests) == 2
|
||||
assert len(cache_requests) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("status_code", "expected_status"),
|
||||
[(400, 401), (429, 429), (503, 503)],
|
||||
)
|
||||
def test_sync_exchange_maps_error_status_and_redacts_secrets(status_code: int, expected_status: int) -> None:
|
||||
subject_token: Final = "header.secret.subject"
|
||||
client_secret: Final = "secret-error-9306"
|
||||
client, _ = _sync_client(
|
||||
status_code,
|
||||
{
|
||||
"error": "invalid_grant",
|
||||
"error_description": f"subject={subject_token}, secret={client_secret}",
|
||||
},
|
||||
)
|
||||
|
||||
with pytest.raises(OAuthTokenExchangeError) as error:
|
||||
exchange_token(client, _config(client_secret=client_secret), subject_token)
|
||||
|
||||
assert error.value.status_code == expected_status
|
||||
assert error.value.message == "OAuth token exchange failed: invalid_grant: subject=[REDACTED], secret=[REDACTED]"
|
||||
assert client_secret not in error.value.message
|
||||
assert subject_token not in error.value.message
|
||||
|
||||
|
||||
def test_sync_exchange_posts_to_endpoint_with_exact_form() -> None:
|
||||
config: Final = _config(token_endpoint="https://identity.example.com/sync-9306")
|
||||
client, requests = _sync_client(200, {"access_token": "sync-token-9306", "expires_in": 3600})
|
||||
|
||||
token: Final = exchange_token(client, config, "sync-subject-9306")
|
||||
|
||||
assert str(requests[0].url) == config.token_endpoint
|
||||
assert _decode_form(requests[0]) == build_token_exchange_form(config, "sync-subject-9306")
|
||||
assert token == "sync-token-9306"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_exchange_posts_to_endpoint_with_exact_form() -> None:
|
||||
config: Final = _config(token_endpoint="https://identity.example.com/async-9306")
|
||||
requests: Final[list[httpx.Request]] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
requests.append(request)
|
||||
return httpx.Response(
|
||||
status_code=200,
|
||||
json={"access_token": "async-token-9306", "expires_in": 3600},
|
||||
request=request,
|
||||
)
|
||||
|
||||
client: Final = AsyncHTTPHandler(transport=_AsyncCaptureTransport(respond))
|
||||
async with client.client:
|
||||
token: Final = await aexchange_token(client, config, "async-subject-9306")
|
||||
|
||||
assert str(requests[0].url) == config.token_endpoint
|
||||
assert _decode_form(requests[0]) == build_token_exchange_form(config, "async-subject-9306")
|
||||
assert token == "async-token-9306"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("status_code", "expected_status"),
|
||||
[(400, 401), (429, 429), (503, 503)],
|
||||
)
|
||||
async def test_async_exchange_maps_error_status(status_code: int, expected_status: int) -> None:
|
||||
client: Final = AsyncHTTPHandler(
|
||||
transport=_AsyncCaptureTransport(
|
||||
lambda request: httpx.Response(
|
||||
status_code=status_code,
|
||||
json={"error": "invalid_grant", "error_description": "request failed"},
|
||||
request=request,
|
||||
)
|
||||
)
|
||||
)
|
||||
async with client.client:
|
||||
with pytest.raises(OAuthTokenExchangeError) as error:
|
||||
await aexchange_token(client, _config(), "async-status-subject-9306")
|
||||
|
||||
assert error.value.status_code == expected_status
|
||||
|
|
@ -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={},
|
||||
|
|
|
|||
0
tests/unit/llms/microsoft_365_copilot/__init__.py
Normal file
0
tests/unit/llms/microsoft_365_copilot/__init__.py
Normal file
0
tests/unit/llms/microsoft_365_copilot/chat/__init__.py
Normal file
0
tests/unit/llms/microsoft_365_copilot/chat/__init__.py
Normal file
677
tests/unit/llms/microsoft_365_copilot/chat/test_handler.py
Normal file
677
tests/unit/llms/microsoft_365_copilot/chat/test_handler.py
Normal file
|
|
@ -0,0 +1,677 @@
|
|||
import json
|
||||
from collections.abc import Callable, Mapping
|
||||
from copy import deepcopy
|
||||
from typing import Final, cast
|
||||
from urllib.parse import parse_qs
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.microsoft_365_copilot.chat.handler import acompletion, completion
|
||||
from litellm.llms.microsoft_365_copilot.common_utils import Microsoft365CopilotError
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import token_counter
|
||||
|
||||
_MODEL: Final = "microsoft_365_copilot/chat"
|
||||
_MESSAGES_ADAPTER: Final = TypeAdapter(list[AllMessageValues])
|
||||
_DEFAULT_CHAT_RESPONSE: Final = {"messages": [{"text": "prompt echo"}, {"text": "Copilot response"}]}
|
||||
_DEFAULT_OAUTH_RESPONSE: Final = {"access_token": "graph-access-token", "expires_in": 3600}
|
||||
|
||||
|
||||
class _AsyncCaptureTransport(httpx.AsyncBaseTransport):
|
||||
def __init__(self, responder: Callable[[httpx.Request], httpx.Response]) -> None:
|
||||
self._responder = responder
|
||||
|
||||
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return self._responder(request)
|
||||
|
||||
|
||||
def _secret_fields(assertion: str) -> SecretFields:
|
||||
return SecretFields(raw_headers={"Authorization": f"bEaReR {assertion}"})
|
||||
|
||||
|
||||
def _messages() -> list[AllMessageValues]:
|
||||
return _MESSAGES_ADAPTER.validate_python([{"role": "user", "content": "hello Copilot"}])
|
||||
|
||||
|
||||
def _sync_client(
|
||||
oauth_status: int = 200,
|
||||
oauth_payload: object = _DEFAULT_OAUTH_RESPONSE,
|
||||
graph_status: int = 200,
|
||||
graph_payload: object = _DEFAULT_CHAT_RESPONSE,
|
||||
conversation_id: str = "conversation/42",
|
||||
) -> tuple[HTTPHandler, list[httpx.Request]]:
|
||||
requests: Final[list[httpx.Request]] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
requests.append(request)
|
||||
if request.url.host == "identity.example.com":
|
||||
return httpx.Response(
|
||||
status_code=oauth_status,
|
||||
json=oauth_payload,
|
||||
request=request,
|
||||
)
|
||||
if request.url.path == "/beta/copilot/conversations":
|
||||
return httpx.Response(
|
||||
status_code=201,
|
||||
json={"id": conversation_id, "state": "active"},
|
||||
request=request,
|
||||
)
|
||||
if request.url.path.endswith("/chat"):
|
||||
return httpx.Response(
|
||||
status_code=graph_status,
|
||||
json=graph_payload,
|
||||
request=request,
|
||||
)
|
||||
return httpx.Response(status_code=404, json={}, request=request)
|
||||
|
||||
client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond)))
|
||||
return client, requests
|
||||
|
||||
|
||||
def _async_client(
|
||||
oauth_status: int = 200,
|
||||
oauth_payload: object = _DEFAULT_OAUTH_RESPONSE,
|
||||
graph_status: int = 200,
|
||||
graph_payload: object = _DEFAULT_CHAT_RESPONSE,
|
||||
conversation_id: str = "conversation/42",
|
||||
) -> tuple[AsyncHTTPHandler, list[httpx.Request]]:
|
||||
requests: Final[list[httpx.Request]] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
requests.append(request)
|
||||
if request.url.host == "identity.example.com":
|
||||
return httpx.Response(
|
||||
status_code=oauth_status,
|
||||
json=oauth_payload,
|
||||
request=request,
|
||||
)
|
||||
if request.url.path == "/beta/copilot/conversations":
|
||||
return httpx.Response(
|
||||
status_code=201,
|
||||
json={"id": conversation_id, "state": "active"},
|
||||
request=request,
|
||||
)
|
||||
if request.url.path.endswith("/chat"):
|
||||
return httpx.Response(
|
||||
status_code=graph_status,
|
||||
json=graph_payload,
|
||||
request=request,
|
||||
)
|
||||
return httpx.Response(status_code=404, json={}, request=request)
|
||||
|
||||
client: Final = AsyncHTTPHandler(transport=_AsyncCaptureTransport(respond))
|
||||
return client, requests
|
||||
|
||||
|
||||
def _token_exchange_params(
|
||||
exchange_name: str,
|
||||
client_id: str,
|
||||
client_secret: str,
|
||||
) -> Mapping[str, object]:
|
||||
return {
|
||||
"token_exchange_endpoint": f"https://identity.example.com/{exchange_name}/oauth2/v2.0/token",
|
||||
"client_id": client_id,
|
||||
"client_secret": client_secret,
|
||||
"token_exchange_profile": "jwt_bearer_obo",
|
||||
"token_exchange_scope": "https://graph.microsoft.com/.default",
|
||||
}
|
||||
|
||||
|
||||
def _token_requests(requests: list[httpx.Request]) -> tuple[httpx.Request, ...]:
|
||||
return tuple(request for request in requests if request.url.host == "identity.example.com")
|
||||
|
||||
|
||||
def test_sync_jwt_bearer_obo_uses_exact_exchange_and_sends_graph_token() -> None:
|
||||
assertion: Final = "header-sync-9306.payload-sync-9306.signature-sync-9306"
|
||||
params: Final = _token_exchange_params(
|
||||
"sync-handler-9306",
|
||||
"client-sync-handler-9306",
|
||||
"secret-sync-handler-9306",
|
||||
)
|
||||
client, requests = _sync_client()
|
||||
|
||||
response: Final = completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key=None,
|
||||
litellm_params=params,
|
||||
optional_params={},
|
||||
secret_fields=_secret_fields(assertion),
|
||||
client=client,
|
||||
)
|
||||
|
||||
token_request: Final = _token_requests(requests)[0]
|
||||
token_form: Final = {name: values[0] for name, values in parse_qs(token_request.content.decode("utf-8")).items()}
|
||||
graph_requests: Final = tuple(request for request in requests if request.url.host == "graph.microsoft.com")
|
||||
assert str(token_request.url) == "https://identity.example.com/sync-handler-9306/oauth2/v2.0/token"
|
||||
assert token_form == {
|
||||
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
|
||||
"client_id": "client-sync-handler-9306",
|
||||
"client_secret": "secret-sync-handler-9306",
|
||||
"assertion": assertion,
|
||||
"scope": "https://graph.microsoft.com/.default",
|
||||
"requested_token_use": "on_behalf_of",
|
||||
}
|
||||
assert graph_requests[0].method == "POST"
|
||||
assert graph_requests[0].url.path == "/beta/copilot/conversations"
|
||||
assert json.loads(graph_requests[0].content) == {}
|
||||
assert graph_requests[0].headers["Authorization"] == "Bearer graph-access-token"
|
||||
assert graph_requests[1].url.raw_path == b"/beta/copilot/conversations/conversation%2F42/chat"
|
||||
assert json.loads(graph_requests[1].content) == {
|
||||
"message": {"text": "hello Copilot"},
|
||||
"locationHint": {"timeZone": "UTC"},
|
||||
}
|
||||
assert graph_requests[1].headers["Authorization"] == "Bearer graph-access-token"
|
||||
assert assertion not in graph_requests[0].headers["Authorization"]
|
||||
assert isinstance(response, ModelResponse)
|
||||
assert response.choices[0].message.content == "Copilot response"
|
||||
assert response.model == _MODEL
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_jwt_bearer_obo_uses_exact_exchange_and_sends_graph_token() -> None:
|
||||
assertion: Final = "header-async-9306.payload-async-9306.signature-async-9306"
|
||||
params: Final = _token_exchange_params(
|
||||
"async-handler-9306",
|
||||
"client-async-handler-9306",
|
||||
"secret-async-handler-9306",
|
||||
)
|
||||
client, requests = _async_client()
|
||||
|
||||
async with client.client:
|
||||
response: Final = await acompletion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key=None,
|
||||
litellm_params=params,
|
||||
optional_params={},
|
||||
secret_fields=_secret_fields(assertion),
|
||||
client=client,
|
||||
)
|
||||
|
||||
token_request: Final = _token_requests(requests)[0]
|
||||
token_form: Final = {name: values[0] for name, values in parse_qs(token_request.content.decode("utf-8")).items()}
|
||||
graph_requests: Final = tuple(request for request in requests if request.url.host == "graph.microsoft.com")
|
||||
assert str(token_request.url) == "https://identity.example.com/async-handler-9306/oauth2/v2.0/token"
|
||||
assert token_form == {
|
||||
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
|
||||
"client_id": "client-async-handler-9306",
|
||||
"client_secret": "secret-async-handler-9306",
|
||||
"assertion": assertion,
|
||||
"scope": "https://graph.microsoft.com/.default",
|
||||
"requested_token_use": "on_behalf_of",
|
||||
}
|
||||
assert graph_requests[0].headers["Authorization"] == "Bearer graph-access-token"
|
||||
assert graph_requests[1].url.raw_path == b"/beta/copilot/conversations/conversation%2F42/chat"
|
||||
assert graph_requests[1].headers["Authorization"] == "Bearer graph-access-token"
|
||||
assert isinstance(response, ModelResponse)
|
||||
assert response.choices[0].message.content == "Copilot response"
|
||||
|
||||
|
||||
def test_token_exchange_cache_hits_and_keys_by_subject_token_and_client_secret() -> None:
|
||||
client, requests = _sync_client()
|
||||
exchange_name: Final = "cache-handler-9306"
|
||||
client_id: Final = "client-cache-handler-9306"
|
||||
client_secret: Final = "secret-cache-handler-original-9306"
|
||||
first_assertion: Final = "header-cache-one.payload-cache-one.signature-cache-one"
|
||||
second_assertion: Final = "header-cache-two.payload-cache-two.signature-cache-two"
|
||||
|
||||
def call(assertion: str, secret: str) -> None:
|
||||
completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key=None,
|
||||
litellm_params=_token_exchange_params(exchange_name, client_id, secret),
|
||||
optional_params={},
|
||||
secret_fields=_secret_fields(assertion),
|
||||
client=client,
|
||||
)
|
||||
|
||||
call(first_assertion, client_secret)
|
||||
call(first_assertion, client_secret)
|
||||
call(second_assertion, client_secret)
|
||||
call(second_assertion, "secret-cache-handler-changed-9306")
|
||||
|
||||
assert len(_token_requests(requests)) == 3
|
||||
|
||||
|
||||
def test_direct_delegated_graph_token_skips_token_exchange() -> None:
|
||||
client, requests = _sync_client()
|
||||
|
||||
response: Final = completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key="delegated-graph-token-9306",
|
||||
litellm_params={},
|
||||
optional_params={},
|
||||
client=client,
|
||||
)
|
||||
|
||||
graph_requests: Final = tuple(request for request in requests if request.url.host == "graph.microsoft.com")
|
||||
assert len(_token_requests(requests)) == 0
|
||||
assert len(graph_requests) == 2
|
||||
assert all(request.headers["Authorization"] == "Bearer delegated-graph-token-9306" for request in graph_requests)
|
||||
assert isinstance(response, ModelResponse)
|
||||
assert response.choices[0].message.content == "Copilot response"
|
||||
|
||||
|
||||
def test_token_exchange_credentials_take_precedence_over_direct_api_key() -> None:
|
||||
assertion: Final = "header-priority.payload-priority.signature-priority"
|
||||
client, requests = _sync_client()
|
||||
messages: Final = _messages()
|
||||
litellm_params: Final = _token_exchange_params(
|
||||
"priority-9306",
|
||||
"client-priority-9306",
|
||||
"secret-priority-9306",
|
||||
)
|
||||
optional_params: Final = {"time_zone": "UTC"}
|
||||
secret_fields: Final = _secret_fields(assertion)
|
||||
original_messages: Final = deepcopy(messages)
|
||||
original_litellm_params: Final = dict(litellm_params)
|
||||
original_optional_params: Final = dict(optional_params)
|
||||
original_headers: Final = dict(cast(Mapping[str, str], secret_fields["raw_headers"]))
|
||||
|
||||
completion(
|
||||
model=_MODEL,
|
||||
messages=messages,
|
||||
api_key="direct-token-must-not-win-9306",
|
||||
litellm_params=litellm_params,
|
||||
optional_params=optional_params,
|
||||
secret_fields=secret_fields,
|
||||
client=client,
|
||||
)
|
||||
|
||||
graph_requests: Final = tuple(request for request in requests if request.url.host == "graph.microsoft.com")
|
||||
assert len(_token_requests(requests)) == 1
|
||||
assert len(graph_requests) == 2
|
||||
assert all(request.headers["Authorization"] == "Bearer graph-access-token" for request in graph_requests)
|
||||
assert messages == original_messages
|
||||
assert litellm_params == original_litellm_params
|
||||
assert optional_params == original_optional_params
|
||||
assert secret_fields["raw_headers"] == original_headers
|
||||
|
||||
|
||||
def test_partial_token_exchange_configuration_fails_without_http_calls() -> None:
|
||||
client, requests = _sync_client()
|
||||
|
||||
with pytest.raises(Microsoft365CopilotError) as error:
|
||||
completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key="direct-token-must-not-fallback-9306",
|
||||
litellm_params={
|
||||
"token_exchange_endpoint": "https://identity.example.com/partial-9306/token",
|
||||
"client_id": "client-partial-9306",
|
||||
},
|
||||
optional_params={},
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert error.value.status_code == 400
|
||||
assert len(requests) == 0
|
||||
|
||||
|
||||
def test_api_base_override_clears_exchange_config_and_fails_closed_without_token_post() -> None:
|
||||
from litellm.router_utils.clientside_credential_handler import get_dynamic_litellm_params
|
||||
|
||||
assertion: Final = "header-base-override.payload-base-override.signature-base-override"
|
||||
oauth_fields: Final = (
|
||||
"token_exchange_audience",
|
||||
"token_exchange_endpoint",
|
||||
"token_exchange_profile",
|
||||
"token_exchange_scope",
|
||||
)
|
||||
litellm_params: Final = get_dynamic_litellm_params(
|
||||
litellm_params={
|
||||
"model": _MODEL,
|
||||
"api_base": "https://graph.microsoft.com/beta",
|
||||
**_token_exchange_params(
|
||||
"base-override-9306",
|
||||
"client-base-override-9306",
|
||||
"secret-base-override-9306",
|
||||
),
|
||||
"token_exchange_audience": "https://graph.microsoft.com",
|
||||
},
|
||||
request_kwargs={"api_base": "https://caller-controlled.example"},
|
||||
)
|
||||
assert litellm_params["api_base"] == "https://caller-controlled.example"
|
||||
assert all(field not in litellm_params for field in oauth_fields)
|
||||
assert "client_id" not in litellm_params
|
||||
assert "client_secret" not in litellm_params
|
||||
|
||||
client, requests = _sync_client()
|
||||
with pytest.raises(Microsoft365CopilotError) as error:
|
||||
completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key=None,
|
||||
litellm_params=litellm_params,
|
||||
optional_params={},
|
||||
secret_fields=_secret_fields(assertion),
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert error.value.status_code == 401
|
||||
assert _token_requests(requests) == ()
|
||||
|
||||
|
||||
def test_unknown_token_exchange_profile_fails_without_http_calls() -> None:
|
||||
client, requests = _sync_client()
|
||||
litellm_params: Final = {
|
||||
**_token_exchange_params("unknown-profile-9306", "client-unknown-profile-9306", "secret-unknown-profile-9306"),
|
||||
"token_exchange_profile": "unknown",
|
||||
}
|
||||
|
||||
with pytest.raises(Microsoft365CopilotError) as error:
|
||||
completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key=None,
|
||||
litellm_params=litellm_params,
|
||||
optional_params={},
|
||||
secret_fields=_secret_fields("header.unknown.profile"),
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert error.value.status_code == 400
|
||||
assert error.value.message == ("microsoft_365_copilot token_exchange_profile must be 'rfc8693' or 'jwt_bearer_obo'")
|
||||
assert len(requests) == 0
|
||||
|
||||
|
||||
def test_missing_auth_configuration_fails_without_http_calls() -> None:
|
||||
client, requests = _sync_client()
|
||||
|
||||
with pytest.raises(Microsoft365CopilotError) as error:
|
||||
completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key=None,
|
||||
litellm_params={},
|
||||
optional_params={},
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert error.value.status_code == 401
|
||||
assert len(requests) == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"secret_fields",
|
||||
[
|
||||
None,
|
||||
SecretFields(raw_headers={"Authorization": "Bearer caller-api-key"}),
|
||||
],
|
||||
)
|
||||
def test_token_exchange_requires_a_caller_bearer_token(secret_fields: SecretFields | None) -> None:
|
||||
client, requests = _sync_client()
|
||||
|
||||
with pytest.raises(Microsoft365CopilotError) as error:
|
||||
completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key="direct-token-must-not-fallback-9306",
|
||||
litellm_params=_token_exchange_params(
|
||||
"no-assertion-9306",
|
||||
"client-no-assertion-9306",
|
||||
"secret-no-assertion-9306",
|
||||
),
|
||||
optional_params={},
|
||||
secret_fields=secret_fields,
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert error.value.status_code == 401
|
||||
assert str(error.value) == (
|
||||
"microsoft_365_copilot with OAuth token exchange requires the caller's IdP-issued access token "
|
||||
"in the Authorization header"
|
||||
)
|
||||
assert len(requests) == 0
|
||||
|
||||
|
||||
def test_token_exchange_invalid_grant_maps_to_401_without_leaking_credentials() -> None:
|
||||
assertion: Final = "header-invalid-grant.payload-invalid-grant.signature-invalid-grant"
|
||||
client_secret: Final = "secret-invalid-grant-9306"
|
||||
client, requests = _sync_client(
|
||||
oauth_status=400,
|
||||
oauth_payload={
|
||||
"error": "invalid_grant",
|
||||
"error_description": "The assertion is invalid",
|
||||
},
|
||||
)
|
||||
|
||||
with pytest.raises(Microsoft365CopilotError) as error:
|
||||
completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key=None,
|
||||
litellm_params=_token_exchange_params(
|
||||
"invalid-grant-9306",
|
||||
"client-invalid-grant-9306",
|
||||
client_secret,
|
||||
),
|
||||
optional_params={},
|
||||
secret_fields=_secret_fields(assertion),
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert error.value.status_code == 401
|
||||
assert str(error.value) == "OAuth token exchange failed: invalid_grant: The assertion is invalid"
|
||||
assert assertion not in str(error.value)
|
||||
assert client_secret not in str(error.value)
|
||||
assert len(requests) == 1
|
||||
|
||||
|
||||
def test_graph_403_extracts_last_reply_from_stringified_conversation() -> None:
|
||||
license_message: Final = "It looks like you do not have a valid license"
|
||||
graph_error: Final = {
|
||||
"error": {
|
||||
"code": "UnknownError",
|
||||
"message": json.dumps({"messages": [{"text": "prompt echo"}, {"text": license_message}]}),
|
||||
}
|
||||
}
|
||||
client, requests = _sync_client(graph_status=403, graph_payload=graph_error)
|
||||
|
||||
with pytest.raises(Microsoft365CopilotError) as error:
|
||||
completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key="delegated-graph-token-403",
|
||||
litellm_params={},
|
||||
optional_params={},
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert error.value.status_code == 403
|
||||
assert str(error.value) == license_message
|
||||
assert len(requests) == 2
|
||||
|
||||
|
||||
def test_graph_error_redacts_caller_token_client_secret_and_access_token() -> None:
|
||||
assertion: Final = "header-redaction.payload-redaction.signature-redaction"
|
||||
client_secret: Final = "secret-redaction-9306"
|
||||
graph_error: Final = {
|
||||
"error": {
|
||||
"message": f"token=graph-access-token assertion={assertion} secret={client_secret}",
|
||||
}
|
||||
}
|
||||
client, _ = _sync_client(graph_status=403, graph_payload=graph_error)
|
||||
|
||||
with pytest.raises(Microsoft365CopilotError) as error:
|
||||
completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key=None,
|
||||
litellm_params=_token_exchange_params(
|
||||
"redaction-9306",
|
||||
"client-redaction-9306",
|
||||
client_secret,
|
||||
),
|
||||
optional_params={},
|
||||
secret_fields=_secret_fields(assertion),
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert "graph-access-token" not in str(error.value)
|
||||
assert assertion not in str(error.value)
|
||||
assert client_secret not in str(error.value)
|
||||
|
||||
|
||||
def test_streaming_wraps_full_reply_as_fake_stream() -> None:
|
||||
client, _ = _sync_client()
|
||||
|
||||
response: Final = litellm.completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key="delegated-stream-token-9306",
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert isinstance(response, CustomStreamWrapper)
|
||||
chunks: Final = tuple(response)
|
||||
content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
|
||||
assert content == "Copilot response"
|
||||
usage_chunk: Final = chunks[-1]
|
||||
assert usage_chunk.usage.prompt_tokens == token_counter(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
)
|
||||
assert usage_chunk.usage.completion_tokens == token_counter(
|
||||
model=_MODEL,
|
||||
text="Copilot response",
|
||||
count_response_tokens=True,
|
||||
)
|
||||
assert usage_chunk.usage.total_tokens == (usage_chunk.usage.prompt_tokens + usage_chunk.usage.completion_tokens)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_streaming_wraps_full_reply_as_fake_stream() -> None:
|
||||
client, _ = _async_client()
|
||||
|
||||
async with client.client:
|
||||
response: Final = await litellm.acompletion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key="delegated-async-stream-token-9306",
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert isinstance(response, CustomStreamWrapper)
|
||||
chunks: Final = tuple([chunk async for chunk in response])
|
||||
contents: Final = tuple(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
|
||||
assert "".join(contents) == "Copilot response"
|
||||
usage_chunk: Final = chunks[-1]
|
||||
assert usage_chunk.usage.prompt_tokens == token_counter(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
)
|
||||
assert usage_chunk.usage.completion_tokens == token_counter(
|
||||
model=_MODEL,
|
||||
text="Copilot response",
|
||||
count_response_tokens=True,
|
||||
)
|
||||
assert usage_chunk.usage.total_tokens == (usage_chunk.usage.prompt_tokens + usage_chunk.usage.completion_tokens)
|
||||
|
||||
|
||||
def test_litellm_completion_forwards_token_exchange_settings() -> None:
|
||||
assertion: Final = "header-dispatch-9306.payload-dispatch-9306.signature-dispatch-9306"
|
||||
client, requests = _sync_client()
|
||||
messages: Final[list[AllMessageValues]] = cast(
|
||||
list[AllMessageValues],
|
||||
[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "dispatch"},
|
||||
{"type": "text", "text": "path"},
|
||||
],
|
||||
}
|
||||
],
|
||||
)
|
||||
original_messages: Final = deepcopy(messages)
|
||||
headers: Final[dict[str, str]] = {"X-Caller-Header": "unchanged"}
|
||||
original_headers: Final = dict(headers)
|
||||
|
||||
response: Final = litellm.completion(
|
||||
model=_MODEL,
|
||||
messages=messages,
|
||||
api_key="direct-dispatch-token-must-not-win-9306",
|
||||
api_base="https://caller-controlled.example",
|
||||
extra_headers=headers,
|
||||
token_exchange_endpoint="https://identity.example.com/dispatch-9306/oauth2/v2.0/token",
|
||||
token_exchange_profile="rfc8693",
|
||||
token_exchange_scope="scope-dispatch-9306",
|
||||
token_exchange_audience="audience-dispatch-9306",
|
||||
client_id="client-dispatch-9306",
|
||||
client_secret="secret-dispatch-9306",
|
||||
secret_fields=_secret_fields(assertion),
|
||||
time_zone="America/New_York",
|
||||
client=client,
|
||||
)
|
||||
|
||||
graph_requests: Final = tuple(request for request in requests if request.url.host == "graph.microsoft.com")
|
||||
token_requests: Final = _token_requests(requests)
|
||||
assert isinstance(response, ModelResponse)
|
||||
assert len(graph_requests) == 2
|
||||
assert len(token_requests) == 1
|
||||
assert str(token_requests[0].url) == "https://identity.example.com/dispatch-9306/oauth2/v2.0/token"
|
||||
assert {name: values[0] for name, values in parse_qs(token_requests[0].content.decode("utf-8")).items()} == {
|
||||
"grant_type": "urn:ietf:params:oauth:grant-type:token-exchange",
|
||||
"client_id": "client-dispatch-9306",
|
||||
"client_secret": "secret-dispatch-9306",
|
||||
"subject_token": assertion,
|
||||
"subject_token_type": "urn:ietf:params:oauth:token-type:access_token",
|
||||
"scope": "scope-dispatch-9306",
|
||||
"audience": "audience-dispatch-9306",
|
||||
}
|
||||
assert all(request.headers["Authorization"] == "Bearer graph-access-token" for request in graph_requests)
|
||||
assert assertion not in graph_requests[0].headers["Authorization"]
|
||||
assert json.loads(graph_requests[1].content) == {
|
||||
"message": {"text": "dispatch\npath"},
|
||||
"locationHint": {"timeZone": "America/New_York"},
|
||||
}
|
||||
assert messages == original_messages
|
||||
assert headers == original_headers
|
||||
|
||||
|
||||
def test_litellm_completion_accepts_max_tokens_without_sending_it_to_graph() -> None:
|
||||
client, requests = _sync_client()
|
||||
|
||||
response: Final = litellm.completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key="delegated-max-token",
|
||||
max_tokens=16,
|
||||
max_completion_tokens=32,
|
||||
client=client,
|
||||
)
|
||||
|
||||
graph_chat_requests: Final = tuple(request for request in requests if request.url.path.endswith("/chat"))
|
||||
assert isinstance(response, ModelResponse)
|
||||
assert json.loads(graph_chat_requests[-1].content) == {
|
||||
"message": {"text": "hello Copilot"},
|
||||
"locationHint": {"timeZone": "UTC"},
|
||||
}
|
||||
|
||||
|
||||
def test_litellm_completion_still_rejects_temperature() -> None:
|
||||
client, _ = _sync_client()
|
||||
|
||||
with pytest.raises(litellm.UnsupportedParamsError):
|
||||
litellm.completion(
|
||||
model=_MODEL,
|
||||
messages=_messages(),
|
||||
api_key="delegated-temperature",
|
||||
temperature=0.2,
|
||||
client=client,
|
||||
)
|
||||
|
|
@ -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)
|
||||
27
tests/unit/llms/microsoft_365_copilot/test_common_utils.py
Normal file
27
tests/unit/llms/microsoft_365_copilot/test_common_utils.py
Normal file
|
|
@ -0,0 +1,27 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.microsoft_365_copilot.common_utils import extract_caller_assertion
|
||||
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("authorization", "expected"),
|
||||
[
|
||||
("bEaReR header.payload.signature", "header.payload.signature"),
|
||||
("Bearer header..signature", None),
|
||||
("Bearer not-a-jwt", None),
|
||||
("Basic header.payload.signature", None),
|
||||
("Bearer header.payload.signature extra", None),
|
||||
],
|
||||
)
|
||||
def test_extract_caller_assertion_requires_a_three_part_bearer_jwt(
|
||||
authorization: str,
|
||||
expected: str | None,
|
||||
) -> None:
|
||||
secret_fields: Final[SecretFields] = SecretFields(raw_headers={"aUtHoRiZaTiOn": authorization})
|
||||
|
||||
assertion: Final = extract_caller_assertion(secret_fields)
|
||||
|
||||
assert assertion == expected
|
||||
|
|
@ -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``),
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)])
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"}],
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -1,12 +1,22 @@
|
|||
import { fireEvent, renderHook, screen, waitFor, within, renderWithProviders } from "../../../tests/test-utils";
|
||||
import {
|
||||
chooseSelectOption,
|
||||
fireEvent,
|
||||
renderHook,
|
||||
screen,
|
||||
waitFor,
|
||||
within,
|
||||
renderWithProviders,
|
||||
} from "../../../tests/test-utils";
|
||||
import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import type { Team } from "../key_team_helpers/key_list";
|
||||
import { credentialCreateCall, type CredentialItem } from "../networking";
|
||||
import { credentialCreateCall, modelCreateCall, type CredentialItem } from "../networking";
|
||||
import { Providers } from "../provider_info_helpers";
|
||||
import { projectMountedValues, useMountRegistry, type MountedFormValues } from "../common_components/MountedFormField";
|
||||
import { useForm } from "react-hook-form";
|
||||
import AddModelForm from "./AddModelForm";
|
||||
import { handleAddModelSubmit } from "./handle_add_model_submit";
|
||||
import { toast } from "@/lib/toast";
|
||||
|
||||
vi.mock("../molecules/models/ProviderLogo", () => ({
|
||||
ProviderLogo: ({ provider, className }: { provider: string; className?: string }) => (
|
||||
|
|
@ -35,6 +45,7 @@ vi.mock("../networking", async () => {
|
|||
}),
|
||||
testConnectionRequest: vi.fn().mockResolvedValue({ status: "success" }),
|
||||
credentialCreateCall: vi.fn().mockResolvedValue({ success: true }),
|
||||
modelCreateCall: vi.fn().mockResolvedValue({ success: true }),
|
||||
getProviderCreateMetadata: vi.fn().mockResolvedValue([
|
||||
{
|
||||
provider: "OpenAI",
|
||||
|
|
@ -43,6 +54,27 @@ vi.mock("../networking", async () => {
|
|||
default_model_placeholder: "gpt-3.5-turbo",
|
||||
credential_fields: [],
|
||||
},
|
||||
{
|
||||
provider: "MICROSOFT_365_COPILOT",
|
||||
provider_display_name: "Microsoft 365 Copilot",
|
||||
litellm_provider: "microsoft_365_copilot",
|
||||
default_model_placeholder: "microsoft_365_copilot/chat",
|
||||
credential_fields: [
|
||||
{ key: "token_exchange_endpoint", label: "Token Endpoint URL", field_type: "text" },
|
||||
{
|
||||
key: "token_exchange_profile",
|
||||
label: "Exchange Grant",
|
||||
field_type: "select",
|
||||
options: ["jwt_bearer_obo", "rfc8693"],
|
||||
default_value: "jwt_bearer_obo",
|
||||
},
|
||||
{ key: "client_id", label: "Client ID", field_type: "text" },
|
||||
{ key: "client_secret", label: "Client Secret", field_type: "password" },
|
||||
{ key: "token_exchange_scope", label: "Scope", field_type: "text" },
|
||||
{ key: "token_exchange_audience", label: "Audience", field_type: "text" },
|
||||
{ key: "api_key", label: "Delegated Access Token", field_type: "password" },
|
||||
],
|
||||
},
|
||||
]),
|
||||
};
|
||||
});
|
||||
|
|
@ -57,6 +89,27 @@ vi.mock("@/app/(dashboard)/hooks/providers/useProviderFields", () => ({
|
|||
default_model_placeholder: "gpt-3.5-turbo",
|
||||
credential_fields: [],
|
||||
},
|
||||
{
|
||||
provider: "MICROSOFT_365_COPILOT",
|
||||
provider_display_name: "Microsoft 365 Copilot",
|
||||
litellm_provider: "microsoft_365_copilot",
|
||||
default_model_placeholder: "microsoft_365_copilot/chat",
|
||||
credential_fields: [
|
||||
{ key: "token_exchange_endpoint", label: "Token Endpoint URL", field_type: "text" },
|
||||
{
|
||||
key: "token_exchange_profile",
|
||||
label: "Exchange Grant",
|
||||
field_type: "select",
|
||||
options: ["jwt_bearer_obo", "rfc8693"],
|
||||
default_value: "jwt_bearer_obo",
|
||||
},
|
||||
{ key: "client_id", label: "Client ID", field_type: "text" },
|
||||
{ key: "client_secret", label: "Client Secret", field_type: "password" },
|
||||
{ key: "token_exchange_scope", label: "Scope", field_type: "text" },
|
||||
{ key: "token_exchange_audience", label: "Audience", field_type: "text" },
|
||||
{ key: "api_key", label: "Delegated Access Token", field_type: "password" },
|
||||
],
|
||||
},
|
||||
],
|
||||
isLoading: false,
|
||||
error: null,
|
||||
|
|
@ -501,6 +554,128 @@ describe("AddModelForm", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("credential-only provider auth types", () => {
|
||||
const renderAsAdmin = async () => {
|
||||
const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized"));
|
||||
mockUseAuthorized.default.mockReturnValue(mockAuthorizedUser("proxy_admin", "user-1", true));
|
||||
const props = createTestProps();
|
||||
props.selectedProvider = Providers.MICROSOFT_365_COPILOT;
|
||||
renderWithProviders(<AddModelForm {...props} />);
|
||||
await screen.findByText("Existing Credentials");
|
||||
return props;
|
||||
};
|
||||
|
||||
const submitModel = async (props: ReturnType<typeof createTestProps>) => {
|
||||
const credentialName = props.form.getValues("litellm_credential_name");
|
||||
const apiKey = props.form.getValues("api_key");
|
||||
const modelValues = {
|
||||
...(credentialName ? { litellm_credential_name: credentialName } : {}),
|
||||
...(apiKey ? { api_key: apiKey } : {}),
|
||||
custom_llm_provider: "MICROSOFT_365_COPILOT",
|
||||
model_mappings: [
|
||||
{
|
||||
public_name: "copilot-model",
|
||||
litellm_model: "microsoft_365_copilot/chat",
|
||||
},
|
||||
],
|
||||
};
|
||||
await handleAddModelSubmit(modelValues, "test-access-token", { resetFields: vi.fn() });
|
||||
};
|
||||
|
||||
it("creates and attaches an OAuth exchange credential without mounting its fields on the model", async () => {
|
||||
const user = userEvent.setup();
|
||||
vi.mocked(credentialCreateCall).mockClear();
|
||||
vi.mocked(modelCreateCall).mockClear();
|
||||
const props = await renderAsAdmin();
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "Create credential" }));
|
||||
const dialog = await screen.findByRole("dialog");
|
||||
expect(within(dialog).getByPlaceholderText("Select a provider")).toHaveValue(Providers.MICROSOFT_365_COPILOT);
|
||||
expect(within(dialog).getByRole("combobox", { name: "Auth Type:" })).toHaveTextContent(
|
||||
"OAuth token exchange (on-behalf-of)",
|
||||
);
|
||||
await chooseSelectOption(
|
||||
user,
|
||||
within(dialog).getByRole("combobox", { name: "Exchange Grant" }),
|
||||
"jwt_bearer_obo",
|
||||
);
|
||||
fireEvent.change(within(dialog).getByLabelText("Credential Name:"), { target: { value: "m365-obo" } });
|
||||
fireEvent.change(within(dialog).getByLabelText("Token Endpoint URL"), {
|
||||
target: { value: "https://login.example.test/token" },
|
||||
});
|
||||
fireEvent.change(within(dialog).getByLabelText("Client ID"), { target: { value: "client-id" } });
|
||||
fireEvent.change(within(dialog).getByLabelText("Client Secret"), { target: { value: "client-secret" } });
|
||||
fireEvent.change(within(dialog).getByLabelText("Scope"), { target: { value: "scope" } });
|
||||
fireEvent.change(within(dialog).getByLabelText("Audience"), { target: { value: "audience" } });
|
||||
await user.click(within(dialog).getByRole("button", { name: "Add Credential" }));
|
||||
|
||||
await waitFor(() =>
|
||||
expect(credentialCreateCall).toHaveBeenCalledWith("test-access-token", {
|
||||
credential_name: "m365-obo",
|
||||
credential_values: {
|
||||
token_exchange_endpoint: "https://login.example.test/token",
|
||||
token_exchange_profile: "jwt_bearer_obo",
|
||||
client_id: "client-id",
|
||||
client_secret: "client-secret",
|
||||
token_exchange_scope: "scope",
|
||||
token_exchange_audience: "audience",
|
||||
},
|
||||
credential_info: { custom_llm_provider: Providers.MICROSOFT_365_COPILOT },
|
||||
}),
|
||||
);
|
||||
expect(props.form.getValues("litellm_credential_name")).toBe("m365-obo");
|
||||
expect(props.mountedValues()).not.toHaveProperty("token_exchange_endpoint");
|
||||
|
||||
props.handleOk.mockImplementation(async () => {
|
||||
await submitModel(props);
|
||||
return true;
|
||||
});
|
||||
await user.click(screen.getByRole("button", { name: "Add Model" }));
|
||||
expect(props.handleOk).toHaveBeenCalledOnce();
|
||||
await waitFor(() => expect(modelCreateCall).toHaveBeenCalledOnce());
|
||||
const modelPayload = vi.mocked(modelCreateCall).mock.calls[0][1];
|
||||
expect(modelPayload.litellm_params).toMatchObject({ litellm_credential_name: "m365-obo" });
|
||||
expect(Object.keys(modelPayload.litellm_params).some((key) => key.startsWith("token_exchange_"))).toBe(false);
|
||||
});
|
||||
|
||||
it("blocks model creation when OAuth exchange has no credential", async () => {
|
||||
const props = await renderAsAdmin();
|
||||
const errorToast = vi.spyOn(toast, "error");
|
||||
|
||||
await userEvent.click(screen.getByRole("button", { name: "Add Model" }));
|
||||
|
||||
expect(props.handleOk).not.toHaveBeenCalled();
|
||||
expect(errorToast).toHaveBeenCalledWith("Create or select a credential for this auth type");
|
||||
errorToast.mockRestore();
|
||||
});
|
||||
|
||||
it("keeps static delegated tokens inline in the model request", async () => {
|
||||
const user = userEvent.setup();
|
||||
vi.mocked(modelCreateCall).mockClear();
|
||||
const props = await renderAsAdmin();
|
||||
await chooseSelectOption(
|
||||
user,
|
||||
screen.getByRole("combobox", { name: "Auth Type:" }),
|
||||
"Static delegated access token",
|
||||
);
|
||||
await waitFor(() =>
|
||||
expect(screen.getByRole("combobox", { name: "Auth Type:" })).toHaveTextContent("Static delegated access token"),
|
||||
);
|
||||
fireEvent.change(screen.getByLabelText("Delegated Access Token"), { target: { value: "delegated-token" } });
|
||||
props.handleOk.mockImplementation(async () => {
|
||||
await submitModel(props);
|
||||
return true;
|
||||
});
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "Add Model" }));
|
||||
expect(props.handleOk).toHaveBeenCalledOnce();
|
||||
await waitFor(() => expect(modelCreateCall).toHaveBeenCalledOnce());
|
||||
expect(vi.mocked(modelCreateCall).mock.calls[0][1].litellm_params).toMatchObject({
|
||||
api_key: "delegated-token",
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("cache control bindings reach the parent form store", () => {
|
||||
const renderWithForm = async () => {
|
||||
const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized"));
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ import ConditionalPublicModelName from "./conditional_public_model_name";
|
|||
import LiteLLMModelNameField from "./litellm_model_name";
|
||||
import ConnectionErrorDisplay from "./model_connection_test";
|
||||
import ProviderSpecificFields from "./provider_specific_fields";
|
||||
import { authTypesFor } from "./provider_auth_types";
|
||||
import { TEST_MODES } from "./add_model_modes";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { credentialsKeys } from "@/app/(dashboard)/hooks/credentials/useCredentials";
|
||||
|
|
@ -99,11 +100,16 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
|
|||
const { data: tagsList } = useTags();
|
||||
const selectedCredentialName = useWatch({ control: form.control, name: "litellm_credential_name" });
|
||||
const queryClient = useQueryClient();
|
||||
const [isCredentialModalOpen, setIsCredentialModalOpen] = useState(false);
|
||||
const [credentialModalAuthTypeId, setCredentialModalAuthTypeId] = useState<string | undefined>();
|
||||
const [selectedAuthTypeId, setSelectedAuthTypeId] = useState("");
|
||||
const canCreateCredential = isProxyAdminRole(userRole ?? "");
|
||||
const [isFederatedCredentialModalOpen, setIsFederatedCredentialModalOpen] = useState(false);
|
||||
const canCreateFederatedCredential =
|
||||
isProxyAdminRole(userRole ?? "") && federatedProviderOf(selectedProvider) !== null;
|
||||
const selectedAuthType = authTypesFor(selectedProvider).find(({ id }) => id === selectedAuthTypeId);
|
||||
|
||||
const handleCreateFederatedCredential = async (values: Record<string, unknown>) => {
|
||||
const handleCreateCredential = async (values: Record<string, unknown>) => {
|
||||
const credential = buildCredential(values, withoutRestrictedFields(values));
|
||||
try {
|
||||
await credentialCreateCall(accessToken, credential);
|
||||
|
|
@ -112,11 +118,25 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
|
|||
return;
|
||||
}
|
||||
toast.success("Credential added successfully");
|
||||
setIsCredentialModalOpen(false);
|
||||
setIsFederatedCredentialModalOpen(false);
|
||||
await queryClient.invalidateQueries({ queryKey: credentialsKeys.all });
|
||||
form.setValue("litellm_credential_name", credential.credential_name, { shouldDirty: true });
|
||||
};
|
||||
|
||||
const openCredentialModal = (authTypeId?: string) => {
|
||||
setCredentialModalAuthTypeId(authTypeId);
|
||||
setIsCredentialModalOpen(true);
|
||||
};
|
||||
|
||||
const handleSubmit = async () => {
|
||||
if (selectedAuthType?.credentialOnly && !selectedCredentialName) {
|
||||
toast.error("Create or select a credential for this auth type");
|
||||
return false;
|
||||
}
|
||||
return handleOk();
|
||||
};
|
||||
|
||||
const handleTestConnection = async () => {
|
||||
setIsTestingConnection(true);
|
||||
setConnectionTestId(`test-${Date.now()}`);
|
||||
|
|
@ -188,7 +208,7 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
|
|||
<form
|
||||
onSubmit={(event) => {
|
||||
event.preventDefault();
|
||||
void handleOk().then((submitted) => {
|
||||
void handleSubmit().then((submitted) => {
|
||||
if (submitted) {
|
||||
setTeamAdminSelectedTeam(null);
|
||||
}
|
||||
|
|
@ -331,7 +351,12 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
|
|||
<span className="px-4 text-muted-foreground text-sm">OR</span>
|
||||
<div className="grow border-t border-border"></div>
|
||||
</div>
|
||||
<ProviderSpecificFields selectedProvider={selectedProvider} />
|
||||
<ProviderSpecificFields
|
||||
selectedProvider={selectedProvider}
|
||||
context="model"
|
||||
onAuthTypeChange={setSelectedAuthTypeId}
|
||||
onCreateCredential={canCreateCredential ? openCredentialModal : undefined}
|
||||
/>
|
||||
{canCreateFederatedCredential && (
|
||||
<div className="mb-4 flex flex-col items-start gap-2">
|
||||
<span className="text-sm text-muted-foreground">
|
||||
|
|
@ -476,6 +501,17 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
|
|||
</CardContent>
|
||||
</Card>
|
||||
|
||||
{isCredentialModalOpen && (
|
||||
<CredentialModal
|
||||
open
|
||||
mode="add"
|
||||
initialProvider={selectedProvider}
|
||||
initialAuthTypeId={credentialModalAuthTypeId}
|
||||
providerLocked
|
||||
onCancel={() => setIsCredentialModalOpen(false)}
|
||||
onSubmit={handleCreateCredential}
|
||||
/>
|
||||
)}
|
||||
{isFederatedCredentialModalOpen && (
|
||||
<CredentialModal
|
||||
open
|
||||
|
|
@ -484,7 +520,7 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
|
|||
initialAuthMethod="federation"
|
||||
providerLocked
|
||||
onCancel={() => setIsFederatedCredentialModalOpen(false)}
|
||||
onSubmit={handleCreateFederatedCredential}
|
||||
onSubmit={handleCreateCredential}
|
||||
/>
|
||||
)}
|
||||
{/* Test Connection Results Modal */}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,81 @@
|
|||
import { describe, expect, it } from "vitest";
|
||||
|
||||
import {
|
||||
PROVIDER_AUTH_TYPES,
|
||||
authTypeFieldKeys,
|
||||
authTypesFor,
|
||||
hiddenAuthFieldKeys,
|
||||
inferAuthTypeId,
|
||||
type ProviderAuthType,
|
||||
} from "./provider_auth_types";
|
||||
import { Providers } from "../provider_info_helpers";
|
||||
|
||||
describe("provider auth types", () => {
|
||||
it("resolves provider enum keys and display names", () => {
|
||||
expect(authTypesFor("MICROSOFT_365_COPILOT")).toEqual(PROVIDER_AUTH_TYPES.MICROSOFT_365_COPILOT);
|
||||
expect(authTypesFor(Providers.MICROSOFT_365_COPILOT)).toEqual(PROVIDER_AUTH_TYPES.MICROSOFT_365_COPILOT);
|
||||
expect(authTypesFor(null)).toEqual([]);
|
||||
expect(authTypesFor("unknown")).toEqual([]);
|
||||
});
|
||||
|
||||
it("returns the unique field keys for every auth type", () => {
|
||||
expect(authTypeFieldKeys("MICROSOFT_365_COPILOT")).toEqual([
|
||||
"token_exchange_endpoint",
|
||||
"token_exchange_profile",
|
||||
"client_id",
|
||||
"client_secret",
|
||||
"token_exchange_scope",
|
||||
"token_exchange_audience",
|
||||
"api_key",
|
||||
]);
|
||||
});
|
||||
|
||||
it("infers auth types from required fields and falls back when only optional defaults are present", () => {
|
||||
const authTypes = authTypesFor("MICROSOFT_365_COPILOT");
|
||||
const delegatedValues: Record<string, unknown> = {
|
||||
token_exchange_profile: "jwt_bearer_obo",
|
||||
token_exchange_scope: "https://graph.microsoft.com/.default",
|
||||
api_key: "delegated-token",
|
||||
};
|
||||
const emptyValues: Record<string, unknown> = {
|
||||
token_exchange_endpoint: "",
|
||||
client_id: null,
|
||||
client_secret: undefined,
|
||||
token_exchange_profile: "jwt_bearer_obo",
|
||||
token_exchange_scope: "https://graph.microsoft.com/.default",
|
||||
api_key: "",
|
||||
};
|
||||
|
||||
expect(inferAuthTypeId(authTypes, delegatedValues)).toBe("static_token");
|
||||
expect(inferAuthTypeId(authTypes, emptyValues)).toBe("oauth_token_exchange");
|
||||
expect(inferAuthTypeId([], {})).toBe("");
|
||||
});
|
||||
|
||||
it("hides only non-selected keys not shared with the selected type", () => {
|
||||
const authTypes: readonly ProviderAuthType[] = [
|
||||
{
|
||||
id: "first",
|
||||
label: "First",
|
||||
description: "First auth type",
|
||||
fieldKeys: ["shared", "first_key"],
|
||||
requiredFieldKeys: ["first_key"],
|
||||
},
|
||||
{
|
||||
id: "second",
|
||||
label: "Second",
|
||||
description: "Second auth type",
|
||||
fieldKeys: ["shared", "second_key"],
|
||||
requiredFieldKeys: ["second_key"],
|
||||
},
|
||||
{
|
||||
id: "third",
|
||||
label: "Third",
|
||||
description: "Third auth type",
|
||||
fieldKeys: ["third_key"],
|
||||
requiredFieldKeys: ["third_key"],
|
||||
},
|
||||
];
|
||||
|
||||
expect(hiddenAuthFieldKeys(authTypes, "first")).toEqual(["second_key", "third_key"]);
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,76 @@
|
|||
import { Providers } from "../provider_info_helpers";
|
||||
|
||||
export interface ProviderAuthType {
|
||||
readonly id: string;
|
||||
readonly label: string;
|
||||
readonly description: string;
|
||||
readonly fieldKeys: readonly string[];
|
||||
readonly requiredFieldKeys: readonly string[];
|
||||
readonly fixedValues?: Readonly<Record<string, string>>;
|
||||
readonly credentialOnly?: boolean;
|
||||
}
|
||||
|
||||
const EMPTY_PROVIDER_AUTH_TYPES: readonly ProviderAuthType[] = [];
|
||||
|
||||
export const PROVIDER_AUTH_TYPES: Partial<Record<keyof typeof Providers, readonly ProviderAuthType[]>> = {
|
||||
MICROSOFT_365_COPILOT: [
|
||||
{
|
||||
id: "oauth_token_exchange",
|
||||
label: "OAuth token exchange (on-behalf-of)",
|
||||
credentialOnly: true,
|
||||
description:
|
||||
"LiteLLM exchanges each caller's IdP-issued JWT at your IdP's token endpoint for a delegated token. Microsoft Graph only accepts Microsoft Entra tokens, so for Microsoft 365 Copilot the token endpoint must be Entra and callers must send an Entra-issued JWT for this app. Users can still sign in through any IdP federated with Entra.",
|
||||
fieldKeys: [
|
||||
"token_exchange_endpoint",
|
||||
"token_exchange_profile",
|
||||
"client_id",
|
||||
"client_secret",
|
||||
"token_exchange_scope",
|
||||
"token_exchange_audience",
|
||||
],
|
||||
requiredFieldKeys: ["token_exchange_endpoint", "client_id", "client_secret"],
|
||||
},
|
||||
{
|
||||
id: "static_token",
|
||||
label: "Static delegated access token",
|
||||
description: "Sends one pre-acquired delegated Microsoft Graph token for every caller.",
|
||||
fieldKeys: ["api_key"],
|
||||
requiredFieldKeys: ["api_key"],
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
export const authTypesFor = (provider: string | null): readonly ProviderAuthType[] => {
|
||||
const providerKeys: readonly (keyof typeof Providers)[] = Object.keys(Providers) as (keyof typeof Providers)[];
|
||||
const providerKey =
|
||||
provider === null ? undefined : providerKeys.find((key) => key === provider || Providers[key] === provider);
|
||||
|
||||
return providerKey === undefined
|
||||
? EMPTY_PROVIDER_AUTH_TYPES
|
||||
: PROVIDER_AUTH_TYPES[providerKey] ?? EMPTY_PROVIDER_AUTH_TYPES;
|
||||
};
|
||||
|
||||
export const authTypeFieldKeys = (provider: string | null): readonly string[] =>
|
||||
authTypesFor(provider)
|
||||
.flatMap(({ fieldKeys }) => fieldKeys)
|
||||
.filter((fieldKey, index, fieldKeys) => fieldKeys.indexOf(fieldKey) === index);
|
||||
|
||||
export const inferAuthTypeId = (authTypes: readonly ProviderAuthType[], values: Record<string, unknown>): string =>
|
||||
authTypes.find(({ requiredFieldKeys }) =>
|
||||
requiredFieldKeys.some(
|
||||
(fieldKey) => values[fieldKey] !== undefined && values[fieldKey] !== null && values[fieldKey] !== "",
|
||||
),
|
||||
)?.id ??
|
||||
authTypes[0]?.id ??
|
||||
"";
|
||||
|
||||
export const hiddenAuthFieldKeys = (authTypes: readonly ProviderAuthType[], selectedId: string): readonly string[] => {
|
||||
const selectedFieldKeys = authTypes.find(({ id }) => id === selectedId)?.fieldKeys ?? [];
|
||||
|
||||
return authTypes
|
||||
.filter(({ id }) => id !== selectedId)
|
||||
.flatMap(({ fieldKeys }) => fieldKeys)
|
||||
.filter(
|
||||
(fieldKey, index, fieldKeys) => !selectedFieldKeys.includes(fieldKey) && fieldKeys.indexOf(fieldKey) === index,
|
||||
);
|
||||
};
|
||||
|
|
@ -1,7 +1,10 @@
|
|||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import type { ReactNode } from "react";
|
||||
import { beforeAll, describe, expect, it, vi } from "vitest";
|
||||
import { useFormContext } from "react-hook-form";
|
||||
import { chooseSelectOption } from "@/../tests/test-utils";
|
||||
import { Providers } from "../provider_info_helpers";
|
||||
import type { MountedFormValues } from "../common_components/MountedFormField";
|
||||
import { MountedFormHost } from "../../../tests/mounted-form-host";
|
||||
|
|
@ -109,6 +112,41 @@ vi.mock("../networking", async () => {
|
|||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
provider: "MICROSOFT_365_COPILOT",
|
||||
provider_display_name: Providers.MICROSOFT_365_COPILOT,
|
||||
litellm_provider: "microsoft_365_copilot",
|
||||
default_model_placeholder: "microsoft_365_copilot/chat",
|
||||
credential_fields: [
|
||||
{
|
||||
key: "token_exchange_endpoint",
|
||||
label: "Token Endpoint URL",
|
||||
placeholder: "https://login.microsoftonline.com/<tenant-id>/oauth2/v2.0/token",
|
||||
field_type: "text",
|
||||
},
|
||||
{
|
||||
key: "token_exchange_profile",
|
||||
label: "Exchange Grant",
|
||||
field_type: "select",
|
||||
options: ["jwt_bearer_obo", "rfc8693"],
|
||||
default_value: "jwt_bearer_obo",
|
||||
},
|
||||
{ key: "client_id", label: "Client ID", field_type: "text" },
|
||||
{ key: "client_secret", label: "Client Secret", field_type: "password" },
|
||||
{
|
||||
key: "token_exchange_scope",
|
||||
label: "Scope",
|
||||
field_type: "text",
|
||||
default_value: "https://graph.microsoft.com/.default",
|
||||
},
|
||||
{
|
||||
key: "token_exchange_audience",
|
||||
label: "Audience",
|
||||
field_type: "text",
|
||||
},
|
||||
{ key: "api_key", label: "Delegated Access Token", field_type: "password" },
|
||||
],
|
||||
},
|
||||
]),
|
||||
};
|
||||
});
|
||||
|
|
@ -144,6 +182,16 @@ const VertexCredentialsProbe = () => {
|
|||
return <output data-testid="vertex-credentials">{String(watch("vertex_credentials") ?? "")}</output>;
|
||||
};
|
||||
|
||||
const ValidationForm = ({ children }: { readonly children: ReactNode }) => {
|
||||
const form = useFormContext<MountedFormValues>();
|
||||
return (
|
||||
<form onSubmit={form.handleSubmit(() => {})}>
|
||||
{children}
|
||||
<button type="submit">Submit</button>
|
||||
</form>
|
||||
);
|
||||
};
|
||||
|
||||
describe("ProviderSpecificFields", () => {
|
||||
it("reads a picked service-account file into the vertex credentials field", async () => {
|
||||
const queryClient = createQueryClient();
|
||||
|
|
@ -212,6 +260,7 @@ describe("ProviderSpecificFields", () => {
|
|||
|
||||
const apiKeyLabel = await screen.findByLabelText("OpenAI API Key");
|
||||
expect(apiKeyLabel).toBeInTheDocument();
|
||||
expect(screen.queryByLabelText("Auth Type:")).not.toBeInTheDocument();
|
||||
|
||||
const apiBaseInput = screen.getByPlaceholderText("https://api.openai.com/v1");
|
||||
expect(apiBaseInput).toBeInTheDocument();
|
||||
|
|
@ -419,4 +468,192 @@ describe("ProviderSpecificFields", () => {
|
|||
expect(apiVersionInput).toHaveValue("2025-01-01-preview");
|
||||
});
|
||||
});
|
||||
|
||||
it("defaults Microsoft 365 Copilot to OAuth token exchange fields", async () => {
|
||||
const queryClient = createQueryClient();
|
||||
render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MountedFormHost>
|
||||
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent(
|
||||
"OAuth token exchange (on-behalf-of)",
|
||||
);
|
||||
expect(screen.getByLabelText("Token Endpoint URL")).toBeInTheDocument();
|
||||
expect(screen.getByLabelText("Client ID")).toBeInTheDocument();
|
||||
expect(screen.getByLabelText("Client Secret")).toBeInTheDocument();
|
||||
expect(screen.queryByLabelText("Delegated Access Token")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("offers a credential-only OAuth mode in model context", async () => {
|
||||
const onCreateCredential = vi.fn();
|
||||
const queryClient = createQueryClient();
|
||||
render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MountedFormHost>
|
||||
<ProviderSpecificFields
|
||||
selectedProvider="MICROSOFT_365_COPILOT"
|
||||
context="model"
|
||||
onCreateCredential={onCreateCredential}
|
||||
/>
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
expect(
|
||||
await screen.findByText("This auth type is saved as an LLM credential and attached to this model."),
|
||||
).toBeInTheDocument();
|
||||
expect(screen.queryByLabelText("Token Endpoint URL")).not.toBeInTheDocument();
|
||||
await userEvent.click(screen.getByRole("button", { name: "Create credential" }));
|
||||
expect(onCreateCredential).toHaveBeenCalledWith("oauth_token_exchange");
|
||||
});
|
||||
|
||||
it("hides credential-only auth modes from non-admin model forms", async () => {
|
||||
const user = userEvent.setup();
|
||||
const queryClient = createQueryClient();
|
||||
render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MountedFormHost>
|
||||
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" context="model" />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
const authTypeSelect = await screen.findByRole("combobox", { name: "Auth Type:" });
|
||||
expect(authTypeSelect).toHaveTextContent("Static delegated access token");
|
||||
await user.click(authTypeSelect);
|
||||
expect((await screen.findAllByRole("option")).map((option) => option.textContent)).toEqual([
|
||||
"Static delegated access token",
|
||||
]);
|
||||
|
||||
expect(screen.queryByLabelText("Token Endpoint URL")).not.toBeInTheDocument();
|
||||
expect(screen.queryByLabelText("Client ID")).not.toBeInTheDocument();
|
||||
expect(screen.queryByLabelText("Client Secret")).not.toBeInTheDocument();
|
||||
expect(screen.getByLabelText("Delegated Access Token")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("switches Microsoft 365 Copilot fields to a static delegated access token", async () => {
|
||||
const user = userEvent.setup();
|
||||
const queryClient = createQueryClient();
|
||||
render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MountedFormHost>
|
||||
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
await screen.findByLabelText("Token Endpoint URL");
|
||||
await chooseSelectOption(
|
||||
user,
|
||||
screen.getByRole("combobox", { name: "Auth Type:" }),
|
||||
"Static delegated access token",
|
||||
);
|
||||
|
||||
expect(await screen.findByLabelText("Delegated Access Token")).toBeInTheDocument();
|
||||
expect(screen.queryByLabelText("Token Endpoint URL")).not.toBeInTheDocument();
|
||||
expect(screen.queryByLabelText("Client ID")).not.toBeInTheDocument();
|
||||
expect(screen.queryByLabelText("Client Secret")).not.toBeInTheDocument();
|
||||
expect(screen.queryByLabelText("Scope")).not.toBeInTheDocument();
|
||||
expect(screen.queryByLabelText("Audience")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("infers a static delegated token from a preloaded api_key", async () => {
|
||||
const queryClient = createQueryClient();
|
||||
render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MountedFormHost defaultValues={{ api_key: "delegated-token" }}>
|
||||
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent(
|
||||
"Static delegated access token",
|
||||
);
|
||||
expect(screen.getByLabelText("Delegated Access Token")).toHaveValue("delegated-token");
|
||||
expect(screen.queryByLabelText("Token Endpoint URL")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("re-infers the auth type when the selected provider changes", async () => {
|
||||
const user = userEvent.setup();
|
||||
const queryClient = createQueryClient();
|
||||
const { rerender } = render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MountedFormHost defaultValues={{ api_key: "delegated-token" }}>
|
||||
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
await chooseSelectOption(
|
||||
user,
|
||||
await screen.findByRole("combobox", { name: "Auth Type:" }),
|
||||
"OAuth token exchange (on-behalf-of)",
|
||||
);
|
||||
|
||||
rerender(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MountedFormHost defaultValues={{ api_key: "delegated-token" }}>
|
||||
<ProviderSpecificFields selectedProvider="OpenAI" />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
await screen.findByLabelText("OpenAI API Key");
|
||||
|
||||
rerender(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MountedFormHost defaultValues={{ api_key: "delegated-token" }}>
|
||||
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent(
|
||||
"Static delegated access token",
|
||||
);
|
||||
});
|
||||
|
||||
it("requires only the required OAuth token exchange fields", async () => {
|
||||
const user = userEvent.setup();
|
||||
const queryClient = createQueryClient();
|
||||
render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MountedFormHost>
|
||||
<ValidationForm>
|
||||
<ProviderSpecificFields selectedProvider="MICROSOFT_365_COPILOT" />
|
||||
</ValidationForm>
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
const tokenEndpoint = await screen.findByLabelText("Token Endpoint URL");
|
||||
const clientId = screen.getByLabelText("Client ID");
|
||||
const clientSecret = screen.getByLabelText("Client Secret");
|
||||
expect(tokenEndpoint).toBeInTheDocument();
|
||||
expect(clientId).toBeInTheDocument();
|
||||
expect(clientSecret).toBeInTheDocument();
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "Submit" }));
|
||||
|
||||
expect(await screen.findAllByText("Required")).toHaveLength(3);
|
||||
});
|
||||
|
||||
it("does not render an Auth Type selector for OpenAI", async () => {
|
||||
const queryClient = createQueryClient();
|
||||
render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MountedFormHost>
|
||||
<ProviderSpecificFields selectedProvider="OpenAI" />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
await screen.findByLabelText("OpenAI API Key");
|
||||
|
||||
expect(screen.queryByRole("combobox", { name: "Auth Type:" })).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import {
|
|||
type MountedFieldControlProps,
|
||||
type MountedFormValues,
|
||||
} from "../common_components/MountedFormField";
|
||||
import { authTypesFor, hiddenAuthFieldKeys, inferAuthTypeId } from "./provider_auth_types";
|
||||
import { ProviderCredentialFieldMetadata } from "../networking";
|
||||
import { Providers } from "../provider_info_helpers";
|
||||
import { labelWithHint } from "@/components/shared/form/LabelWithHint";
|
||||
|
|
@ -21,6 +22,10 @@ import type { ProviderFieldValidators } from "../model_add/credential_federation
|
|||
interface ProviderSpecificFieldsProps {
|
||||
selectedProvider: string | null;
|
||||
hiddenFieldKeys?: readonly string[];
|
||||
context?: "model" | "credential";
|
||||
onCreateCredential?: (authTypeId: string) => void;
|
||||
initialAuthTypeId?: string;
|
||||
onAuthTypeChange?: (authTypeId: string) => void;
|
||||
fieldValidators?: ProviderFieldValidators;
|
||||
}
|
||||
|
||||
|
|
@ -83,14 +88,40 @@ const mapFieldMetadataToUiField = (field: ProviderCredentialFieldMetadata): Prov
|
|||
|
||||
const providerFieldsByDisplayName: Record<string, ProviderCredentialField[]> = {};
|
||||
|
||||
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
|
||||
const ProviderSpecificFieldsContent: React.FC<ProviderSpecificFieldsProps> = ({
|
||||
selectedProvider,
|
||||
hiddenFieldKeys,
|
||||
context = "credential",
|
||||
onCreateCredential,
|
||||
initialAuthTypeId,
|
||||
onAuthTypeChange,
|
||||
fieldValidators,
|
||||
}) => {
|
||||
const selectedProviderEnum = Providers[selectedProvider as keyof typeof Providers] as Providers;
|
||||
const form = useFormContext<MountedFormValues>();
|
||||
const authTypes = authTypesFor(selectedProvider);
|
||||
const availableAuthTypes = React.useMemo(
|
||||
() =>
|
||||
authTypes.filter(
|
||||
({ credentialOnly }) => context !== "model" || onCreateCredential !== undefined || !credentialOnly,
|
||||
),
|
||||
[authTypes, context, onCreateCredential],
|
||||
);
|
||||
const [selectedAuthTypeId, setSelectedAuthTypeId] = React.useState(() =>
|
||||
initialAuthTypeId && availableAuthTypes.some(({ id }) => id === initialAuthTypeId)
|
||||
? initialAuthTypeId
|
||||
: inferAuthTypeId(availableAuthTypes, form.getValues()),
|
||||
);
|
||||
const selectedAuthType = availableAuthTypes.find(({ id }) => id === selectedAuthTypeId) ?? availableAuthTypes[0];
|
||||
const credentialOnlySelected = context === "model" && selectedAuthType?.credentialOnly === true;
|
||||
const authTypeSelectId = React.useId();
|
||||
const credentialsFileRef = React.useRef<HTMLInputElement>(null);
|
||||
|
||||
React.useEffect(() => {
|
||||
if (selectedAuthType) {
|
||||
onAuthTypeChange?.(selectedAuthType.id);
|
||||
}
|
||||
}, [onAuthTypeChange, selectedAuthType]);
|
||||
const pickCredentialsFile =
|
||||
(onLoaded: (contents: string) => void) => (event: React.ChangeEvent<HTMLInputElement>) => {
|
||||
const file = event.target.files?.[0];
|
||||
|
|
@ -172,11 +203,34 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
|
|||
return mapped;
|
||||
}, [selectedProviderEnum, selectedProvider, providerMetadata]);
|
||||
|
||||
const allFields = React.useMemo(
|
||||
() => (hiddenFieldKeys ? providerFields.filter((field) => !hiddenFieldKeys.includes(field.key)) : providerFields),
|
||||
[providerFields, hiddenFieldKeys],
|
||||
const authTypeHiddenFieldKeys = React.useMemo(
|
||||
() => hiddenAuthFieldKeys(authTypes, selectedAuthType?.id ?? ""),
|
||||
[authTypes, selectedAuthType?.id],
|
||||
);
|
||||
|
||||
const allFields = React.useMemo(() => {
|
||||
if (credentialOnlySelected) {
|
||||
return [];
|
||||
}
|
||||
|
||||
if (availableAuthTypes.length === 0) {
|
||||
return hiddenFieldKeys ? providerFields.filter((field) => !hiddenFieldKeys.includes(field.key)) : providerFields;
|
||||
}
|
||||
|
||||
const fieldsToHide = [...(hiddenFieldKeys ?? []), ...authTypeHiddenFieldKeys];
|
||||
|
||||
return providerFields
|
||||
.filter((field) => !fieldsToHide.includes(field.key))
|
||||
.map((field) => (selectedAuthType?.requiredFieldKeys.includes(field.key) ? { ...field, required: true } : field));
|
||||
}, [
|
||||
providerFields,
|
||||
hiddenFieldKeys,
|
||||
availableAuthTypes,
|
||||
authTypeHiddenFieldKeys,
|
||||
selectedAuthType,
|
||||
credentialOnlySelected,
|
||||
]);
|
||||
|
||||
const hasApiVersionField = React.useMemo(() => allFields.some((field) => field.key === "api_version"), [allFields]);
|
||||
const lastInferredApiVersionRef = React.useRef<string | null>(null);
|
||||
|
||||
|
|
@ -291,6 +345,44 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
|
|||
|
||||
return (
|
||||
<>
|
||||
{availableAuthTypes.length > 0 && selectedAuthType && (
|
||||
<div className="mb-4 flex flex-col gap-2">
|
||||
<label htmlFor={authTypeSelectId} className="text-sm font-medium">
|
||||
{labelWithHint("Auth Type:", "Select how LiteLLM authenticates to this provider.")}
|
||||
</label>
|
||||
<Select
|
||||
items={availableAuthTypes.map(({ id, label }) => ({ value: id, label }))}
|
||||
value={selectedAuthType.id}
|
||||
onValueChange={(value) => setSelectedAuthTypeId(value ?? "")}
|
||||
>
|
||||
<SelectTrigger id={authTypeSelectId} className="w-full">
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{availableAuthTypes.map((authType) => (
|
||||
<SelectItem key={authType.id} value={authType.id}>
|
||||
{authType.label}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
<p className="text-sm text-muted-foreground">{selectedAuthType.description}</p>
|
||||
</div>
|
||||
)}
|
||||
{credentialOnlySelected ? (
|
||||
<>
|
||||
<p className="mb-4 text-sm">This auth type is saved as an LLM credential and attached to this model.</p>
|
||||
<Button type="button" onClick={() => onCreateCredential?.(selectedAuthType.id)}>
|
||||
Create credential
|
||||
</Button>
|
||||
</>
|
||||
) : (
|
||||
Object.entries(selectedAuthType?.fixedValues ?? {}).map(([fieldKey, fixedValue]) => (
|
||||
<MountedFormField key={`${selectedProvider}:${fieldKey}`} name={fieldKey} defaultValue={fixedValue} bare>
|
||||
{() => null}
|
||||
</MountedFormField>
|
||||
))
|
||||
)}
|
||||
{isLoading && allFields.length === 0 && <p className="text-sm mb-2">Loading provider fields...</p>}
|
||||
{loadError && allFields.length === 0 && (
|
||||
<p className="text-sm mb-2 text-destructive">
|
||||
|
|
@ -341,4 +433,8 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
|
|||
);
|
||||
};
|
||||
|
||||
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = (props) => (
|
||||
<ProviderSpecificFieldsContent key={props.selectedProvider ?? "no-provider"} {...props} />
|
||||
);
|
||||
|
||||
export default ProviderSpecificFields;
|
||||
|
|
|
|||
|
|
@ -40,6 +40,26 @@ vi.mock("../networking", async () => {
|
|||
{ key: "api_key", label: "Azure API Key", field_type: "password" },
|
||||
],
|
||||
},
|
||||
{
|
||||
provider: "MICROSOFT_365_COPILOT",
|
||||
provider_display_name: Providers.MICROSOFT_365_COPILOT,
|
||||
litellm_provider: "microsoft_365_copilot",
|
||||
credential_fields: [
|
||||
{ key: "token_exchange_endpoint", label: "Token Endpoint URL", field_type: "text" },
|
||||
{
|
||||
key: "token_exchange_profile",
|
||||
label: "Exchange Grant",
|
||||
field_type: "select",
|
||||
options: ["jwt_bearer_obo", "rfc8693"],
|
||||
default_value: "jwt_bearer_obo",
|
||||
},
|
||||
{ key: "client_id", label: "Client ID", field_type: "text" },
|
||||
{ key: "client_secret", label: "Client Secret", field_type: "password" },
|
||||
{ key: "token_exchange_scope", label: "Scope", field_type: "text" },
|
||||
{ key: "token_exchange_audience", label: "Audience", field_type: "text" },
|
||||
{ key: "api_key", label: "Delegated Access Token", field_type: "password" },
|
||||
],
|
||||
},
|
||||
]),
|
||||
};
|
||||
});
|
||||
|
|
@ -84,6 +104,19 @@ const unknownSourceCredential: CredentialItem = {
|
|||
credential_info: { custom_llm_provider: "anthropic" },
|
||||
};
|
||||
|
||||
const microsoftCopilotCredential: CredentialItem = {
|
||||
credential_name: "microsoft-copilot",
|
||||
credential_values: {
|
||||
token_exchange_endpoint: "https://identity.example.com/stored-token",
|
||||
token_exchange_profile: "jwt_bearer_obo",
|
||||
client_id: "stored-client",
|
||||
client_secret: "stored-secret",
|
||||
token_exchange_scope: "https://graph.microsoft.com/.default",
|
||||
token_exchange_audience: "https://graph.microsoft.com",
|
||||
},
|
||||
credential_info: { custom_llm_provider: "MICROSOFT_365_COPILOT" },
|
||||
};
|
||||
|
||||
const renderModal = (props: Partial<React.ComponentProps<typeof CredentialModal>> = {}) => {
|
||||
const onSubmit = vi.fn();
|
||||
render(
|
||||
|
|
@ -109,6 +142,34 @@ const chooseProvider = async (user: ReturnType<typeof userEvent.setup>, provider
|
|||
};
|
||||
|
||||
describe("CredentialModal with Anthropic workload identity federation", () => {
|
||||
it("deletes stored exchange fields when switching to a static token", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSubmit = renderModal({ mode: "edit", existingCredential: microsoftCopilotCredential });
|
||||
|
||||
expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent(
|
||||
"OAuth token exchange (on-behalf-of)",
|
||||
);
|
||||
await screen.findByLabelText("Token Endpoint URL");
|
||||
await chooseOption(user, /^Auth Type:/, "Static delegated access token");
|
||||
fill("Delegated Access Token", "delegated-graph-token");
|
||||
await user.click(screen.getByRole("button", { name: "Update Credential" }));
|
||||
|
||||
const [values, valuesToDelete] = onSubmit.mock.calls[0];
|
||||
expect(values).toEqual({
|
||||
credential_name: "microsoft-copilot",
|
||||
custom_llm_provider: "MICROSOFT_365_COPILOT",
|
||||
api_key: "delegated-graph-token",
|
||||
});
|
||||
expect([...valuesToDelete].sort()).toEqual([
|
||||
"client_id",
|
||||
"client_secret",
|
||||
"token_exchange_audience",
|
||||
"token_exchange_endpoint",
|
||||
"token_exchange_profile",
|
||||
"token_exchange_scope",
|
||||
]);
|
||||
});
|
||||
|
||||
it("does not carry the previous provider's base URL into the Anthropic form or its payload", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSubmit = renderModal();
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import { SimpleTooltip } from "@/components/ui/tooltip";
|
|||
import { Button } from "@/components/ui/button";
|
||||
import { useState } from "react";
|
||||
import { FormProvider, useForm } from "react-hook-form";
|
||||
import { authTypeFieldKeys } from "../add_model/provider_auth_types";
|
||||
import ProviderSpecificFields from "../add_model/provider_specific_fields";
|
||||
import { requiredRule } from "../common_components/formRules";
|
||||
import { labelWithHint } from "@/components/shared/form/LabelWithHint";
|
||||
|
|
@ -59,6 +60,7 @@ interface CredentialModalProps {
|
|||
existingCredential?: CredentialItem | null;
|
||||
initialProvider?: string | null;
|
||||
initialAuthMethod?: AuthMethod;
|
||||
initialAuthTypeId?: string;
|
||||
providerLocked?: boolean;
|
||||
}
|
||||
|
||||
|
|
@ -106,6 +108,7 @@ export default function CredentialModal({
|
|||
existingCredential = null,
|
||||
initialProvider = null,
|
||||
initialAuthMethod,
|
||||
initialAuthTypeId,
|
||||
providerLocked = false,
|
||||
}: CredentialModalProps) {
|
||||
const isEdit = mode === "edit";
|
||||
|
|
@ -163,10 +166,16 @@ export default function CredentialModal({
|
|||
onSubmit({ ...meta, ...buildCreateCredentialValues(withoutRestrictedFields(values), selection) }, []);
|
||||
return;
|
||||
}
|
||||
const patch = sameProvider(selectedProvider, storedProvider)
|
||||
? buildCredentialPatch(storedValues, withoutRestrictedFields(values), storedSelection, selection)
|
||||
: buildProviderChangePatch(storedValues, withoutRestrictedFields(values), selection);
|
||||
onSubmit({ ...meta, ...patch.credential_values }, patch.credential_values_to_delete);
|
||||
const sameStoredProvider = sameProvider(selectedProvider, storedProvider);
|
||||
const projectedValues = withoutRestrictedFields(values);
|
||||
const patch = sameStoredProvider
|
||||
? buildCredentialPatch(storedValues, projectedValues, storedSelection, selection)
|
||||
: buildProviderChangePatch(storedValues, projectedValues, selection);
|
||||
const authTypeValuesToDelete = sameStoredProvider
|
||||
? authTypeFieldKeys(selectedProvider).filter((key) => key in storedValues && !(key in projectedValues))
|
||||
: [];
|
||||
const valuesToDelete = Array.from(new Set([...patch.credential_values_to_delete, ...authTypeValuesToDelete]));
|
||||
onSubmit({ ...meta, ...patch.credential_values }, valuesToDelete);
|
||||
};
|
||||
|
||||
const closeAndReset = () => {
|
||||
|
|
@ -283,6 +292,7 @@ export default function CredentialModal({
|
|||
<ProviderSpecificFields
|
||||
selectedProvider={selectedProvider}
|
||||
hiddenFieldKeys={selection.authMethod === "federation" ? API_KEY_FIELDS : NO_HIDDEN_FIELDS}
|
||||
initialAuthTypeId={initialAuthTypeId}
|
||||
fieldValidators={providerFieldValidators(selection)}
|
||||
/>
|
||||
|
||||
|
|
|
|||
|
|
@ -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("");
|
||||
|
|
|
|||
|
|
@ -132,6 +132,7 @@ export enum Providers {
|
|||
LM_STUDIO = "Lm Studio",
|
||||
LLAMA = "Meta Llama",
|
||||
MARITALK = "Maritalk",
|
||||
MICROSOFT_365_COPILOT = "Microsoft 365 Copilot",
|
||||
MiniMax = "MiniMax",
|
||||
MistralAI = "Mistral AI",
|
||||
MOONSHOT = "Moonshot",
|
||||
|
|
@ -252,6 +253,7 @@ export const provider_map: Record<string, string> = {
|
|||
LLAMAFILE: "llamafile",
|
||||
LLAMA: "meta_llama",
|
||||
LM_STUDIO: "lm_studio",
|
||||
MICROSOFT_365_COPILOT: "microsoft_365_copilot",
|
||||
MARITALK: "maritalk",
|
||||
MiniMax: "minimax",
|
||||
MistralAI: "mistral",
|
||||
|
|
@ -364,6 +366,7 @@ export const providerLogoMap: Partial<Record<Providers, string>> = {
|
|||
[Providers.LM_STUDIO]: lmstudioLogo.src,
|
||||
[Providers.LLAMA]: metaLlamaLogo.src,
|
||||
[Providers.MiniMax]: minimaxLogo.src,
|
||||
[Providers.MICROSOFT_365_COPILOT]: microsoftAzureLogo.src,
|
||||
[Providers.MistralAI]: mistralLogo.src,
|
||||
[Providers.MOONSHOT]: moonshotLogo.src,
|
||||
[Providers.MORPH]: morphLogo.src,
|
||||
|
|
@ -458,6 +461,7 @@ const providerPlaceholderMap: Partial<Record<Providers, string>> = {
|
|||
[Providers.FalAI]: "fal_ai/fal-ai/flux-pro/v1.1-ultra",
|
||||
[Providers.Google_AI_Studio]: "gemini-pro",
|
||||
[Providers.JinaAI]: "jina_ai/",
|
||||
[Providers.MICROSOFT_365_COPILOT]: "microsoft_365_copilot/chat",
|
||||
[Providers.NVIDIA_RIVA]: "nvidia_riva/nvidia/parakeet-ctc-1_1b-asr",
|
||||
[Providers.Oracle]: "oci/xai.grok-4",
|
||||
[Providers.RunwayML]: "runwayml/gen4_turbo",
|
||||
|
|
|
|||
16
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
16
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue