diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008120000_user_provider_credentials/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008120000_user_provider_credentials/migration.sql new file mode 100644 index 00000000000..ea4fe9bc0ab --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008120000_user_provider_credentials/migration.sql @@ -0,0 +1,18 @@ +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_UserProviderCredentials" ( + "id" TEXT NOT NULL, + "user_id" TEXT NOT NULL, + "credential_name" TEXT NOT NULL, + "provider" TEXT NOT NULL, + "credential_b64" TEXT NOT NULL, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + + CONSTRAINT "LiteLLM_UserProviderCredentials_pkey" PRIMARY KEY ("id") +); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_UserProviderCredentials_user_id_credential_name_key" ON "LiteLLM_UserProviderCredentials"("user_id", "credential_name"); + +-- CreateIndex +CREATE INDEX IF NOT EXISTS "LiteLLM_UserProviderCredentials_credential_name_idx" ON "LiteLLM_UserProviderCredentials"("credential_name"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 35d026bb6f6..f3ebee52dcd 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -453,6 +453,20 @@ model LiteLLM_MCPUserCredentials { @@unique([user_id, server_id]) } +// Per-user provider connections (e.g. GitHub Copilot OAuth) keyed by credential name +model LiteLLM_UserProviderCredentials { + id String @id @default(uuid()) + user_id String + credential_name String + provider String + credential_b64 String + created_at DateTime @default(now()) @map("created_at") + updated_at DateTime @default(now()) @updatedAt @map("updated_at") + + @@unique([user_id, credential_name]) + @@index([credential_name]) +} + // Per-user environment variable values for MCP servers. // values_b64 is an encrypted JSON object: {VAR_NAME: "value", ...}. model LiteLLM_MCPUserEnvVars { diff --git a/litellm/__init__.py b/litellm/__init__.py index d4ab0d648a7..a7c8f0dd5eb 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1396,6 +1396,8 @@ from .integrations import * from .llms.custom_httpx.async_client_cleanup import close_litellm_async_clients from .exceptions import ( AuthenticationError, + CallerCredentialAuthenticationError as CallerCredentialAuthenticationError, + CallerCredentialRateLimitError as CallerCredentialRateLimitError, InvalidRequestError, BadRequestError, ImageFetchError, diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index c25d523a8c3..c97baae293a 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -107,11 +107,15 @@ def _drop_logging_obj_from_kwargs(request_kwargs: dict[str, object]) -> dict[str def _is_response_cache_excluded(model: str | None, kwargs: Mapping[str, object]) -> bool: + from litellm.llms.github_copilot.per_user_auth import is_github_copilot_per_user_request + 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 + (isinstance(custom_llm_provider, str) and custom_llm_provider in RESPONSE_CACHE_EXCLUDED_PROVIDERS) + or model_provider in RESPONSE_CACHE_EXCLUDED_PROVIDERS + or is_github_copilot_per_user_request(kwargs) + ) def _is_chat_completion_cached_dict(cached_result: dict) -> bool: @@ -317,6 +321,7 @@ class LLMCachingHandler: None """ if _is_response_cache_excluded(model=model, kwargs=kwargs): + self.request_kwargs = _drop_logging_obj_from_kwargs(kwargs) return None # Check if caching should be performed BEFORE doing expensive operations @@ -444,6 +449,7 @@ class LLMCachingHandler: # Check if caching should be performed BEFORE doing expensive kwargs copy if _is_response_cache_excluded(model=model, kwargs=kwargs): + self.request_kwargs = _drop_logging_obj_from_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 () @@ -867,11 +873,11 @@ class LLMCachingHandler: new_kwargs.pop("metadata", None) model_value: Final = new_kwargs.get("model") model: Final = model_value if isinstance(model_value, str) else None + self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs) 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) cached_result: object | None = None if call_type == CallTypes.aembedding.value: new_kwargs["input"] = self.handle_kwargs_input_list_or_str(new_kwargs) diff --git a/litellm/constants.py b/litellm/constants.py index 5b202f9281c..4ae4e2f509e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -13,6 +13,13 @@ MICROSOFT_365_COPILOT_DEFAULT_TOKEN_EXCHANGE_SCOPE: Final = "https://graph.micro # 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" +GITHUB_COPILOT_PER_USER_AUTH_TYPE: Final = "per_user_oauth" +GITHUB_COPILOT_AUTH_TYPE_KEY: Final = "github_copilot_auth_type" +GITHUB_COPILOT_USER_TOKEN_SAFETY_MARGIN_SECONDS: Final = 60 +GITHUB_COPILOT_USER_CREDENTIAL_CACHE_TTL_SECONDS: Final = 60 +GITHUB_COPILOT_DEVICE_FLOW_CACHE_PREFIX: Final = "github_copilot_device_flow" +USER_PROVIDER_CREDENTIAL_CACHE_PREFIX: Final = "user_provider_credential" +USER_PROVIDER_CREDENTIAL_NOT_CONNECTED: Final = "__not_connected__" SERVER_STREAMING_CLASSIFICATION_KEY: Final = "litellm_server_streaming_classification" diff --git a/litellm/exceptions.py b/litellm/exceptions.py index b349b1a1bb5..439eadce13d 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -172,6 +172,12 @@ class AuthenticationError(openai.AuthenticationError): return _message +class CallerCredentialAuthenticationError(AuthenticationError): + """401 raised when the CALLING user's own stored provider credential is + missing or rejected (e.g. a per-user OAuth connection). Scoped to one + caller's credential, so routers must not cool down the shared deployment.""" + + # raise when invalid models passed, example gpt-8 class NotFoundError(openai.NotFoundError): def __init__( @@ -531,6 +537,12 @@ class RateLimitError(openai.RateLimitError): return _message +class CallerCredentialRateLimitError(RateLimitError): + """429 raised when a per-user credential exchange is rate limited by the + provider (e.g. GitHub's token endpoint). Scoped to one caller's + credential, so routers must not cool down the shared deployment.""" + + # sub class of rate limit error - meant to give more granularity for error handling context window exceeded errors class ContextWindowExceededError(BadRequestError): def __init__( diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index 627466ba565..a3b64fc22dd 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -293,6 +293,7 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset( "LiteLLM_BackgroundInteractionSettlement", "LiteLLM_ErrorLogs", "LiteLLM_UserNotifications", + "LiteLLM_UserProviderCredentials", "LiteLLM_TeamMembership", "LiteLLM_OrganizationMembership", "LiteLLM_InvitationLink", diff --git a/litellm/litellm_core_utils/credential_accessor.py b/litellm/litellm_core_utils/credential_accessor.py index 7b4e8240b69..95e36a64ef7 100644 --- a/litellm/litellm_core_utils/credential_accessor.py +++ b/litellm/litellm_core_utils/credential_accessor.py @@ -15,7 +15,9 @@ class CredentialAccessor: ) @staticmethod - def get_credential_values(credential_name: str) -> dict: + def get_credential_values( + credential_name: str, + ) -> dict[str, object]: # mutable-ok: callers merge the returned credential map """Safe accessor for credentials.""" credential: Final = CredentialAccessor.find_credential(credential_name) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index 0aeedafb956..04755f9e1e9 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -75,6 +75,8 @@ OPTIONAL_KWARGS_KEYS: Final = ( "itpm", "otpm", "use_xai_oauth", + "github_copilot_auth_type", + "github_copilot_user_session", "fireworks_forward_user_id", PROVIDER_AFFINITY_HEADER_KWARG_KEY, } diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index 724be5510f5..c9e5167e7a9 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -795,13 +795,36 @@ def _get_openai_compatible_provider_info( api_base = api_base or get_secret("GALADRIEL_API_BASE") or "https://api.galadriel.com/v1" dynamic_api_key = api_key or get_secret_str("GALADRIEL_API_KEY") elif custom_llm_provider == "github_copilot": - ( - api_base, - dynamic_api_key, - custom_llm_provider, - ) = litellm.GithubCopilotConfig().get_openai_compatible_provider_info( - model, api_base, api_key, custom_llm_provider + from litellm.llms.github_copilot.common_utils import DEFAULT_GITHUB_COPILOT_API_BASE + from litellm.llms.github_copilot.per_user_auth import ( + github_copilot_per_user_credential_name, + github_copilot_user_session_from, ) + + user_session: Final = github_copilot_user_session_from( + cast("dict[str, object]", litellm_params) # cast-ok: litellm_params arrives as an untyped request dict + ) + if user_session is not None: + api_base = user_session.api_base # rebind-ok: resolves provider args in place + dynamic_api_key = user_session.token # rebind-ok: resolves provider args in place + elif ( + github_copilot_per_user_credential_name( + cast("dict[str, object]", litellm_params) # cast-ok: untyped params dict + ) + is not None + ): + # per-user deployments resolve no token at provider-info time; the + # caller's session supplies base + key per request + api_base = api_base or DEFAULT_GITHUB_COPILOT_API_BASE # rebind-ok: resolves provider args in place + dynamic_api_key = None # rebind-ok: resolves provider args in place + else: + ( + api_base, # rebind-ok: resolves provider args in place + dynamic_api_key, # rebind-ok: resolves provider args in place + custom_llm_provider, # rebind-ok: resolves provider args in place + ) = litellm.GithubCopilotConfig().get_openai_compatible_provider_info( + model, api_base, api_key, custom_llm_provider + ) elif custom_llm_provider == "chatgpt": ( api_base, diff --git a/litellm/llms/anthropic/pass_through/messages/handler.py b/litellm/llms/anthropic/pass_through/messages/handler.py index 4e7a154be67..47744d6e93b 100644 --- a/litellm/llms/anthropic/pass_through/messages/handler.py +++ b/litellm/llms/anthropic/pass_through/messages/handler.py @@ -7,7 +7,7 @@ import asyncio import contextvars -from collections.abc import AsyncIterator, Coroutine, Iterator +from collections.abc import AsyncIterator, Coroutine, Iterator, Sequence from functools import partial from typing import Any, Final, cast @@ -23,6 +23,7 @@ from litellm.llms.base_llm.anthropic_messages.transformation import ( from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.types.llms.anthropic import AllAnthropicToolsValues from litellm.types.llms.anthropic_messages.anthropic_request import AnthropicMetadata from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, @@ -467,6 +468,7 @@ def anthropic_messages_handler( custom_llm_provider=custom_llm_provider, api_base=litellm_params.api_base, api_key=litellm_params.api_key, + litellm_params=litellm_params, ) # Store agentic loop params in logging object for agentic hooks @@ -513,7 +515,9 @@ def anthropic_messages_handler( metadata=metadata, stop_sequences=stop_sequences, stream=stream, - system=system, + system=cast( # cast-ok: anthropic accepts list system blocks; handler signature is narrower + "str | None", system + ), temperature=temperature, thinking=thinking, tool_choice=tool_choice, @@ -557,11 +561,15 @@ def anthropic_messages_handler( metadata=metadata, stop_sequences=stop_sequences, stream=stream, - system=system, + system=cast( # cast-ok: anthropic accepts list system blocks; handler signature is narrower + "str | None", system + ), temperature=temperature, thinking=thinking, tool_choice=tool_choice, - tools=tools, + tools=cast( # cast-ok: dict tools already fit the handler's union + "list[AllAnthropicToolsValues | dict[str, object]] | None", tools + ), top_k=top_k, top_p=top_p, _is_async=is_async, @@ -583,11 +591,13 @@ def anthropic_messages_handler( metadata=metadata, stop_sequences=stop_sequences, stream=stream, - system=system, + system=cast( # cast-ok: anthropic accepts list system blocks; handler signature is narrower + "str | None", system + ), temperature=temperature, thinking=thinking, tool_choice=tool_choice, - tools=tools, + tools=cast("list[dict[str, object]] | None", tools), # cast-ok: dict tools fit the adapter's declared union top_k=top_k, top_p=top_p, _is_async=is_async, @@ -619,7 +629,12 @@ def anthropic_messages_handler( ) return base_llm_http_handler.anthropic_messages_handler( model=model, - messages=strip_provider_specific_fields_from_anthropic_messages(messages), + messages=cast( # cast-ok: anthropic messages are dicts at runtime + "list[dict[str, object]]", + strip_provider_specific_fields_from_anthropic_messages( + cast(Sequence[object], messages) # cast-ok: anthropic payloads arrive as untyped dicts + ), + ), anthropic_messages_provider_config=anthropic_messages_provider_config, anthropic_messages_optional_request_params=dict(anthropic_messages_optional_request_params), _is_async=is_async, diff --git a/litellm/llms/github_copilot/authenticator.py b/litellm/llms/github_copilot/authenticator.py index a58f6b08a89..59e1804f453 100644 --- a/litellm/llms/github_copilot/authenticator.py +++ b/litellm/llms/github_copilot/authenticator.py @@ -45,6 +45,22 @@ _DEVICE_CODE: Final = TypeAdapter(_DeviceCode) _ACCESS_TOKEN_POLL: Final = TypeAdapter(_AccessTokenPoll) +def github_api_headers( + access_token: str | None = None, +) -> dict[str, str]: # mutable-ok: returned straight into httpx handlers whose headers params require dict + headers: Final = { + "accept": "application/json", + "editor-version": "vscode/1.85.1", + "editor-plugin-version": "copilot/1.155.0", + "user-agent": "GithubCopilot/1.155.0", + "accept-encoding": "gzip,deflate,br", + "content-type": "application/json", + } + if access_token: + headers["authorization"] = f"token {access_token}" + return headers + + class Authenticator: def __init__(self) -> None: """Initialize the GitHub Copilot authenticator with configurable token paths.""" @@ -230,21 +246,7 @@ class Authenticator: Returns: Dict[str, str]: Headers for GitHub API requests. """ - headers: Final = { - "accept": "application/json", - "editor-version": "vscode/1.85.1", - "editor-plugin-version": "copilot/1.155.0", - "user-agent": "GithubCopilot/1.155.0", - "accept-encoding": "gzip,deflate,br", - } - - if access_token: - headers["authorization"] = f"token {access_token}" - - if "content-type" not in headers: - headers["content-type"] = "application/json" - - return headers + return github_api_headers(access_token) def _get_device_code(self) -> _DeviceCode: """ diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py index b17bd74365c..9553327db97 100644 --- a/litellm/llms/github_copilot/chat/transformation.py +++ b/litellm/llms/github_copilot/chat/transformation.py @@ -1,7 +1,7 @@ import json import os -from collections.abc import Sequence -from typing import TYPE_CHECKING, Final +from collections.abc import Mapping, Sequence +from typing import TYPE_CHECKING, Final, cast # noqa: TID251 # spread dict loses the role literal without a cast import httpx @@ -17,9 +17,11 @@ from ..common_utils import ( GetAPIKeyError, get_copilot_default_headers, ) +from ..per_user_auth import require_github_copilot_user_session if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.tokenizer import Encoding class GithubCopilotConfig(OpenAIConfig): @@ -85,14 +87,22 @@ class GithubCopilotConfig(OpenAIConfig): for message in messages: if message.get("role") == "system": # Convert system message to assistant message - transformed_message = message.copy() - transformed_message["role"] = "assistant" - transformed_messages.append(transformed_message) + transformed_messages.append( + cast(AllMessageValues, {**message, "role": "assistant"}) # cast-ok: dict matches the union + ) else: transformed_messages.append(message) return transformed_messages + def _copilot_headers(self, session_token: str | None) -> Mapping[str, str]: + if session_token is not None: + return get_copilot_default_headers(session_token) + try: + return get_copilot_default_headers(self.authenticator.get_api_key()) + except GetAPIKeyError: + return {} + def validate_environment( self, headers: dict, @@ -103,18 +113,19 @@ class GithubCopilotConfig(OpenAIConfig): api_key: str | None = None, api_base: str | None = None, ) -> dict: - # Get base headers from parent - validated_headers = super().validate_environment( - headers, model, messages, optional_params, litellm_params, api_key, api_base + parent_headers: Final = cast( # cast-ok: the parent validate_environment returns a plain dict at runtime + "dict[str, str]", + super().validate_environment(headers, model, messages, optional_params, litellm_params, api_key, api_base), ) - - # Add Copilot-specific headers (editor-version, user-agent, etc.) - try: - copilot_api_key: Final = self.authenticator.get_api_key() - copilot_headers: Final = get_copilot_default_headers(copilot_api_key) - validated_headers = {**copilot_headers, **validated_headers} - except GetAPIKeyError: - pass # Will be handled later in the request flow + user_session: Final = require_github_copilot_user_session( + cast("dict[str, object]", litellm_params) # cast-ok: litellm_params arrives as an untyped request dict + ) + session_token: Final = user_session.token if user_session is not None else None + validated_headers: Final = {**self._copilot_headers(session_token), **parent_headers} + if session_token is not None: + # the caller's stored session token wins unconditionally, even over a + # caller-supplied api_key or the parent's placeholder Authorization + validated_headers["Authorization"] = f"Bearer {session_token}" # Add X-Initiator header based on message roles initiator: Final = self._determine_initiator(messages) @@ -293,7 +304,7 @@ class GithubCopilotConfig(OpenAIConfig): messages: list[AllMessageValues], optional_params: dict, litellm_params: dict, - encoding: object, + encoding: "Encoding | None", api_key: str | None = None, json_mode: bool | None = None, ) -> "ModelResponse": diff --git a/litellm/llms/github_copilot/common_utils.py b/litellm/llms/github_copilot/common_utils.py index 110a6584ec2..515b7bf4fb5 100644 --- a/litellm/llms/github_copilot/common_utils.py +++ b/litellm/llms/github_copilot/common_utils.py @@ -2,6 +2,7 @@ Constants for Copilot integration """ +from collections.abc import MutableMapping from typing import Final from uuid import uuid4 @@ -57,7 +58,9 @@ class GetAPIKeyError(GithubCopilotError): pass -def get_copilot_default_headers(api_key: str) -> dict: +def get_copilot_default_headers( + api_key: str, +) -> dict[str, str]: # mutable-ok: callers merge caller headers into the returned dict """ Get default headers for GitHub Copilot Responses API. @@ -75,3 +78,15 @@ def get_copilot_default_headers(api_key: str) -> dict: "x-request-id": str(uuid4()), "x-vscode-user-agent-library-version": "electron-fetch", } + + +def pin_session_authorization( + headers: MutableMapping[str, str], # mutable-ok: the helper fixes caller-owned headers in place + session_token: str, +) -> None: + """The stored per-user session token owns Authorization outright: drop every + caller-supplied bearer header (any casing) so it can never displace the + session's on the wire.""" + for key in tuple(k for k in headers if k.lower() == "authorization"): + del headers[key] + headers["Authorization"] = f"Bearer {session_token}" diff --git a/litellm/llms/github_copilot/embedding/transformation.py b/litellm/llms/github_copilot/embedding/transformation.py index c1acf6fccaa..a6250afdd03 100644 --- a/litellm/llms/github_copilot/embedding/transformation.py +++ b/litellm/llms/github_copilot/embedding/transformation.py @@ -8,7 +8,12 @@ https://github.com/caozhiyuan/copilot-api """ import os -from typing import TYPE_CHECKING, Any, Final +from typing import ( + TYPE_CHECKING, + Any, + Final, + cast, # noqa: TID251 # narrows untyped request dicts at the session boundary +) import httpx from pydantic import ConfigDict, TypeAdapter @@ -25,7 +30,9 @@ from ..common_utils import ( DEFAULT_GITHUB_COPILOT_API_BASE, GetAPIKeyError, get_copilot_default_headers, + pin_session_authorization, ) +from ..per_user_auth import require_github_copilot_user_session if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -61,6 +68,15 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): """ Validate environment and set up headers for GitHub Copilot API. """ + user_session: Final = require_github_copilot_user_session( + cast("dict[str, object]", litellm_params) # cast-ok: litellm_params arrives as an untyped request dict + ) + if user_session is not None: + session_headers: Final = cast( # cast-ok: both spreads are str-valued header dicts + "dict[str, str]", {**get_copilot_default_headers(user_session.token), **headers} + ) + pin_session_authorization(session_headers, user_session.token) + return session_headers try: # Get GitHub Copilot API key via OAuth api_key = self.authenticator.get_api_key() @@ -101,9 +117,12 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): """ Get the complete URL for GitHub Copilot Embedding API endpoint. """ - # Use provided api_base or fall back to authenticator's base or default + user_session: Final = require_github_copilot_user_session( + cast("dict[str, object]", litellm_params) # cast-ok: litellm_params arrives as an untyped request dict + ) effective_api_base = ( - api_base + (user_session.api_base if user_session is not None else None) + or api_base or self.authenticator.get_api_base() or os.getenv("GITHUB_COPILOT_API_BASE") or DEFAULT_GITHUB_COPILOT_API_BASE diff --git a/litellm/llms/github_copilot/messages/transformation.py b/litellm/llms/github_copilot/messages/transformation.py index b36b437e2e5..be138755f85 100644 --- a/litellm/llms/github_copilot/messages/transformation.py +++ b/litellm/llms/github_copilot/messages/transformation.py @@ -1,4 +1,4 @@ -from typing import Any, Final +from typing import Any, Final, cast # noqa: TID251 # narrows untyped request dicts at the session boundary from litellm.exceptions import AuthenticationError from litellm.llms.anthropic.pass_through.messages.transformation import ( @@ -10,7 +10,9 @@ from ..common_utils import ( DEFAULT_GITHUB_COPILOT_API_BASE, GetAPIKeyError, get_copilot_default_headers, + pin_session_authorization, ) +from ..per_user_auth import require_github_copilot_user_session _MESSAGES_PROXY_API_VERSION: Final = "2026-06-01" @@ -68,9 +70,18 @@ class GithubCopilotAnthropicMessagesConfig(AnthropicMessagesConfig): # session, never the caller-supplied api_base. rstrip so a # tenant-specific base with a trailing slash does not yield a # double-slash URL once "/v1/messages" is appended downstream. - dynamic_api_base: Final = (self.authenticator.get_api_base() or DEFAULT_GITHUB_COPILOT_API_BASE).rstrip("/") + user_session: Final = require_github_copilot_user_session( + cast("dict[str, object]", litellm_params) # cast-ok: litellm_params arrives as an untyped request dict + ) + dynamic_api_base: Final = ( + user_session.api_base + if user_session is not None + else (self.authenticator.get_api_base() or DEFAULT_GITHUB_COPILOT_API_BASE) + ).rstrip("/") try: - dynamic_api_key: Final = self.authenticator.get_api_key() + dynamic_api_key: Final = ( + user_session.token if user_session is not None else self.authenticator.get_api_key() + ) except GetAPIKeyError as e: raise AuthenticationError( model=model, @@ -83,6 +94,11 @@ class GithubCopilotAnthropicMessagesConfig(AnthropicMessagesConfig): for key, value in copilot_headers.items(): if key not in headers: headers[key] = value + if user_session is not None: + pin_session_authorization( + cast("dict[str, str]", headers), # cast-ok: headers is a str-valued request dict at runtime + user_session.token, + ) headers["openai-intent"] = "messages-proxy" headers["x-interaction-type"] = "messages-proxy" @@ -116,7 +132,13 @@ class GithubCopilotAnthropicMessagesConfig(AnthropicMessagesConfig): reuse it to avoid a second authenticator read, falling back to a fresh resolution only if it was not provided. """ - resolved = (api_base or self.authenticator.get_api_base() or DEFAULT_GITHUB_COPILOT_API_BASE).rstrip("/") - if not resolved.endswith("/v1/messages"): - resolved = f"{resolved}/v1/messages" - return resolved + user_session: Final = require_github_copilot_user_session( + cast("dict[str, object]", litellm_params) # cast-ok: litellm_params arrives as an untyped request dict + ) + resolved: Final = ( + (user_session.api_base if user_session is not None else None) + or api_base + or self.authenticator.get_api_base() + or DEFAULT_GITHUB_COPILOT_API_BASE + ).rstrip("/") + return resolved if resolved.endswith("/v1/messages") else f"{resolved}/v1/messages" diff --git a/litellm/llms/github_copilot/per_user_auth.py b/litellm/llms/github_copilot/per_user_auth.py new file mode 100644 index 00000000000..10e5e12dec3 --- /dev/null +++ b/litellm/llms/github_copilot/per_user_auth.py @@ -0,0 +1,741 @@ +"""Per-user GitHub Copilot OAuth: exchange a caller's stored GitHub token for a +short-lived Copilot token, plus the device-flow helpers the proxy endpoints use +to establish that connection. The shared device login (Authenticator) is never +touched on the per-user path. +""" + +import asyncio +import hashlib +import os +import threading +import time +from collections.abc import Callable, Mapping +from dataclasses import dataclass, field +from typing import ( + Final, + Literal, + Protocol, + TypeAlias, + cast, # noqa: TID251 # casts pin the untyped litellm clients/caches to the local Protocols +) +from urllib.parse import urlparse + +import httpx +from pydantic import ConfigDict, TypeAdapter, ValidationError, with_config +from typing_extensions import NotRequired, ReadOnly, TypedDict + +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.constants import ( + GITHUB_COPILOT_AUTH_TYPE_KEY, + GITHUB_COPILOT_PER_USER_AUTH_TYPE, + GITHUB_COPILOT_USER_TOKEN_SAFETY_MARGIN_SECONDS, +) +from litellm.exceptions import ( + APIConnectionError, + BadRequestError, + CallerCredentialAuthenticationError, + CallerCredentialRateLimitError, + ServiceUnavailableError, +) + +from .authenticator import ( + DEFAULT_GITHUB_ACCESS_TOKEN_URL, + DEFAULT_GITHUB_API_KEY_URL, + DEFAULT_GITHUB_CLIENT_ID, + DEFAULT_GITHUB_DEVICE_CODE_URL, + github_api_headers, +) +from .common_utils import DEFAULT_GITHUB_COPILOT_API_BASE + +GITHUB_COPILOT_USER_SESSION_KWARG_KEY: Final = "github_copilot_user_session" +_LLM_PROVIDER: Final = "github_copilot" +_DEVICE_FLOW_SCOPE: Final = "read:user" +_SESSION_CACHE_MAX_SIZE: Final = 1000 +_SYNC_LOCK_COUNT: Final = 64 + + +class _SessionCache(Protocol): + def get_cache(self, key: str) -> object: ... + + def set_cache(self, key: str, value: object, ttl: float | None = None) -> object: ... + + def delete_cache(self, key: str) -> object: ... + + +class _SyncGetClient(Protocol): + def get( + self, + url: str, + params: dict[str, str] | None = None, # mutable-ok: mirrors the untyped client's dict parameters + headers: dict[str, str] | None = None, # mutable-ok: mirrors the untyped client's dict parameters + ) -> httpx.Response: ... + + +class _AsyncGitHubClient(Protocol): + async def get( + self, + url: str, + params: dict[str, str] | None = None, # mutable-ok: mirrors the untyped client's dict parameters + headers: dict[str, str] | None = None, # mutable-ok: mirrors the untyped client's dict parameters + ) -> httpx.Response: ... + + async def post( + self, + url: str, + data: dict[str, str] | str | bytes | None = None, # mutable-ok: mirrors the untyped client's dict parameters + json: dict[str, str] | None = None, # mutable-ok: mirrors the untyped client's dict parameters + params: dict[str, str] | None = None, # mutable-ok: mirrors the untyped client's dict parameters + headers: dict[str, str] | None = None, # mutable-ok: mirrors the untyped client's dict parameters + ) -> httpx.Response: ... + + +_SESSION_CACHE: Final = cast( # cast-ok: pins the untyped module cache to _SessionCache + _SessionCache, + InMemoryCache( + max_size_in_memory=_SESSION_CACHE_MAX_SIZE, + default_ttl=GITHUB_COPILOT_USER_TOKEN_SAFETY_MARGIN_SECONDS, + ), +) +_EXCHANGE_LOCKS: Final = tuple(threading.Lock() for _ in range(_SYNC_LOCK_COUNT)) +_IN_FLIGHT: Final[ # mutable-ok: single-flight dedup registry mutated under its own locking + dict[str, "asyncio.Future[GithubCopilotUserSession]"] +] = {} + + +@with_config(ConfigDict(extra="allow", strict=True)) +class _CopilotTokenPayload(TypedDict): + token: ReadOnly[str] + expires_at: ReadOnly[NotRequired[object]] + endpoints: ReadOnly[NotRequired[object]] + + +@with_config(ConfigDict(extra="allow", strict=True)) +class _DeviceCodePayload(TypedDict): + device_code: ReadOnly[str] + user_code: ReadOnly[str] + verification_uri: ReadOnly[str] + expires_in: ReadOnly[object] + interval: ReadOnly[object] + + +@with_config(ConfigDict(extra="allow", strict=True)) +class _AccessTokenPollPayload(TypedDict): + access_token: ReadOnly[NotRequired[str]] + error: ReadOnly[NotRequired[object]] + error_description: ReadOnly[NotRequired[object]] + interval: ReadOnly[NotRequired[object]] + + +@with_config(ConfigDict(extra="allow", strict=True)) +class _CopilotEndpointsPayload(TypedDict): + api: ReadOnly[NotRequired[object]] + + +@with_config(ConfigDict(extra="allow", strict=True)) +class _GithubUserPayload(TypedDict): + login: ReadOnly[str] + + +_COPILOT_TOKEN: Final = TypeAdapter(_CopilotTokenPayload) +_DEVICE_CODE: Final = TypeAdapter(_DeviceCodePayload) +_ACCESS_TOKEN_POLL: Final = TypeAdapter(_AccessTokenPollPayload) +_GITHUB_USER: Final = TypeAdapter(_GithubUserPayload) +_COPILOT_ENDPOINTS: Final = TypeAdapter(_CopilotEndpointsPayload) +_CACHED_SESSION_ENTRY: Final = TypeAdapter(tuple[str, str, float]) + + +@dataclass(frozen=True, slots=True, repr=False) +class GithubCopilotUserSession: + """A resolved per-user Copilot token plus the validated host it may be sent + to. Only this module constructs it, so ``isinstance`` checks downstream are + unforgeable by request JSON.""" + + token: str = field(repr=False) + api_base: str + + def __repr__(self) -> str: + return f"GithubCopilotUserSession(api_base={self.api_base!r}, token=REDACTED)" + + __str__ = __repr__ + + +@dataclass(frozen=True, slots=True) +class GithubCopilotDeviceFlowStart: + device_code: str = field(repr=False) + user_code: str + verification_uri: str + expires_in: int + interval: int + + +DeviceFlowPollStatus: TypeAlias = Literal["pending", "slow_down", "expired", "denied", "connected"] + + +@dataclass(frozen=True, slots=True) +class GithubCopilotDeviceFlowPoll: + status: DeviceFlowPollStatus + interval: int | None = None + access_token: str | None = field(default=None, repr=False) + + +def validated_copilot_api_base(endpoints_api: object) -> str: + """Only a genuine *.githubcopilot.com HTTPS endpoint may receive a Copilot + token; anything else falls back to the default base.""" + if not isinstance(endpoints_api, str): + return DEFAULT_GITHUB_COPILOT_API_BASE + parsed: Final = urlparse(endpoints_api) + hostname: Final = (parsed.hostname or "").lower() + if ( + parsed.scheme.lower() != "https" + or not (hostname == "githubcopilot.com" or hostname.endswith(".githubcopilot.com")) + or parsed.username is not None + or parsed.password is not None + ): + return DEFAULT_GITHUB_COPILOT_API_BASE + try: + port: Final = parsed.port + except ValueError: + return DEFAULT_GITHUB_COPILOT_API_BASE + if port not in (None, 443): + return DEFAULT_GITHUB_COPILOT_API_BASE + return endpoints_api.rstrip("/") + + +def _session_cache_key(user_id: str, github_token: str) -> str: + return hashlib.sha256(f"{user_id}\x00{github_token}".encode()).hexdigest() + + +def evict_copilot_user_session(user_id: str, github_token: str) -> None: + _SESSION_CACHE.delete_cache(_session_cache_key(user_id, github_token)) + + +def _cached_session(cache_key: str) -> GithubCopilotUserSession | None: + cached: Final[object] = _SESSION_CACHE.get_cache(cache_key) + try: + token, api_base, expires_at = _CACHED_SESSION_ENTRY.validate_python(cached) + except ValidationError: + return None + if expires_at - GITHUB_COPILOT_USER_TOKEN_SAFETY_MARGIN_SECONDS <= time.time(): + return None + return GithubCopilotUserSession(token=token, api_base=api_base) + + +def _cache_session(cache_key: str, session: GithubCopilotUserSession, expires_at: float) -> None: + ttl: Final = expires_at - GITHUB_COPILOT_USER_TOKEN_SAFETY_MARGIN_SECONDS - time.time() + if ttl <= 0: + return + _SESSION_CACHE.set_cache(cache_key, (session.token, session.api_base, expires_at), ttl=ttl) + + +def _reject_connection(credential_name: str) -> CallerCredentialAuthenticationError: + return CallerCredentialAuthenticationError( + message=( + f"GitHub rejected the GitHub Copilot connection for credential '{credential_name}'. " + f"Reconnect GitHub Copilot for credential '{credential_name}' in the LiteLLM UI (LLM Credentials)" + ), + llm_provider=_LLM_PROVIDER, + model="", + ) + + +def _session_from_response(response: httpx.Response, credential_name: str, cache_key: str) -> GithubCopilotUserSession: + if response.status_code in (401, 403, 404): + _SESSION_CACHE.delete_cache(cache_key) + raise _reject_connection(credential_name) + if response.status_code == 429: + raise CallerCredentialRateLimitError( + message="GitHub rate limited the GitHub Copilot token exchange", + llm_provider=_LLM_PROVIDER, + model="", + ) + if not response.is_success: + raise ServiceUnavailableError( + message=f"GitHub Copilot token exchange failed with status {response.status_code}", + llm_provider=_LLM_PROVIDER, + model="", + ) + try: + payload: Final = _COPILOT_TOKEN.validate_python(response.json()) + except Exception as e: + raise ServiceUnavailableError( + message="GitHub Copilot token exchange returned an invalid response", + llm_provider=_LLM_PROVIDER, + model="", + ) from e + + endpoints: Final[object] = payload.get("endpoints") + endpoints_api: Final[object] = ( + _COPILOT_ENDPOINTS.validate_python(endpoints).get("api") if isinstance(endpoints, Mapping) else None + ) + api_base: Final = validated_copilot_api_base(endpoints_api) + session: Final = GithubCopilotUserSession(token=payload["token"], api_base=api_base) + expires_at: Final = _expires_at_seconds(payload.get("expires_at")) + if expires_at > 0: + _cache_session(cache_key, session, expires_at) + return session + + +def _to_int(raw: object) -> int: + if isinstance(raw, bool): + return 0 + if isinstance(raw, int): + return raw + if isinstance(raw, float): + return int(raw) + if isinstance(raw, str): + try: + return int(raw) + except ValueError: + return 0 + return 0 + + +def _expires_at_seconds(raw: object) -> float: + if isinstance(raw, bool): + return 0.0 + if isinstance(raw, (int, float)): + return float(raw) + if isinstance(raw, str): + try: + return float(raw) + except ValueError: + return 0.0 + return 0.0 + + +def _do_exchange(github_token: str, credential_name: str, cache_key: str) -> GithubCopilotUserSession: + import litellm + + try: + client: Final = cast( # cast-ok: pins the untyped litellm client to the local Protocol + _SyncGetClient, litellm.module_level_client + ) + response: Final = client.get( + DEFAULT_GITHUB_API_KEY_URL, + headers=github_api_headers(github_token), + ) + except Exception as e: + raise APIConnectionError( + message="GitHub Copilot token exchange request failed", + llm_provider=_LLM_PROVIDER, + model="", + ) from e + return _session_from_response(response, credential_name, cache_key) + + +def _new_async_github_client() -> _AsyncGitHubClient: + from litellm.llms.custom_httpx import http_handler + + return cast( # cast-ok: pins the untyped litellm client factory to the local Protocol + Callable[[str], _AsyncGitHubClient], + getattr(http_handler, "get_async_httpx_client"), # noqa: B009 # untyped factory + )(_LLM_PROVIDER) + + +async def _do_aexchange(github_token: str, credential_name: str, cache_key: str) -> GithubCopilotUserSession: + client: Final = _new_async_github_client() + try: + response: Final = await client.get( + DEFAULT_GITHUB_API_KEY_URL, + headers=github_api_headers(github_token), + ) + except Exception as e: + raise APIConnectionError( + message="GitHub Copilot token exchange request failed", + llm_provider=_LLM_PROVIDER, + model="", + ) from e + return _session_from_response(response, credential_name, cache_key) + + +def exchange_github_token( + user_id: str, + github_token: str, + credential_name: str, +) -> GithubCopilotUserSession: + cache_key: Final = _session_cache_key(user_id, github_token) + cached: Final = _cached_session(cache_key) + if cached is not None: + return cached + lock: Final = _EXCHANGE_LOCKS[int(cache_key, 16) % _SYNC_LOCK_COUNT] + with lock: + rechecked: Final = _cached_session(cache_key) + if rechecked is not None: + return rechecked + return _do_exchange(github_token, credential_name, cache_key) + + +async def aexchange_github_token( + user_id: str, + github_token: str, + credential_name: str, +) -> GithubCopilotUserSession: + cache_key: Final = _session_cache_key(user_id, github_token) + cached: Final = _cached_session(cache_key) + if cached is not None: + return cached + loop: Final = asyncio.get_running_loop() + in_flight: Final = _IN_FLIGHT.get(cache_key) + if in_flight is not None and not in_flight.done() and in_flight.get_loop() is loop: + return await asyncio.shield(in_flight) + future: Final = loop.create_future() + _IN_FLIGHT[cache_key] = future + try: + session: Final = await _do_aexchange(github_token, credential_name, cache_key) + except BaseException as e: + if not future.done(): + future.set_exception(e) + future.exception() # consume so unawaited waiters do not log "exception was never retrieved" + raise + else: + if not future.done(): + future.set_result(session) + return session + finally: + if _IN_FLIGHT.get(cache_key) is future: + _IN_FLIGHT.pop(cache_key, None) + + +async def astart_device_flow() -> GithubCopilotDeviceFlowStart: + client: Final = _new_async_github_client() + device_code_url: Final = os.getenv("GITHUB_COPILOT_DEVICE_CODE_URL", DEFAULT_GITHUB_DEVICE_CODE_URL) + client_id: Final = os.getenv("GITHUB_COPILOT_CLIENT_ID", DEFAULT_GITHUB_CLIENT_ID) + try: + response: Final = await client.post( + device_code_url, + headers=github_api_headers(), + json={"client_id": client_id, "scope": _DEVICE_FLOW_SCOPE}, + ) + except httpx.HTTPStatusError as e: + raise ServiceUnavailableError( + message=f"GitHub device flow start failed with status {e.response.status_code}", + llm_provider=_LLM_PROVIDER, + model="", + ) from e + except Exception as e: + raise APIConnectionError( + message="GitHub device flow start request failed", + llm_provider=_LLM_PROVIDER, + model="", + ) from e + try: + payload: Final = _DEVICE_CODE.validate_python(response.json()) + except Exception as e: + raise ServiceUnavailableError( + message="GitHub device flow start returned an invalid response", + llm_provider=_LLM_PROVIDER, + model="", + ) from e + return GithubCopilotDeviceFlowStart( + device_code=payload["device_code"], + user_code=payload["user_code"], + verification_uri=payload["verification_uri"], + expires_in=_to_int(payload["expires_in"]), + interval=_to_int(payload["interval"]), + ) + + +async def apoll_device_flow(device_code: str) -> GithubCopilotDeviceFlowPoll: + client: Final = _new_async_github_client() + access_token_url: Final = os.getenv("GITHUB_COPILOT_ACCESS_TOKEN_URL", DEFAULT_GITHUB_ACCESS_TOKEN_URL) + client_id: Final = os.getenv("GITHUB_COPILOT_CLIENT_ID", DEFAULT_GITHUB_CLIENT_ID) + try: + response: Final = await client.post( + access_token_url, + headers=github_api_headers(), + json={ + "client_id": client_id, + "device_code": device_code, + "grant_type": "urn:ietf:params:oauth:grant-type:device_code", + }, + ) + except httpx.HTTPStatusError as e: + return _poll_payload_from_response(e.response) + except Exception as e: + raise APIConnectionError( + message="GitHub device flow poll request failed", + llm_provider=_LLM_PROVIDER, + model="", + ) from e + return _poll_payload_from_response(response) + + +def _poll_payload_from_response(response: httpx.Response) -> GithubCopilotDeviceFlowPoll: + if not response.is_success: + raise ServiceUnavailableError( + message=f"GitHub device flow poll failed with status {response.status_code}", + llm_provider=_LLM_PROVIDER, + model="", + ) + try: + payload: Final = _ACCESS_TOKEN_POLL.validate_python(response.json()) + except Exception as e: + raise ServiceUnavailableError( + message="GitHub device flow poll returned an invalid response", + llm_provider=_LLM_PROVIDER, + model="", + ) from e + + access_token: Final = payload.get("access_token") + if isinstance(access_token, str) and access_token: + return GithubCopilotDeviceFlowPoll(status="connected", access_token=access_token) + + error: Final = payload.get("error") + raw_interval: Final = payload.get("interval") + interval: Final = _to_int(raw_interval) or None + match error: + case "authorization_pending": + return GithubCopilotDeviceFlowPoll(status="pending", interval=interval) + case "slow_down": + return GithubCopilotDeviceFlowPoll(status="slow_down", interval=interval) + case "expired_token": + return GithubCopilotDeviceFlowPoll(status="expired", interval=interval) + case "access_denied": + return GithubCopilotDeviceFlowPoll(status="denied", interval=interval) + case _: + raise ServiceUnavailableError( + message="GitHub device flow poll returned an unrecognized error", + llm_provider=_LLM_PROVIDER, + model="", + ) + + +async def afetch_github_login(github_token: str) -> str: + client: Final = _new_async_github_client() + try: + response: Final = await client.get( + "https://api.github.com/user", + headers=github_api_headers(github_token), + ) + except Exception as e: + raise APIConnectionError( + message="GitHub user lookup request failed", + llm_provider=_LLM_PROVIDER, + model="", + ) from e + if not response.is_success: + raise ServiceUnavailableError( + message=f"GitHub user lookup failed with status {response.status_code}", + llm_provider=_LLM_PROVIDER, + model="", + ) + try: + payload: Final = _GITHUB_USER.validate_python(response.json()) + except Exception as e: + raise ServiceUnavailableError( + message="GitHub user lookup returned an invalid response", + llm_provider=_LLM_PROVIDER, + model="", + ) from e + return payload["login"] + + +async def acheck_copilot_seat(user_id: str, github_token: str, credential_name: str) -> GithubCopilotUserSession: + return await aexchange_github_token( + user_id=user_id, + github_token=github_token, + credential_name=credential_name, + ) + + +def github_copilot_auth_mode(litellm_credential_name: object, auth_type_value: object) -> bool: + """Whether this call runs in per-user mode. A resolved credential's stored + ``credential_values`` are the only authority once a credential is named and + found; the request's own ``github_copilot_auth_type`` can never flip the + mode either way on a named credential.""" + if isinstance(litellm_credential_name, str) and litellm_credential_name: + from litellm.litellm_core_utils.credential_accessor import CredentialAccessor + + credential: Final = CredentialAccessor.find_credential(litellm_credential_name) + if credential is not None: + values: Final[object] = getattr(credential, "credential_values", None) + return isinstance(values, Mapping) and ( + cast( # cast-ok: value is Mapping-checked or dict-shaped at runtime + Mapping[object, object], values + ).get(GITHUB_COPILOT_AUTH_TYPE_KEY) + == GITHUB_COPILOT_PER_USER_AUTH_TYPE + ) + return auth_type_value == GITHUB_COPILOT_PER_USER_AUTH_TYPE + if auth_type_value == GITHUB_COPILOT_PER_USER_AUTH_TYPE: + raise BadRequestError( + message=f"{GITHUB_COPILOT_AUTH_TYPE_KEY} {GITHUB_COPILOT_PER_USER_AUTH_TYPE} requires litellm_credential_name", + llm_provider=_LLM_PROVIDER, + model="", + ) + return False + + +def _connect_error(credential_name: str) -> CallerCredentialAuthenticationError: + return CallerCredentialAuthenticationError( + message=(f"Connect GitHub Copilot for credential '{credential_name}' in the LiteLLM UI (LLM Credentials)"), + llm_provider=_LLM_PROVIDER, + model="", + ) + + +def _caller_connection(kwargs: Mapping[str, object], credential_name: str) -> tuple[str, str]: + secret_fields_raw: Final[object] = kwargs.get("secret_fields") + secret_fields: Final = ( + cast(Mapping[object, object], secret_fields_raw) # cast-ok: value is Mapping-checked or dict-shaped at runtime + if isinstance(secret_fields_raw, Mapping) + else None + ) + user_id: Final[object] = ( + secret_fields.get("user_provider_credentials_user_id") if secret_fields is not None else None + ) + credentials_raw: Final[object] = ( + secret_fields.get("user_provider_credentials") if secret_fields is not None else None + ) + credentials: Final = ( + cast(Mapping[object, object], credentials_raw) # cast-ok: value is Mapping-checked or dict-shaped at runtime + if isinstance(credentials_raw, Mapping) + else None + ) + token: Final[object] = credentials.get(credential_name) if credentials is not None else None + if not isinstance(user_id, str) or not user_id or not isinstance(token, str) or not token: + raise _connect_error(credential_name) + return user_id, token + + +def _params_value(source: object, key: str) -> object: + if isinstance(source, Mapping): + return cast(Mapping[object, object], source).get(key) # cast-ok: isinstance above, Mapping values are object + return getattr(source, key, None) + + +def github_copilot_per_user_credential_name(litellm_params: object) -> str | None: + """The credential name iff these litellm_params select per-user mode.""" + credential_name: Final[object] = _params_value(litellm_params, "litellm_credential_name") + auth_type: Final[object] = _params_value(litellm_params, GITHUB_COPILOT_AUTH_TYPE_KEY) + if not github_copilot_auth_mode(credential_name, auth_type): + return None + return credential_name if isinstance(credential_name, str) else "" + + +def github_copilot_user_session_from(source: object) -> GithubCopilotUserSession | None: + if source is None: + return None + candidate: Final = ( + cast(Mapping[object, object], source).get( # cast-ok: value is Mapping-checked or dict-shaped at runtime + GITHUB_COPILOT_USER_SESSION_KWARG_KEY + ) + if isinstance(source, Mapping) + else getattr(source, GITHUB_COPILOT_USER_SESSION_KWARG_KEY, None) + ) + return candidate if isinstance(candidate, GithubCopilotUserSession) else None + + +def require_github_copilot_user_session(litellm_params: object) -> GithubCopilotUserSession | None: + """The attached session, or None for shared mode. Per-user params with no + attached session only reach here via paths that skipped the attach step, so + they get the same not-connected 401 instead of the shared device login.""" + session: Final = github_copilot_user_session_from(litellm_params) + if session is not None: + return session + credential_name: Final = github_copilot_per_user_credential_name(litellm_params) + if credential_name is not None: + raise _connect_error(credential_name) + return None + + +def is_github_copilot_per_user_request(kwargs: Mapping[str, object]) -> bool: + """True when this call carries a resolved per-user Copilot session, meaning it + must bypass the shared response cache (cache keys ignore caller identity). + The session kwarg is attached in ``utils.wrapper``/``wrapper_async`` after + ``load_credentials_from_list``, before the cache lookup runs.""" + return isinstance(kwargs.get(GITHUB_COPILOT_USER_SESSION_KWARG_KEY), GithubCopilotUserSession) + + +def _without_authorization(headers: Mapping[str, object]) -> dict[str, object]: + return {k: v for k, v in headers.items() if k.lower() != "authorization"} + + +def _str_keyed_mapping(value: object) -> Mapping[str, object] | None: + return ( + cast("Mapping[str, object]", value) # cast-ok: request mappings are str-keyed + if isinstance(value, Mapping) + else None + ) + + +def _strip_caller_authorization( + kwargs: dict[str, object], # mutable-ok: kwargs is the request mutation channel +) -> None: + """Once a per-user session is attached, its token owns Authorization on the + wire. Caller-supplied bearer headers would be re-merged by the HTTP handler + after the transformations run, so they are rebuilt here without any + Authorization key; the caller's mappings are left untouched.""" + for key in ("extra_headers", "headers"): + value = kwargs.get(key) + if isinstance(value, Mapping): + stripped = _without_authorization( + cast("Mapping[str, object]", value) # cast-ok: Mapping-checked request header dict + ) + if len(stripped) != len( + cast("Mapping[object, object]", value) # cast-ok: Mapping-checked request header dict + ): + kwargs[key] = stripped # rebind-ok: kwargs is the request mutation channel + optional_params: Final = kwargs.get("optional_params") + if isinstance(optional_params, dict): + eh: Final = _str_keyed_mapping( + cast("dict[str, object]", optional_params).get( # cast-ok: optional_params is a plain str-keyed dict + "extra_headers" + ) + ) + if eh is not None: + kwargs["optional_params"] = { # rebind-ok: kwargs is the request mutation channel + **optional_params, + "extra_headers": _without_authorization(eh), + } + litellm_params: Final = kwargs.get("litellm_params") + if isinstance(litellm_params, dict): + eh2: Final = _str_keyed_mapping( + cast("dict[str, object]", litellm_params).get( # cast-ok: litellm_params is a plain str-keyed dict + "extra_headers" + ) + ) + if eh2 is not None: + kwargs["litellm_params"] = { # rebind-ok: kwargs is the request mutation channel + **litellm_params, + "extra_headers": _without_authorization(eh2), + } + + +def attach_github_copilot_user_session( + kwargs: dict[str, object], # mutable-ok: kwargs is the request mutation channel +) -> None: + credential_name_obj: Final[object] = kwargs.get("litellm_credential_name") + auth_type_obj: Final[object] = kwargs.get(GITHUB_COPILOT_AUTH_TYPE_KEY) + if not github_copilot_auth_mode(credential_name_obj, auth_type_obj): + return + if isinstance(kwargs.get(GITHUB_COPILOT_USER_SESSION_KWARG_KEY), GithubCopilotUserSession): + return + credential_name: Final = credential_name_obj if isinstance(credential_name_obj, str) else "" + user_id, github_token = _caller_connection(kwargs, credential_name) + session: Final = exchange_github_token( + user_id=user_id, + github_token=github_token, + credential_name=credential_name, + ) + kwargs[GITHUB_COPILOT_USER_SESSION_KWARG_KEY] = session # rebind-ok: kwargs is the request mutation channel + _strip_caller_authorization(kwargs) + + +async def aattach_github_copilot_user_session( + kwargs: dict[str, object], # mutable-ok: kwargs is the request mutation channel +) -> None: + credential_name_obj: Final[object] = kwargs.get("litellm_credential_name") + auth_type_obj: Final[object] = kwargs.get(GITHUB_COPILOT_AUTH_TYPE_KEY) + if not github_copilot_auth_mode(credential_name_obj, auth_type_obj): + return + if isinstance(kwargs.get(GITHUB_COPILOT_USER_SESSION_KWARG_KEY), GithubCopilotUserSession): + return + credential_name: Final = credential_name_obj if isinstance(credential_name_obj, str) else "" + user_id, github_token = _caller_connection(kwargs, credential_name) + session: Final = await aexchange_github_token( + user_id=user_id, + github_token=github_token, + credential_name=credential_name, + ) + kwargs[GITHUB_COPILOT_USER_SESSION_KWARG_KEY] = session # rebind-ok: kwargs is the request mutation channel + _strip_caller_authorization(kwargs) diff --git a/litellm/llms/github_copilot/responses/transformation.py b/litellm/llms/github_copilot/responses/transformation.py index 7bdf253f28f..d468889b626 100644 --- a/litellm/llms/github_copilot/responses/transformation.py +++ b/litellm/llms/github_copilot/responses/transformation.py @@ -9,7 +9,12 @@ https://github.com/caozhiyuan/copilot-api """ import os -from typing import TYPE_CHECKING, Any, Final +from typing import ( + TYPE_CHECKING, + Any, + Final, + cast, # noqa: TID251 # narrows untyped request dicts at the session boundary +) import litellm from litellm._logging import verbose_logger @@ -30,7 +35,9 @@ from ..common_utils import ( DEFAULT_GITHUB_COPILOT_API_BASE, GetAPIKeyError, get_copilot_default_headers, + pin_session_authorization, ) +from ..per_user_auth import require_github_copilot_user_session if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -199,11 +206,14 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): - copilot-vision-request if vision content detected - User-provided extra_headers (merged with priority) """ + user_session: Final = require_github_copilot_user_session( + cast("dict[str, object]", litellm_params) # cast-ok: litellm_params arrives as an untyped request dict + ) + default_headers: Final = get_copilot_default_headers(user_session.token) if user_session is not None else None try: # Get GitHub Copilot API key via OAuth - api_key: Final = self.authenticator.get_api_key() - - if not api_key: + api_key: Final = self.authenticator.get_api_key() if default_headers is None else "" + if default_headers is None and not api_key: raise AuthenticationError( model=model, llm_provider="github_copilot", @@ -211,10 +221,14 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): ) # Get default headers (from copilot-api configuration) - default_headers: Final = get_copilot_default_headers(api_key) + copilot_headers: Final = default_headers or get_copilot_default_headers(api_key) # Merge with existing headers (user's extra_headers take priority) - merged_headers: Final = {**default_headers, **headers} + merged_headers: Final = cast( # cast-ok: both spreads are str-valued header dicts + "dict[str, str]", {**copilot_headers, **headers} + ) + if user_session is not None: + pin_session_authorization(merged_headers, user_session.token) # Analyze input to determine additional headers input_param: Final = self._get_input_from_params(litellm_params) @@ -249,9 +263,12 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): """ Get the complete URL for GitHub Copilot Responses API endpoint. """ - # Use provided api_base or fall back to authenticator's base or default + user_session: Final = require_github_copilot_user_session( + cast("dict[str, object]", litellm_params) # cast-ok: litellm_params arrives as an untyped request dict + ) effective_api_base = ( - api_base + (user_session.api_base if user_session is not None else None) + or api_base or self.authenticator.get_api_base() or os.getenv("GITHUB_COPILOT_API_BASE") or DEFAULT_GITHUB_COPILOT_API_BASE @@ -341,9 +358,9 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): return "user" # If input is a list, analyze items - if isinstance(input_param, list): + if isinstance(input_param, list): # pyright: ignore[reportUnnecessaryIsInstance] # items arrive from untyped request params for item in input_param: - if not isinstance(item, dict): + if not isinstance(item, dict): # pyright: ignore[reportUnnecessaryIsInstance] # items arrive from untyped request params continue # Check if item has no role (agent-initiated) @@ -352,7 +369,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): # Check if role is assistant (agent-initiated) role = item.get("role") - if isinstance(role, str) and role.lower() == "assistant": + if isinstance(role, str) and role.lower() == "assistant": # pyright: ignore[reportUnnecessaryIsInstance] # role arrives from untyped request params return "agent" # Default to user-initiated diff --git a/litellm/main.py b/litellm/main.py index 4c72cf26359..b80819354f0 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -642,10 +642,17 @@ async def acompletion( "enable_json_schema_validation": enable_json_schema_validation, } if custom_llm_provider is None: + _supplemental: Final = cast( # cast-ok: kwargs values are request-scoped objects + "dict[str, object]", {k: kwargs[k] for k in OPTIONAL_KWARGS_KEYS if k in kwargs} + ) _, custom_llm_provider, _, _ = get_llm_provider( model=model, custom_llm_provider=custom_llm_provider, - api_base=kwargs.get("api_base") or base_url, + api_base=cast( # cast-ok: api_base arrives as a str in the request kwargs + "str | None", kwargs.get("api_base") + ) + or base_url, + litellm_params=(GenericLiteLLMParams.model_validate(_supplemental) if _supplemental else None), ) fallbacks = fallbacks or litellm.model_fallbacks @@ -2649,13 +2656,22 @@ def _complete_custom_openai( from litellm.llms.github_copilot.authenticator import Authenticator from litellm.llms.github_copilot.common_utils import ( get_copilot_default_headers, + pin_session_authorization, + ) + from litellm.llms.github_copilot.per_user_auth import ( + require_github_copilot_user_session, ) - copilot_auth: Final = Authenticator() - copilot_api_key: Final = copilot_auth.get_api_key() - copilot_headers: Final = get_copilot_default_headers(copilot_api_key) + user_session: Final = require_github_copilot_user_session(litellm_params) + copilot_headers: Final = get_copilot_default_headers( + user_session.token if user_session is not None else Authenticator().get_api_key() + ) if extra_headers: - copilot_headers.update(extra_headers) + copilot_headers.update( + cast("dict[str, str]", extra_headers) # cast-ok: extra_headers is a str-valued request dict + ) + if user_session is not None: + pin_session_authorization(copilot_headers, user_session.token) extra_headers = copilot_headers use_base_llm_http_handler: Final = get_secret_bool("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER") @@ -5790,6 +5806,8 @@ def completion( *ANTHROPIC_WIF_KWARGS_KEYS, *OPENAI_WIF_KWARGS_KEYS, *OAUTH_TOKEN_EXCHANGE_KWARGS_KEYS, + "github_copilot_auth_type", + "github_copilot_user_session", PROVIDER_AFFINITY_HEADER_KWARG_KEY, "fireworks_forward_user_id", ) @@ -6483,17 +6501,24 @@ def embedding( default_params: Final = [*openai_params, "aembedding", "extra_headers"] non_default_params: Final = filter_out_litellm_params(kwargs, excluding=default_params) + litellm_params_dict: Final = get_litellm_params(**kwargs) + model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider( model=model, custom_llm_provider=custom_llm_provider, api_base=api_base, api_key=api_key, + litellm_params=GenericLiteLLMParams(**kwargs), ) if dynamic_api_key is not None: api_key = dynamic_api_key - allowed_openai_params: Final[list[str] | None] = kwargs.get("allowed_openai_params", None) + allowed_openai_params: Final[list[str] | None] = ( + cast( # cast-ok: allowed_openai_params arrives as a str list in the request kwargs + "list[str] | None", kwargs.get("allowed_openai_params", None) + ) + ) optional_params: Final = get_optional_params_embeddings( model=model, user=user, @@ -6518,8 +6543,6 @@ def embedding( model_info=kwargs.get("model_info"), ) - litellm_params_dict: Final = get_litellm_params(**kwargs) - logging: Final[LiteLLMLoggingObj] = litellm_logging_obj logging.update_environment_variables( model=model, diff --git a/litellm/models/credentials.py b/litellm/models/credentials.py index 730778a323b..d569ea12617 100644 --- a/litellm/models/credentials.py +++ b/litellm/models/credentials.py @@ -22,7 +22,9 @@ class CredentialBase(LiteLLMBaseModel): class CredentialItem(CredentialBase): - credential_values: dict + credential_values: dict[ # mutable-ok: values are decrypted and merged in place + str, object + ] # PATCH-only instruction naming keys to drop from the stored credential_values. It describes an # edit rather than the credential, so it stays out of dumps: those feed config loading, the DB # write, and the in-memory list, none of which have a place for it. @@ -56,3 +58,68 @@ class UpdateCredentialItem(LiteLLMBaseModel): credential_values: Mapping[str, object] | None = None model_id: str | None = None credential_values_to_delete: tuple[str, ...] | None = None + + +class UserProviderConnection(LiteLLMBaseModel): + model_config = ConfigDict(frozen=True) + + credential_name: str + provider: str + connected: bool + github_login: str | None = None + connected_at: str | None = None + + +class UserProviderConnectionsResponse(LiteLLMBaseModel): + model_config = ConfigDict(frozen=True) + + connections: list[UserProviderConnection] # mutable-ok: pydantic response model field + + +class UserConnectionStartResponse(LiteLLMBaseModel): + model_config = ConfigDict(frozen=True) + + user_code: str + verification_uri: str + expires_in: int + interval: int + flow_handle: str + + +class UserConnectionPollRequest(LiteLLMBaseModel): + model_config = ConfigDict(frozen=True) + + flow_handle: str + + +class UserConnectionFlowHandle(LiteLLMBaseModel): + """The stateless device-flow ticket: the pending ``device_code`` travels + encrypted inside ``flow_handle`` instead of a worker-local cache entry, so a + poll landing on any worker can complete the connection.``""" + + model_config = ConfigDict(frozen=True) + + user_id: str + credential_name: str + device_code: str = Field(repr=False) + interval: int + expires_at: float + + +UserConnectionPollStatus: TypeAlias = Literal[ + "pending", "slow_down", "expired", "denied", "connected", "no_copilot_seat" +] + + +class UserConnectionPollResponse(LiteLLMBaseModel): + model_config = ConfigDict(frozen=True) + + status: UserConnectionPollStatus + interval: int | None = None + github_login: str | None = None + + +class UserConnectionDeleteResponse(LiteLLMBaseModel): + model_config = ConfigDict(frozen=True) + + status: str diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 54705a8e171..44f62fa2b26 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -438,6 +438,12 @@ _BANNED_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = ( # SDK-only field; also rejected outright in is_request_body_safe. "model_list", "vertex_ai_credentials", + # Per-user GitHub Copilot connection slots: the mode is decided by the + # stored credential's values and the caller's connection by the proxy's own + # secret_fields, so a body-supplied value could only spoof either. + "github_copilot_auth_type", + "user_provider_credentials", + "github_copilot_user_session", # Observability credentials, hosts, and project identifiers: derived # from the canonical ``_supported_callback_params`` allowlist so new # integrations are covered automatically. Sorted for stable iteration diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 2c1643c7c42..bddda7da546 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -70,6 +70,16 @@ _PROXY_ADMIN_VIEW_ONLY_BLOCKED_KEY_SUFFIXES: Final = ("/regenerate", "/reset_spe _AUTH_ENFORCED_PASS_THROUGH_ROUTE_GROUPS: Final = frozenset(("openai_routes", "llm_api_routes")) +# Method-scoped because the generic credential CRUD handlers share these paths: +# the route check cannot admit PATCH/DELETE on a name like "user_connections" +# or "*/user_connection*" without also opening the admin CRUD handlers. +_INTERNAL_USER_CREDENTIAL_CONNECTION_ROUTES: Final = ( + ("GET", "/credentials/user_connections"), + ("POST", "/credentials/{credential_name:path}/user_connection/start"), + ("POST", "/credentials/{credential_name:path}/user_connection/poll"), + ("DELETE", "/credentials/{credential_name:path}/user_connection"), +) + class RouteChecks: @staticmethod @@ -331,7 +341,10 @@ class RouteChecks: ) elif ( _user_role == LitellmUserRoles.INTERNAL_USER.value - and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.internal_user_routes.value) + and ( + RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.internal_user_routes.value) + or RouteChecks.is_internal_user_credential_connection_route(route=route, request=request) + ) or user_is_org_admin(request_data=request_data, user_object=user_obj) and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.org_admin_allowed_routes.value) or _user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value @@ -663,7 +676,9 @@ class RouteChecks: return None try: - method: Final = request.method + method: Final = cast( # cast-ok: request.method is str at runtime; tests hand a MagicMock + object, request.method + ) except (AttributeError, KeyError): return None if not isinstance(method, str): @@ -673,6 +688,16 @@ class RouteChecks: _get_request_method = get_request_method + @staticmethod + def is_internal_user_credential_connection_route(route: str, request: Request | None) -> bool: + method: Final = RouteChecks.get_request_method(request) + if method is None: + return False + return any( + method == expected_method and RouteChecks.check_route_access(route=route, allowed_routes=(pattern,)) + for expected_method, pattern in _INTERNAL_USER_CREDENTIAL_CONNECTION_ROUTES + ) + @staticmethod def is_auth_enforced_pass_through_route(route: str, method: str | None = None) -> bool: """ diff --git a/litellm/proxy/common_utils/credential_hydration.py b/litellm/proxy/common_utils/credential_hydration.py index 2a66a3ec2c4..a4e858692e6 100644 --- a/litellm/proxy/common_utils/credential_hydration.py +++ b/litellm/proxy/common_utils/credential_hydration.py @@ -9,7 +9,7 @@ import asyncio from collections.abc import Mapping from itertools import chain from types import MappingProxyType -from typing import Final +from typing import Final, cast # noqa: TID251 # narrows untyped credential values at the decrypt boundary import litellm from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper @@ -71,7 +71,10 @@ def decrypted_or_stored(key: str, value: str) -> str: def _decrypted(db_credential: CredentialItem) -> CredentialItem: """The stored credential with every value decrypted, leaving already-plaintext values alone.""" decrypted_values: Final = MappingProxyType( - {key: decrypted_or_stored(key, value) for key, value in db_credential.credential_values.items()} + { + key: decrypted_or_stored(key, cast("str", value)) # cast-ok: stored values are str + for key, value in db_credential.credential_values.items() + } ) return CredentialItem( credential_name=db_credential.credential_name, diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 877fd058eb6..83af63bf24a 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -2,7 +2,9 @@ CRUD endpoints for storing reusable credentials. """ +import time from collections.abc import Mapping +from datetime import datetime from typing import ( Annotated, Final, @@ -10,10 +12,15 @@ from typing import ( ) from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response, status -from pydantic import TypeAdapter +from pydantic import TypeAdapter, ValidationError import litellm +from litellm._internal_context import with_service_target from litellm._logging import verbose_proxy_logger +from litellm.constants import ( + GITHUB_COPILOT_AUTH_TYPE_KEY, + GITHUB_COPILOT_PER_USER_AUTH_TYPE, +) from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.litellm_logging import get_masked_values from litellm.llms.anthropic.wif import ( @@ -22,7 +29,17 @@ from litellm.llms.anthropic.wif import ( UnbuildableIdentitySource, anthropic_internal_issuer_jwks, ) -from litellm.models.credentials import CredentialView, UpdateCredentialItem +from litellm.models.credentials import ( + CredentialView, + UpdateCredentialItem, + UserConnectionDeleteResponse, + UserConnectionFlowHandle, + UserConnectionPollRequest, + UserConnectionPollResponse, + UserConnectionStartResponse, + UserProviderConnection, + UserProviderConnectionsResponse, +) from litellm.proxy._types import ( CommonProxyErrors, LitellmUserRoles, @@ -36,15 +53,31 @@ from litellm.proxy.common_utils.credential_hydration import ( named_credential_wif_fields, stored_credential_provider, ) -from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_value_helper, + encrypt_value_helper, +) from litellm.proxy.utils import handle_exception_on_proxy, jsonify_object from litellm.repositories.base_repository import is_unique_violation from litellm.repositories.credentials_repository import CredentialsRepository from litellm.types.router import server_owned_wif_fields_named from litellm.types.utils import CreateCredentialItem, CredentialItem +from .user_provider_credentials import ( + GithubCopilotUserConnectionPayload, + decode_user_provider_credential, + delete_user_provider_credential, + delete_user_provider_credentials_for_credential, + invalidate_user_provider_credential_cache, + list_user_provider_credentials, + list_user_provider_credentials_for_credential, + set_user_provider_credential_cache, + upsert_user_provider_credential, +) + router: Final = APIRouter() _CREDENTIAL_DICT_ADAPTER: Final = TypeAdapter(dict[str, object]) +_GITHUB_COPILOT_PROVIDER: Final = "github_copilot" _DISPLAY_NAME_MAX_LENGTH: Final = 255 @@ -140,9 +173,12 @@ class CredentialHelperUtils: @staticmethod def encrypt_credential_values(credential: CredentialItem, new_encryption_key: str | None = None) -> CredentialItem: """Encrypt values in credential.credential_values and add to DB""" - encrypted_credential_values: Final = {} + encrypted_credential_values: Final[dict[str, object]] = {} # mutable-ok: built one entry per credential value for key, value in (credential.credential_values or {}).items(): - encrypted_credential_values[key] = encrypt_value_helper(value, new_encryption_key) + encrypted_credential_values[key] = encrypt_value_helper( + cast("str", value), # cast-ok: credential values are str at the encryption boundary + new_encryption_key, + ) # Return a new object to avoid mutating the caller's credential, which # is kept in memory and should remain unencrypted. @@ -452,6 +488,297 @@ async def get_credential_by_model( raise handle_exception_on_proxy(e) +def _authenticated_user_id(user_api_key_dict: UserAPIKeyAuth) -> str: + user_id: Final = user_api_key_dict.user_id + if not isinstance(user_id, str) or not user_id: + raise HTTPException(status_code=401, detail="An authenticated user is required") + return user_id + + +def _per_user_credential_or_404(credential_name: str) -> CredentialItem: + credential: Final = CredentialAccessor.find_credential(credential_name) + values: Final = credential.credential_values if credential is not None else None + if ( + credential is None + or not isinstance(values, Mapping) + or values.get(GITHUB_COPILOT_AUTH_TYPE_KEY) != GITHUB_COPILOT_PER_USER_AUTH_TYPE + ): + raise HTTPException(status_code=404, detail="Credential not found") + return credential + + +def _per_user_credential_names() -> tuple[str, ...]: + return tuple( + credential.credential_name + for credential in litellm.credential_list + if cast( # cast-ok: credential_values is a plain dict at runtime + Mapping[object, object], credential.credential_values + ).get(GITHUB_COPILOT_AUTH_TYPE_KEY) + == GITHUB_COPILOT_PER_USER_AUTH_TYPE + ) + + +@router.get( + "/credentials/user_connections", + dependencies=[Depends(user_api_key_auth)], + tags=["credential management"], + response_model=UserProviderConnectionsResponse, +) +async def list_user_connections( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI resolves the dependency from the default +) -> UserProviderConnectionsResponse: + """List the calling user's per-user provider connections.""" + from litellm.proxy.proxy_server import prisma_client + + try: + user_id: Final = _authenticated_user_id(user_api_key_dict) + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + rows: Final = await list_user_provider_credentials(prisma_client, user_id) + github_logins: Final[dict[str, str]] = {} # mutable-ok: accumulates one entry per credential row + connected_at: Final[dict[str, str]] = {} # mutable-ok: accumulates one entry per credential row + for row in rows: + if row.provider != _GITHUB_COPILOT_PROVIDER: + continue + payload = decode_user_provider_credential(row.credential_b64) + if payload is not None: + github_logins[row.credential_name] = payload.github_login + updated = getattr(row, "updated_at", None) + if isinstance(updated, datetime): + connected_at[row.credential_name] = updated.isoformat() + return UserProviderConnectionsResponse( + connections=[ + UserProviderConnection( + credential_name=name, + provider=_GITHUB_COPILOT_PROVIDER, + connected=name in github_logins, + github_login=github_logins.get(name), + connected_at=connected_at.get(name), + ) + for name in _per_user_credential_names() + ] + ) + except Exception as e: # noqa: BLE001 # endpoint boundary: every failure becomes the proxy error contract + raise handle_exception_on_proxy(e) + + +def _decode_flow_handle(flow_handle: str) -> UserConnectionFlowHandle | None: + decrypted: Final = decrypt_value_helper( + value=flow_handle, + key="device_flow_handle", + exception_type="debug", + return_original_value=False, + ) + if decrypted is None: + return None + try: + return UserConnectionFlowHandle.model_validate_json(decrypted) + except ValidationError: + return None + + +@router.post( + "/credentials/{credential_name:path}/user_connection/start", + dependencies=[Depends(user_api_key_auth)], + tags=["credential management"], + response_model=UserConnectionStartResponse, +) +@with_service_target("user_provider_connections") +async def start_user_connection( + request: Request, + fastapi_response: Response, + credential_name: str = Path(..., description="The credential name, percent-decoded"), + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI resolves the dependency from the default +) -> UserConnectionStartResponse: + """Begin a GitHub device flow for the calling user's connection to a per-user credential.""" + from litellm.llms.github_copilot.per_user_auth import astart_device_flow + from litellm.proxy.proxy_server import prisma_client + + try: + user_id: Final = _authenticated_user_id(user_api_key_dict) + _per_user_credential_or_404(credential_name) + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + flow: Final = await astart_device_flow() + handle: Final = encrypt_value_helper( + UserConnectionFlowHandle( + user_id=user_id, + credential_name=credential_name, + device_code=flow.device_code, + interval=flow.interval, + expires_at=time.time() + flow.expires_in, + ).model_dump_json() + ) + return UserConnectionStartResponse( + user_code=flow.user_code, + verification_uri=flow.verification_uri, + expires_in=flow.expires_in, + interval=flow.interval, + flow_handle=handle, + ) + except Exception as e: # noqa: BLE001 # endpoint boundary: every failure becomes the proxy error contract + raise handle_exception_on_proxy(e) + + +@router.post( + "/credentials/{credential_name:path}/user_connection/poll", + dependencies=[Depends(user_api_key_auth)], + tags=["credential management"], + response_model=UserConnectionPollResponse, +) +@with_service_target("user_provider_connections") +async def poll_user_connection( + request: Request, + fastapi_response: Response, + body: UserConnectionPollRequest, + credential_name: str = Path(..., description="The credential name, percent-decoded"), + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI resolves the dependency from the default +) -> UserConnectionPollResponse: + """Poll the device flow once and persist the connection on completion.""" + from litellm.llms.github_copilot.per_user_auth import ( + acheck_copilot_seat, + afetch_github_login, + apoll_device_flow, + ) + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + try: + user_id: Final = _authenticated_user_id(user_api_key_dict) + _per_user_credential_or_404(credential_name) + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + handle: Final = _decode_flow_handle(body.flow_handle) + if ( + handle is None + or handle.user_id != user_id + or handle.credential_name != credential_name + or handle.expires_at <= time.time() + ): + raise HTTPException(status_code=400, detail="invalid or expired flow_handle") + + poll: Final = await apoll_device_flow(handle.device_code) + if poll.status in ("expired", "denied", "pending", "slow_down"): + return UserConnectionPollResponse(status=poll.status, interval=poll.interval) + + github_token: Final = poll.access_token + if not github_token: + raise HTTPException(status_code=502, detail="GitHub device flow completed without a token") + try: + await acheck_copilot_seat(user_id=user_id, github_token=github_token, credential_name=credential_name) + except litellm.CallerCredentialAuthenticationError: + return UserConnectionPollResponse(status="no_copilot_seat") + github_login: Final = await afetch_github_login(github_token) + connection_payload: Final = GithubCopilotUserConnectionPayload( + access_token=github_token, github_login=github_login + ) + await upsert_user_provider_credential( + prisma_client, + user_id, + credential_name, + _GITHUB_COPILOT_PROVIDER, + connection_payload, + ) + if not await set_user_provider_credential_cache( + user_api_key_cache, user_id, credential_name, connection_payload + ): + raise HTTPException( + status_code=503, + detail="GitHub connection saved, but the cache could not be refreshed; " + "requests may be rejected for up to 60 seconds", + ) + return UserConnectionPollResponse(status="connected", github_login=github_login) + except litellm.CallerCredentialRateLimitError as e: + raise HTTPException(status_code=429, detail=str(e)) + except Exception as e: # noqa: BLE001 # every remaining failure maps to the proxy error shape + raise handle_exception_on_proxy(e) + + +@router.delete( + "/credentials/{credential_name:path}/user_connection", + dependencies=[Depends(user_api_key_auth)], + tags=["credential management"], + response_model=UserConnectionDeleteResponse, +) +@with_service_target("user_provider_connections") +async def delete_user_connection( + request: Request, + fastapi_response: Response, + credential_name: str = Path(..., description="The credential name, percent-decoded"), + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI resolves the dependency from the default +) -> UserConnectionDeleteResponse: + """Disconnect the calling user's stored GitHub token for a per-user credential. Idempotent.""" + from litellm.llms.github_copilot.per_user_auth import evict_copilot_user_session + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + try: + user_id: Final = _authenticated_user_id(user_api_key_dict) + _per_user_credential_or_404(credential_name) + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + # Tombstone first: while a token may still be cached, the DB row stays. + # If the tombstone write fails and the row were deleted, the stale token + # would keep working until its TTL with no connection left to re-check. + tombstoned: Final = await invalidate_user_provider_credential_cache( + user_api_key_cache, user_id, credential_name + ) + if not tombstoned: + raise HTTPException( + status_code=503, + detail={"error": "could not revoke cached connection, retry"}, + ) + prior: Final = await delete_user_provider_credential(prisma_client, user_id, credential_name) + if prior is not None: + evict_copilot_user_session(user_id, prior.access_token) + return UserConnectionDeleteResponse(status="disconnected") + except Exception as e: # noqa: BLE001 # endpoint boundary: every failure becomes the proxy error contract + raise handle_exception_on_proxy(e) + + +@with_service_target("user_provider_connections") +async def _purge_user_connections_for_credential(credential_name: str) -> None: + """Drop every user connection under a per-user credential and clear caches. + + Shared cleanup for DELETE (credential removed) and PATCH (renamed or no + longer per-user).""" + from litellm.llms.github_copilot.per_user_auth import evict_copilot_user_session + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + if prisma_client is None: + return + try: + prior_rows: Final = await list_user_provider_credentials_for_credential(prisma_client, credential_name) + user_ids: Final = await delete_user_provider_credentials_for_credential(prisma_client, credential_name) + except Exception: # noqa: BLE001 # best-effort purge must not break credential deletion + verbose_proxy_logger.exception( + "_purge_user_connections_for_credential: failed to drop user connections for %s", credential_name + ) + return + access_tokens: Final[dict[str, str]] = {} # mutable-ok: accumulates one entry per deleted row + for row in prior_rows: + decoded = decode_user_provider_credential(row.credential_b64) + if decoded is not None: + access_tokens[row.user_id] = decoded.access_token + for user_id in user_ids: + await invalidate_user_provider_credential_cache(user_api_key_cache, user_id, credential_name) + token = access_tokens.get(user_id) + if token: + evict_copilot_user_session(user_id, token) + + @router.delete( "/credentials/{credential_name:path}", dependencies=[Depends(user_api_key_auth)], @@ -485,6 +812,7 @@ async def delete_credential( ## DELETE FROM LITELLM ## litellm.credential_list = [cred for cred in litellm.credential_list if cred.credential_name != credential_name] + await _purge_user_connections_for_credential(credential_name) return {"success": True, "message": "Credential deleted successfully"} except Exception as e: raise handle_exception_on_proxy(e) @@ -596,6 +924,13 @@ async def update_credential( # Sync in-memory credential_list (skip if not in memory - e.g., proxy restarted) _sync_in_memory_credential(patch, credential_name) + merged_values: Final = cast( # cast-ok: credential_values is a plain dict at runtime + Mapping[object, object], merged_credential.credential_values + ) + still_per_user: Final = merged_values.get(GITHUB_COPILOT_AUTH_TYPE_KEY) == GITHUB_COPILOT_PER_USER_AUTH_TYPE + if not still_per_user: + await _purge_user_connections_for_credential(credential_name) + return {"success": True, "message": "Credential updated successfully"} except Exception as e: raise handle_exception_on_proxy(e) diff --git a/litellm/proxy/credential_endpoints/user_provider_credentials.py b/litellm/proxy/credential_endpoints/user_provider_credentials.py new file mode 100644 index 00000000000..4253483a228 --- /dev/null +++ b/litellm/proxy/credential_endpoints/user_provider_credentials.py @@ -0,0 +1,351 @@ +"""DB and cache helpers for per-user provider connections. + +A row in ``LiteLLM_UserProviderCredentials`` binds one LiteLLM user to one +admin-named LLM credential (e.g. the user's own GitHub token behind a +``github_copilot`` per-user credential). ``credential_b64`` always stores the +payload encrypted; the request-time cache carries only ciphertext plus a +negative marker so an unconnected user does not hit the DB per request. +""" + +import json +from collections.abc import Mapping, Sequence +from typing import ( + TYPE_CHECKING, + Final, + Protocol, + cast, # noqa: TID251 # narrows the untyped Redis cache to the str-keyed Protocol +) + +from pydantic import BaseModel, ConfigDict, Field, ValidationError + +from litellm._internal_context import with_service_target +from litellm._logging import verbose_proxy_logger +from litellm.caching.dual_cache import DualCache + +if TYPE_CHECKING: + from litellm.caching.redis_cache import RedisCache +from litellm.constants import ( + GITHUB_COPILOT_USER_CREDENTIAL_CACHE_TTL_SECONDS, + USER_PROVIDER_CREDENTIAL_CACHE_PREFIX, + USER_PROVIDER_CREDENTIAL_NOT_CONNECTED, +) +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_value_helper, + encrypt_value_helper, +) +from litellm.repositories.chunked_in import find_many_in +from litellm.repositories.prisma_protocols import TableActions +from litellm.repositories.table_repositories import PrismaTableRepository + +if TYPE_CHECKING: + from prisma import models as prisma_models + + from litellm.proxy.utils import PrismaClient + +_NOT_CONNECTED: Final = USER_PROVIDER_CREDENTIAL_NOT_CONNECTED + + +class GithubCopilotUserConnectionPayload(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + access_token: str = Field(repr=False) + github_login: str + + +class _UserProviderCredentialsRepository(PrismaTableRepository["prisma_models.LiteLLM_UserProviderCredentials"]): + table_name = "litellm_userprovidercredentials" + + +def _table( + prisma_client: "PrismaClient", +) -> TableActions["prisma_models.LiteLLM_UserProviderCredentials"]: + return _UserProviderCredentialsRepository(prisma_client, use_writer=True).table + + +def _cache_key(user_id: str, credential_name: str) -> str: + pair: Final = json.dumps([user_id, credential_name], separators=(",", ":")) + return f"{USER_PROVIDER_CREDENTIAL_CACHE_PREFIX}:{pair}" + + +def _encode(payload: GithubCopilotUserConnectionPayload) -> str: + return encrypt_value_helper(payload.model_dump_json()) + + +def decode_user_provider_credential(stored: str) -> GithubCopilotUserConnectionPayload | None: + decrypted: Final = decrypt_value_helper( + value=stored, + key="user_provider_credential", + exception_type="debug", + return_original_value=False, + ) + if decrypted is None: + return None + try: + return GithubCopilotUserConnectionPayload.model_validate_json(decrypted) + except ValidationError: + return None + + +async def upsert_user_provider_credential( + prisma_client: "PrismaClient", + user_id: str, + credential_name: str, + provider: str, + payload: GithubCopilotUserConnectionPayload, +) -> None: + await _table(prisma_client).upsert( + where={"user_id_credential_name": {"user_id": user_id, "credential_name": credential_name}}, + data={ + "create": { + "user_id": user_id, + "credential_name": credential_name, + "provider": provider, + "credential_b64": _encode(payload), + }, + "update": {"credential_b64": _encode(payload), "provider": provider}, + }, + ) + + +async def get_user_provider_credential( + prisma_client: "PrismaClient", + user_id: str, + credential_name: str, +) -> GithubCopilotUserConnectionPayload | None: + row: Final = await _table(prisma_client).find_unique( + where={"user_id_credential_name": {"user_id": user_id, "credential_name": credential_name}} + ) + if row is None: + return None + return decode_user_provider_credential(row.credential_b64) + + +async def delete_user_provider_credential( + prisma_client: "PrismaClient", + user_id: str, + credential_name: str, +) -> GithubCopilotUserConnectionPayload | None: + existing: Final = await get_user_provider_credential(prisma_client, user_id, credential_name) + await _table(prisma_client).delete_many(where={"user_id": user_id, "credential_name": credential_name}) + return existing + + +async def delete_user_provider_credentials_for_credential( + prisma_client: "PrismaClient", + credential_name: str, +) -> Sequence[str]: + """Delete every user connection under ``credential_name``; returns the + affected user_ids so callers can invalidate their cache entries.""" + rows: Final = await _table(prisma_client).find_many(where={"credential_name": credential_name}) + user_ids: Final = tuple(dict.fromkeys(row.user_id for row in rows)) + await _table(prisma_client).delete_many(where={"credential_name": credential_name}) + return user_ids + + +async def list_user_provider_credentials( + prisma_client: "PrismaClient", + user_id: str, +) -> Sequence["prisma_models.LiteLLM_UserProviderCredentials"]: + return await _table(prisma_client).find_many(where={"user_id": user_id}) + + +async def list_user_provider_credentials_for_credential( + prisma_client: "PrismaClient", + credential_name: str, +) -> Sequence["prisma_models.LiteLLM_UserProviderCredentials"]: + return await _table(prisma_client).find_many(where={"credential_name": credential_name}) + + +class _StringKeyCache(Protocol): + async def async_get_cache( + self, + key: str, + **kwargs: object, # kwargs-ok: RedisCache accepts optional cache kwargs + ) -> object: ... + + async def async_set_cache( + self, + key: str, + value: str, + **kwargs: object, # kwargs-ok: RedisCache accepts ttl/nx kwargs + ) -> None: ... + + async def async_delete_cache( + self, + key: str, + **kwargs: object, # kwargs-ok: RedisCache accepts optional cache kwargs + ) -> None: ... + + +def _string_cache(cache: "RedisCache") -> _StringKeyCache: + return cast(_StringKeyCache, cache) # cast-ok: RedisCache exposes the str-keyed cache protocol + + +async def _try_cache_get(token_cache: "RedisCache", key: str) -> object: + """Redis errors must read as a plain miss: the DB is the source of truth.""" + try: + return await _string_cache(token_cache).async_get_cache(key) + except Exception: # noqa: BLE001 # a Redis outage must fall back to the database, never reject a connected user + verbose_proxy_logger.warning("aget_user_provider_tokens: Redis get failed; falling back to the database") + return None + + +async def _try_cache_set(token_cache: "RedisCache", key: str, value: str, nx: bool = False) -> None: + try: + await _string_cache(token_cache).async_set_cache( + key, value, nx=nx, ttl=GITHUB_COPILOT_USER_CREDENTIAL_CACHE_TTL_SECONDS + ) + except Exception: # noqa: BLE001 # caching is best-effort; a Redis outage is not worth failing the request + verbose_proxy_logger.warning("aget_user_provider_tokens: Redis set failed; skipping the cache write") + + +@with_service_target("user_provider_connections") +async def invalidate_user_provider_credential_cache( + cache: DualCache, + user_id: str, + credential_name: str, +) -> bool: + """Write the not-connected tombstone rather than deleting: a read in flight + before the disconnect must not fill the token back in behind it (fills are + set-if-absent). Returns False when Redis is attached and the write fails, so + the disconnect can refuse to drop the row while a stale token might linger.""" + token_cache: Final = cache.redis_cache + if token_cache is None: + return True + try: + await _string_cache(token_cache).async_set_cache( + _cache_key(user_id, credential_name), + _NOT_CONNECTED, + ttl=GITHUB_COPILOT_USER_CREDENTIAL_CACHE_TTL_SECONDS, + ) + except Exception: # noqa: BLE001 # caller decides whether a Redis outage aborts the disconnect + verbose_proxy_logger.warning("invalidate_user_provider_credential_cache: Redis tombstone failed") + return False + # async_set_cache swallows client errors internally, so a write that never + # landed looks identical to a success. The tombstone is only trusted when the + # key reads back as _NOT_CONNECTED; anything else means revoke failed. + readback: Final = await _try_cache_get(token_cache, _cache_key(user_id, credential_name)) + if readback != _NOT_CONNECTED: + verbose_proxy_logger.warning("invalidate_user_provider_credential_cache: tombstone not visible after write") + return False + return True + + +@with_service_target("user_provider_connections") +async def drop_user_provider_credential_cache( + cache: DualCache, + user_id: str, + credential_name: str, +) -> None: + """Connect path: delete the key outright so the next read fills fresh (a + tombstone would linger for the TTL and hide the new connection).""" + token_cache: Final = cache.redis_cache + if token_cache is None: + return + try: + await _string_cache(token_cache).async_delete_cache(_cache_key(user_id, credential_name)) + except Exception: # noqa: BLE001 # a Redis outage must not fail a connect + verbose_proxy_logger.warning("drop_user_provider_credential_cache: Redis delete failed") + + +@with_service_target("user_provider_connections") +async def set_user_provider_credential_cache( + cache: DualCache, + user_id: str, + credential_name: str, + payload: GithubCopilotUserConnectionPayload, +) -> bool: + """Connect path: overwrite the key with the new connection's ciphertext so a + stale set-if-absent fill from a pre-connect read cannot resurrect a + not-connected marker. The overwrite is verified by decoding the key back + (async_set_cache swallows write errors); on a mismatch the key is deleted, + and False means a stale entry survived both attempts.""" + token_cache: Final = cache.redis_cache + if token_cache is None: + return True + key: Final = _cache_key(user_id, credential_name) + try: + await _string_cache(token_cache).async_set_cache( + key, + _encode(payload), + ttl=GITHUB_COPILOT_USER_CREDENTIAL_CACHE_TTL_SECONDS, + ) + readback: Final = await _try_cache_get(token_cache, key) + if isinstance(readback, str) and decode_user_provider_credential(readback) == payload: + return True + verbose_proxy_logger.warning( + "set_user_provider_credential_cache: Redis overwrite not visible; falling back to delete" + ) + except Exception: # noqa: BLE001 # a Redis outage must not fail a connect + verbose_proxy_logger.warning("set_user_provider_credential_cache: Redis set failed; falling back to delete") + try: + await _string_cache(token_cache).async_delete_cache(key) + except Exception: # noqa: BLE001 # a Redis outage must not fail a connect + verbose_proxy_logger.warning("set_user_provider_credential_cache: Redis delete failed") + after_delete: Final = await _try_cache_get(token_cache, key) + if after_delete is None: + return True + if isinstance(after_delete, str) and decode_user_provider_credential(after_delete) == payload: + return True + verbose_proxy_logger.warning("set_user_provider_credential_cache: stale entry survived delete") + return False + + +@with_service_target("user_provider_connections") +async def aget_user_provider_tokens( + prisma_client: "PrismaClient", + cache: DualCache, + user_id: str, + credential_names: Sequence[str], +) -> Mapping[str, str]: + """Map each per-user credential name to the caller's stored GitHub token. + + Cached values are the stored ciphertext or the ``_NOT_CONNECTED`` marker; + plaintext tokens only ever live in the returned dict. Redis is the only + cache layer used: without it there is no shared invalidation, so a worker + serving a stale local entry after a disconnect is worse than a DB read.""" + token_cache: Final = cache.redis_cache + names: Final = tuple(dict.fromkeys(credential_names)) + if not names: + return {} + cached: Final[dict[str, str]] = {} # mutable-ok: accumulates hits and DB reads + misses: Final[list[str]] = [] # mutable-ok: accumulates cache misses + for name in names: + value = await _try_cache_get(token_cache, _cache_key(user_id, name)) if token_cache is not None else None + if value == _NOT_CONNECTED: + continue + if isinstance(value, str) and value: + cached[name] = value + else: + misses.append(name) + if not misses: + return { + name: token + for name, ciphertext in cached.items() + if (token := _token_from_ciphertext(ciphertext)) is not None + } + + try: + rows: Final = await find_many_in( + _table(prisma_client), + "credential_name", + misses, + where={"user_id": user_id}, + ) + except Exception: # noqa: BLE001 # a DB outage reads as not connected so the caller gets a 401, never a 500 + verbose_proxy_logger.exception("aget_user_provider_tokens: DB read failed for user_id=%s", user_id) + return {} + found: Final = {row.credential_name: row.credential_b64 for row in rows} + for name in misses: + if token_cache is not None: + await _try_cache_set(token_cache, _cache_key(user_id, name), found.get(name, _NOT_CONNECTED), nx=True) + if name in found: + cached[name] = found[name] + return { + name: token for name, ciphertext in cached.items() if (token := _token_from_ciphertext(ciphertext)) is not None + } + + +def _token_from_ciphertext(ciphertext: str) -> str | None: + payload: Final = decode_user_provider_credential(ciphertext) + return payload.access_token if payload is not None else None diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 84e5e113e34..490b393aa68 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1,5 +1,6 @@ import asyncio import copy +import itertools import json import re import time @@ -7,7 +8,14 @@ from collections import OrderedDict from collections.abc import Mapping, MutableMapping, Sequence from datetime import datetime from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, cast +from typing import ( + TYPE_CHECKING, + Any, + Final, + Literal, + TypeAlias, + cast, # noqa: TID251 # narrows untyped deployment and credential payloads +) from fastapi import HTTPException, Request from pydantic import TypeAdapter @@ -120,7 +128,7 @@ def _trace_id_from_otel_span(span: "OtelSpan | None") -> str | None: try: span_context: Final = span.get_span_context() is_valid: Final = span_context.is_valid - trace_id: Final = span_context.trace_id + trace_id: Final = cast(object, span_context.trace_id) # cast-ok: trace_id is int at runtime; tests mock it except AttributeError: return None if not is_valid or not isinstance(trace_id, int): @@ -296,6 +304,7 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = ( "proxy_server_request", "standard_logging_object", "secret_fields", + "github_copilot_user_session", "mock_response", "mock_tool_calls", "disable_global_guardrails", @@ -1948,7 +1957,9 @@ class LiteLLMProxyRequestSetup: return TeamCallbackMetadata( success_callback=team_config.get("success_callback", None), failure_callback=team_config.get("failure_callback", None), - callback_vars=callback_vars_dict, + callback_vars=cast( # cast-ok: callback_vars_dict values are str-coerced above + "dict[str, str]", callback_vars_dict + ), ) @staticmethod @@ -2278,6 +2289,9 @@ async def add_litellm_data_to_request( data=data, headers=_headers, user_api_key_dict=user_api_key_dict ) + # Header mappings can overwrite user_id below; the key/JWT identity is the only + # one allowed to load per-user provider credentials. + authenticated_user_id: Final = user_api_key_dict.user_id user_api_key_dict = LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping( general_settings, user_api_key_dict, _headers ) @@ -2537,23 +2551,17 @@ async def add_litellm_data_to_request( if ( general_settings is not None and general_settings.get("use_x_forwarded_for") is True - and request is not None and hasattr(request, "headers") and "x-forwarded-for" in request.headers ): requester_ip_address = request.headers["x-forwarded-for"] - elif ( - request is not None - and hasattr(request, "client") - and hasattr(request.client, "host") - and request.client is not None - ): + elif hasattr(request, "client") and hasattr(request.client, "host") and request.client is not None: requester_ip_address = request.client.host data[_metadata_variable_name]["requester_ip_address"] = requester_ip_address # Add User-Agent user_agent = "" - if request is not None and hasattr(request, "headers") and "user-agent" in request.headers: + if hasattr(request, "headers") and "user-agent" in request.headers: user_agent = request.headers["user-agent"] data[_metadata_variable_name]["user_agent"] = user_agent @@ -2645,17 +2653,35 @@ async def add_litellm_data_to_request( getattr(request.state, "litellm_roi_estimator", False) is True ) - verbose_proxy_logger.debug("[PROXY] returned data from litellm_pre_call_utils: %s", data) + verbose_proxy_logger.debug( + "[PROXY] returned data from litellm_pre_call_utils: %s", + data, + ) # Team/Project credential overrides from model_config # Placed after the debug log to avoid leaking credential secrets in logs _apply_credential_overrides_from_model_config( data=data, user_api_key_dict=user_api_key_dict, - pre_alias_model_name=_pre_alias_model, + pre_alias_model_name=cast("str | None", _pre_alias_model), # cast-ok: the pre-alias model name arrives as a str llm_router=llm_router, ) + await _resolve_user_provider_credentials_for_request( + data=cast("dict[str, object]", data), # cast-ok: data is the untyped request body dict + authenticated_user_id=authenticated_user_id, + team_id=user_api_key_dict.team_id, + llm_router=llm_router, + router_settings=( + cast("dict[str, object]", rs) # cast-ok: router_settings is a plain config dict + if isinstance( + rs := cast("object", user_api_key_dict.router_settings), # cast-ok: config dict + dict, + ) + else None + ), + ) + ## ENFORCED PARAMS CHECK # loop through each enforced param # example enforced_params ['user', 'metadata', 'metadata.generation_name'] @@ -2681,6 +2707,413 @@ async def add_litellm_data_to_request( return data +_REQUEST_FALLBACK_KEYS: Final = ("fallbacks", "context_window_fallbacks", "content_policy_fallbacks") +_FALLBACK_DISCOVERY_LIMIT: Final = 256 + + +def _per_user_oauth_configured() -> bool: + """Whether any configured credential opted into per-user GitHub OAuth. + + Cheap in-memory precondition for all of the fallback discovery below: with + no per-user credential there is nothing to discover, so every code path in + this section (traversal, request-list size limit) must stay byte-identical + to shared-mode behavior.""" + from litellm.constants import ( + GITHUB_COPILOT_AUTH_TYPE_KEY, + GITHUB_COPILOT_PER_USER_AUTH_TYPE, + ) + + return any( + isinstance(values := getattr(credential, "credential_values", None), Mapping) + and cast("Mapping[object, object]", values).get( # cast-ok: credential_values is a plain dict at runtime + GITHUB_COPILOT_AUTH_TYPE_KEY + ) + == GITHUB_COPILOT_PER_USER_AUTH_TYPE + for credential in (litellm.credential_list or ()) + ) + + +def _all_per_user_credential_names() -> tuple[str, ...]: + """Every credential name configured for per-user GitHub OAuth, so one DB + query can prefetch all of the caller's connections when an auto-router + deployment may route to a Copilot group at request time.""" + from litellm.constants import ( + GITHUB_COPILOT_AUTH_TYPE_KEY, + GITHUB_COPILOT_PER_USER_AUTH_TYPE, + ) + + return tuple( + credential.credential_name + for credential in (litellm.credential_list or ()) + if getattr(credential, "credential_name", None) + and isinstance(values := getattr(credential, "credential_values", None), Mapping) + and cast("Mapping[object, object]", values).get( # cast-ok: credential_values is a plain dict at runtime + GITHUB_COPILOT_AUTH_TYPE_KEY + ) + == GITHUB_COPILOT_PER_USER_AUTH_TYPE + ) + + +def _fallback_entry_target_count(entry: object) -> int: + """Work units one fallback entry carries, floored so nothing counts as zero: + a bare string or any non-dict counts 1; a dict entry counts one unit per + key plus one per target in its value lists, so empty lists and scalar + values still cost what the router scan costs.""" + if isinstance(entry, dict): + entry_dict: Final = cast("dict[object, object]", entry) # cast-ok: fallback entries are plain dicts at runtime + return sum( + max(1, len(cast("Sequence[object]", value))) # cast-ok: fallback values are scalar or list entries + if isinstance(value, list) + else 1 + for value in entry_dict.values() + ) + return 1 + + +def _fallback_lists( + data: Mapping[str, object], + llm_router: litellm.Router, + router_settings: Mapping[str, object] | None, + team_router_settings: Mapping[str, object] | None, +) -> tuple[Sequence[object], ...]: + """Every fallback list the router can route this request through: the three + router-level lists, the same three keys from the request body (which replace + the router's when present, so the union is the safe superset), and the + key/team ``router_settings.fallbacks`` the proxy prefers over router-level + fallbacks in ``_configured_fallbacks``.""" + lists: Final[list[Sequence[object]]] = [] # mutable-ok: accumulates each fallback list found + for router_list in ( + cast("object", llm_router.fallbacks), # cast-ok: router fallback attrs are untyped + cast("object", llm_router.context_window_fallbacks), # cast-ok: router fallback attrs are untyped + cast("object", llm_router.content_policy_fallbacks), # cast-ok: router fallback attrs are untyped + ): + if isinstance(router_list, list): + lists.append(cast("Sequence[object]", router_list)) # cast-ok: fallback lists hold mixed entry shapes + for key in _REQUEST_FALLBACK_KEYS: + entries: object = data.get(key) + if isinstance(entries, list): + lists.append(cast("Sequence[object]", entries)) # cast-ok: request fallback lists hold mixed entry shapes + for settings in (router_settings, team_router_settings): + key_fallbacks: object = settings.get("fallbacks") if settings is not None else None + if isinstance(key_fallbacks, list): + lists.append( + cast("Sequence[object]", key_fallbacks) # cast-ok: settings fallback lists hold mixed entry shapes + ) + return tuple(lists) + + +_FallbackIndex: TypeAlias = tuple[Mapping[str, tuple[str, ...]], tuple[object, ...]] + + +def _index_fallback_list(fallback_list: Sequence[object]) -> _FallbackIndex: + """Index one fallback list once: an exact ``{source_key: targets}`` map for + bare keys (duplicate keys union their targets, a superset of the router's + first-wins resolution), plus the entries only the router matcher can score + ("*" keys, provider-prefixed keys, bare-string generic targets).""" + exact: Final[dict[str, tuple[str, ...]]] = {} # mutable-ok: one entry per bare source key + fuzzy: Final[list[object]] = [] # mutable-ok: accumulates matcher-only entries + for entry in fallback_list: + if isinstance(entry, dict) and entry: + entry_dict = cast("dict[object, object]", entry) # cast-ok: fallback entries are plain dicts at runtime + key = next(iter(entry_dict)) + raw = entry_dict[key] + values = ( + cast("Sequence[object]", raw) # cast-ok: fallback values are scalar or list entries + if isinstance(raw, list) + else ([raw] if raw is not None else []) + ) + targets = tuple(name for value in values if (name := _fallback_target_name(value)) is not None) + if isinstance(key, str) and not ("*" in key or "/" in key): + exact[key] = exact.get(key, ()) + targets + else: + fuzzy.append(entry) + else: + fuzzy.append(entry) + return exact, tuple(fuzzy) + + +def _fuzzy_entry_matched(entry: object, resolved: object) -> bool: + """Whether a fuzzy entry is the one the router matcher returned for a group, + so its targets never get expanded a second time.""" + if isinstance(entry, dict) and entry: + entry_dict: Final = cast("dict[object, object]", entry) # cast-ok: fallback entries are plain dicts at runtime + value: Final = entry_dict[next(iter(entry_dict))] + return value is resolved or value == resolved + if isinstance(entry, str): + return resolved == [entry] + return False + + +def _fallback_edges( + indexed_lists: Sequence[_FallbackIndex], + model_group: str, + fired_fuzzy: set[tuple[int, int]], # mutable-ok: records which fuzzy entries already expanded +) -> frozenset[str]: + """Targets routing would pick for ``model_group``: exact-map hits on the group + and its provider-stripped suffix, plus the router matcher run over the fuzzy + entries that have not fired yet (tracked in ``fired_fuzzy`` by + ``(list_index, entry_index)``), so each entry's targets are expanded once.""" + from litellm.router_utils.fallback_event_handlers import get_fallback_model_group + + stripped: Final = model_group.split("/", 1)[1] if "/" in model_group else None + targets: Final[set[str]] = set() # mutable-ok: accumulates fallback targets + for list_index, (exact, fuzzy) in enumerate(indexed_lists): + targets.update(exact.get(model_group, ())) + if stripped is not None: + targets.update(exact.get(stripped, ())) + pending = tuple(index for index in range(len(fuzzy)) if (list_index, index) not in fired_fuzzy) + if not pending: + continue + resolved = get_fallback_model_group(fallbacks=[fuzzy[index] for index in pending], model_group=model_group)[0] + values = resolved if isinstance(resolved, list) else ([resolved] if resolved is not None else []) + targets.update(name for value in values if (name := _fallback_target_name(value)) is not None) + if resolved is not None: + fired_fuzzy.update((list_index, index) for index in pending if _fuzzy_entry_matched(fuzzy[index], resolved)) + return frozenset(targets) + + +def _fallback_target_name(value: object) -> str | None: + """The group a single fallback target names: a bare string, or the + ``{"model": "..."}`` advanced entry the router accepts.""" + if isinstance(value, str): + return value + if isinstance(value, dict): + model_value: Final = cast("dict[str, object]", value).get( # cast-ok: fallback dict entries are str-keyed + "model" + ) + return model_value if isinstance(model_value, str) else None + return None + + +def _fallback_target_groups( + llm_router: litellm.Router, + data: Mapping[str, object], + router_settings: Mapping[str, object] | None, + team_router_settings: Mapping[str, object] | None, + model_group: str, +) -> frozenset[str]: + """Transitive closure over every fallback graph that can route this request: + a chain like A -> B -> C must still surface C's per-user credential. Every + entry the router could match is indexed or matcher-scored, so discovery is + a faithful superset of what routing can reach; the only bound is the 400 on + caller-supplied fallback lists in the resolver. A fuzzy entry whose targets + already expanded once is marked fired in ``fired_fuzzy`` and never matched + or expanded again, so total target visits stay O(total targets).""" + indexed: Final = tuple( + _index_fallback_list(fallback_list) + for fallback_list in _fallback_lists(data, llm_router, router_settings, team_router_settings) + ) + seen: Final[set[str]] = set() # mutable-ok: BFS visited set + frontier: Final[list[str]] = [model_group] # mutable-ok: BFS work list + fired_fuzzy: Final[set[tuple[int, int]]] = set() # mutable-ok: fired (list, entry) indexes + while frontier: + group = frontier.pop() + for target in _fallback_edges(indexed, group, fired_fuzzy): + if target not in seen and target != model_group: + seen.add(target) + frontier.append(target) + return frozenset(seen) + + +def _per_user_credential_names_for_groups( + llm_router: litellm.Router, + model_groups: frozenset[str], + team_id: str | None, +) -> tuple[str, ...]: + from litellm.constants import ( + GITHUB_COPILOT_AUTH_TYPE_KEY, + GITHUB_COPILOT_PER_USER_AUTH_TYPE, + ) + from litellm.router_utils.auto_router_model_naming import classify_strategy_router_model + + names: Final[list[str]] = [] # mutable-ok: accumulates one name per per-user deployment + for group in model_groups: + by_name = llm_router.get_model_list(model_name=group, team_id=team_id) or () + deployment_id_match = llm_router.get_deployment(model_id=group) + deployments: Sequence[object] = ( + *by_name, + *((deployment_id_match,) if deployment_id_match is not None else ()), + ) + for deployment in deployments: + dep_map = ( + cast("Mapping[str, object]", deployment) # cast-ok: deployments are str-keyed dicts or objects + if isinstance(deployment, Mapping) + else None + ) + litellm_params: object = ( + dep_map.get("litellm_params") if dep_map is not None else getattr(deployment, "litellm_params", None) + ) + lp_map = ( + cast("Mapping[str, object]", litellm_params) # cast-ok: deployment litellm_params is a str-keyed dict + if isinstance(litellm_params, Mapping) + else None + ) + deployment_model: object = ( + lp_map.get("model") + if lp_map is not None + else getattr( + litellm_params, + "model", + None, + ) + ) + if isinstance(deployment_model, str) and classify_strategy_router_model(deployment_model) is not None: + return _all_per_user_credential_names() + credential_name_obj: object = ( + lp_map.get("litellm_credential_name") + if lp_map is not None + else getattr( + litellm_params, + "litellm_credential_name", + None, + ) + ) + if not isinstance(credential_name_obj, str) or not credential_name_obj or credential_name_obj in names: + continue + credential = CredentialAccessor.find_credential(credential_name_obj) + if credential is None: + continue + values = cast( # cast-ok: credential_values is a plain dict at runtime + Mapping[object, object], credential.credential_values + ) + if values.get(GITHUB_COPILOT_AUTH_TYPE_KEY) == GITHUB_COPILOT_PER_USER_AUTH_TYPE: + names.append(credential_name_obj) + return tuple(names) + + +async def _team_router_settings(user_api_key_dict_team_id: str | None) -> Mapping[str, object] | None: + """The team's router_settings, which _configured_fallbacks prefers after the + key's (hierarchical Key > Team). Mirrors the cached get_team_object lookup the + proxy uses so no extra DB read lands on the hot path.""" + if not user_api_key_dict_team_id: + return None + try: + from litellm.proxy.auth.auth_checks import get_team_object + from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache + + if prisma_client is None: + return None + team_obj: Final = await get_team_object( + team_id=user_api_key_dict_team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + settings: Final = getattr(team_obj, "router_settings", None) + return ( + cast("Mapping[str, object]", settings) # cast-ok: team router_settings is a plain config dict + if isinstance(settings, Mapping) + else None + ) + except Exception: # noqa: BLE001 # advisory discovery must never break the request + return None + + +async def _resolve_user_provider_credentials_for_request( + data: dict[str, object], # mutable-ok: writes resolved credentials into the nested secret_fields dict + authenticated_user_id: str | None, + team_id: str | None, + llm_router: litellm.Router | None, + router_settings: Mapping[str, object] | None = None, +) -> None: + """Resolve the calling user's per-user provider connections into secret_fields. + + Runs after model alias rewrites so ``data['model']`` is the final requested + model. ``authenticated_user_id`` is the identity established by the validated + API key or JWT, never a header-mapped value, so a caller cannot load another + user's stored token. Never raises for "not connected": shared deployments in + the same group still serve the request, so the provider-side helper is the + one that fails the call with a 401 when per-user mode applies.""" + user_id: Final = authenticated_user_id + model: Final = data.get("model") + if llm_router is None or not isinstance(model, str): + return + # With no per-user credential configured there is nothing to discover: no + # traversal, no size check, identical behavior to a shared-mode-only proxy. + if not _per_user_oauth_configured(): + return + # The one caller-controlled input to discovery is the request body's own + # fallback lists, so the size bound lives here instead of on the scan: admin + # lists (router, key/team router_settings) are always read in full. The + # bound is on aggregate target names, not outer entries, so a single dict + # holding thousands of targets is counted honestly. + request_fallback_targets: Final = sum( + sum( + _fallback_entry_target_count(entry) + for entry in cast("Sequence[object]", entries) # cast-ok: request fallback lists hold mixed entry shapes + ) + for key in _REQUEST_FALLBACK_KEYS + if isinstance(entries := data.get(key), list) + ) + if request_fallback_targets > _FALLBACK_DISCOVERY_LIMIT: + raise HTTPException( + status_code=400, + detail=(f"fallback lists in the request body cannot exceed {_FALLBACK_DISCOVERY_LIMIT} entries"), + ) + from litellm.router_utils.common_utils import resolve_model_group_alias + + # Discovery must see the group the router actually routes to, so include the + # post-alias name from every alias map that may still rewrite data["model"] + # after this function (litellm.model_alias_map, router_settings.model_group_alias). + requested_groups: Final = frozenset( + target + for target in ( + model, + litellm.model_alias_map.get(model), + ( + resolve_model_group_alias(router_settings.get("model_group_alias"), model) + if router_settings is not None + else None + ), + ) + if isinstance(target, str) and target + ) + team_router_settings: Final = await _team_router_settings(user_api_key_dict_team_id=team_id) + model_groups: Final = frozenset( + itertools.chain.from_iterable( + (requested_groups,) + + tuple( + _fallback_target_groups(llm_router, data, router_settings, team_router_settings, group) + for group in requested_groups + ) + ) + ) + credential_names: Final = _per_user_credential_names_for_groups(llm_router, model_groups, team_id) + if credential_names and "litellm_credential_name" in data: + raise HTTPException( + status_code=400, + detail="litellm_credential_name cannot be set in the request body for a model that uses per-user GitHub OAuth", + ) + if not isinstance(user_id, str) or not user_id or not credential_names: + return + + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + from litellm.types.proxy.litellm_pre_call_utils import RedactedDict + + from .credential_endpoints.user_provider_credentials import aget_user_provider_tokens + + if prisma_client is None: + return + try: + tokens: Final = await aget_user_provider_tokens( + prisma_client=prisma_client, + cache=user_api_key_cache, + user_id=user_id, + credential_names=credential_names, + ) + except Exception: # noqa: BLE001 # credential lookup failures degrade to the shared deployment, never break the request + verbose_proxy_logger.exception( + "_resolve_user_provider_credentials_for_request: failed to load user provider credentials" + ) + return + + secret_fields: Final = data.get("secret_fields") + if not isinstance(secret_fields, dict): + return + secret_fields["user_provider_credentials"] = RedactedDict(dict(tokens)) + secret_fields["user_provider_credentials_user_id"] = user_id + + def _warn_stale_team_alias_once(warning_key: str, message: str, *args: str) -> None: if warning_key in _STALE_TEAM_ALIAS_WARNING_KEYS: return diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4eab87cc84e..d8eeaaa2237 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -9364,7 +9364,10 @@ class ProxyConfig: decrypted_credential_values: Final = {} for k, v in credential_object.credential_values.items(): - decrypted_credential_values[k] = decrypted_or_stored(k, v) + decrypted_credential_values[k] = decrypted_or_stored( + k, + cast("str", v), # cast-ok: credential values are str at the decrypt boundary + ) credential_object.credential_values = decrypted_credential_values return credential_object diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 35d026bb6f6..f3ebee52dcd 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -453,6 +453,20 @@ model LiteLLM_MCPUserCredentials { @@unique([user_id, server_id]) } +// Per-user provider connections (e.g. GitHub Copilot OAuth) keyed by credential name +model LiteLLM_UserProviderCredentials { + id String @id @default(uuid()) + user_id String + credential_name String + provider String + credential_b64 String + created_at DateTime @default(now()) @map("created_at") + updated_at DateTime @default(now()) @updatedAt @map("updated_at") + + @@unique([user_id, credential_name]) + @@index([credential_name]) +} + // Per-user environment variable values for MCP servers. // values_b64 is an encrypted JSON object: {VAR_NAME: "value", ...}. model LiteLLM_MCPUserEnvVars { diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 7298f4b1a36..f6bd4bfba54 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -109,8 +109,8 @@ def _has_file_search_tool(tools: Iterable[Mapping[str, object]] | None) -> bool: def mock_responses_api_response( mock_response: str = "In a peaceful grove beneath a silver moon, a unicorn named Lumina discovered a hidden pool that reflected the stars. As she dipped her horn into the water, the pool began to shimmer, revealing a pathway to a magical realm of endless night skies. Filled with wonder, Lumina whispered a wish for all who dream to find their own hidden magic, and as she glanced back, her hoofprints sparkled like stardust.", ): - return ResponsesAPIResponse( - **{ + return ResponsesAPIResponse.model_validate( + { "id": "resp_67ccd2bed1ec8190b14f964abc0542670bb6a6b452d3795b", "object": "response", "created_at": 1741476542, @@ -668,7 +668,11 @@ async def aresponses( # get custom llm provider so we can use this for mapping exceptions if custom_llm_provider is None: _, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=model, api_base=local_vars.get("base_url", None) + model=model, + api_base=local_vars.get("base_url", None), + litellm_params=GenericLiteLLMParams( + **cast("dict[str, object]", kwargs) # cast-ok: kwargs is the untyped request dict + ), ) # Update local_vars with detected provider (fixes #19782) local_vars["custom_llm_provider"] = custom_llm_provider @@ -792,7 +796,7 @@ async def aresponses( # (mirrors litellm/main.py:1371 for chat completions) response.hidden_params["custom_llm_provider"] = custom_llm_provider - if response is None: + if response is None: # pyright: ignore[reportUnnecessaryComparison] # provider handlers can return None at runtime raise ValueError(f"Got an unexpected None response from the Responses API: {response}") return response @@ -1054,7 +1058,6 @@ def _responses_try_dispatch_mcp_gateway( if skip_mcp_handler or not LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway(tools=tools): return None mcp_call_kwargs: Final = { - "input": input, "model": model, "include": include, "instructions": instructions, @@ -1063,13 +1066,11 @@ def _responses_try_dispatch_mcp_gateway( "metadata": metadata, "parallel_tool_calls": parallel_tool_calls, "previous_response_id": previous_response_id, - "reasoning": reasoning, "store": store, "background": background, "stream": stream, "temperature": temperature, "text": text, - "tool_choice": tool_choice, "tools": tools, "top_p": top_p, "truncation": truncation, @@ -1077,13 +1078,25 @@ def _responses_try_dispatch_mcp_gateway( "extra_headers": extra_headers, "extra_query": extra_query, "extra_body": extra_body, - "timeout": timeout, "custom_llm_provider": custom_llm_provider, **kwargs, } if _is_async: - return aresponses_api_with_mcp(**mcp_call_kwargs) - return run_async_function(aresponses_api_with_mcp, **mcp_call_kwargs) + return aresponses_api_with_mcp( + input=input, + reasoning=reasoning, + timeout=timeout, + tool_choice=tool_choice, + **mcp_call_kwargs, + ) + return run_async_function( + aresponses_api_with_mcp, + input=input, + reasoning=reasoning, + timeout=timeout, + tool_choice=tool_choice, + **mcp_call_kwargs, + ) def _responses_try_dispatch_emulated_file_search( @@ -1251,7 +1264,11 @@ def responses( if custom_llm_provider is None: _, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=model, api_base=local_vars.get("base_url", None) + model=model, + api_base=local_vars.get("base_url", None), + litellm_params=GenericLiteLLMParams( + **cast("dict[str, object]", kwargs) # cast-ok: kwargs is the untyped request dict + ), ) local_vars["custom_llm_provider"] = custom_llm_provider @@ -1719,13 +1736,15 @@ async def aget_responses( response = init_response # Update the responses_api_response_id with the model_id - if isinstance(response, ResponsesAPIResponse): - response = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( - responses_api_response=response, - litellm_metadata=kwargs.get("litellm_metadata", {}), - custom_llm_provider=custom_llm_provider, - ) - return response + if not isinstance(response, ResponsesAPIResponse): # pyright: ignore[reportUnnecessaryIsInstance] # handlers can return non-ResponsesAPIResponse objects at runtime + return response + return ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( + responses_api_response=response, + litellm_metadata=cast( # cast-ok: litellm_metadata is a plain dict when present + "dict[str, object]", kwargs.get("litellm_metadata", {}) + ), + custom_llm_provider=custom_llm_provider, + ) except Exception as e: raise litellm.exception_type( model=None, @@ -2159,7 +2178,11 @@ async def acompact_responses( # get custom llm provider so we can use this for mapping exceptions if custom_llm_provider is None: _, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=model, api_base=local_vars.get("base_url", None) + model=model, + api_base=local_vars.get("base_url", None), + litellm_params=GenericLiteLLMParams( + **cast("dict[str, object]", kwargs) # cast-ok: kwargs is the untyped request dict + ), ) # Update local_vars with detected provider (fixes #19782) local_vars["custom_llm_provider"] = custom_llm_provider @@ -2188,14 +2211,15 @@ async def acompact_responses( response = init_response # Update the responses_api_response_id with the model_id - if isinstance(response, ResponsesAPIResponse): - response = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( - responses_api_response=response, - litellm_metadata=kwargs.get("litellm_metadata", {}), - custom_llm_provider=custom_llm_provider, - ) - - return response + if not isinstance(response, ResponsesAPIResponse): # pyright: ignore[reportUnnecessaryIsInstance] # handlers can return non-ResponsesAPIResponse objects at runtime + return response + return ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( + responses_api_response=response, + litellm_metadata=cast( # cast-ok: litellm_metadata is a plain dict when present + "dict[str, object]", kwargs.get("litellm_metadata", {}) + ), + custom_llm_provider=custom_llm_provider, + ) except Exception as e: raise litellm.exception_type( model=model, @@ -2423,6 +2447,7 @@ async def _aresponses_websocket( model=model, api_base=api_base, api_key=api_key, + litellm_params=litellm_params, ) resolved_model: Final = _strip_responses_routing_prefix(provider_model) diff --git a/litellm/router.py b/litellm/router.py index 15b8282b54d..f173cdb37b4 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4225,6 +4225,29 @@ class Router: existing_tags.append(credential_tag) kwargs[metadata_variable_name]["tags"] = existing_tags + from litellm.llms.github_copilot.per_user_auth import ( + github_copilot_per_user_credential_name, + ) + + configured_per_user_name: Final = github_copilot_per_user_credential_name( + cast( # cast-ok: deployment litellm_params is a str-keyed dict + "Mapping[str, object]", deployment.get("litellm_params") + ) + ) + caller_credential_name: Final[object] = cast( # cast-ok: kwargs is the untyped request dict + "object", kwargs.get("litellm_credential_name") + ) + if ( + configured_per_user_name is not None + and isinstance(caller_credential_name, str) + and caller_credential_name != configured_per_user_name + ): + raise litellm.BadRequestError( + message="litellm_credential_name cannot be overridden on a deployment that uses per-user GitHub OAuth", + model=deployment_model_name, + llm_provider="", + ) + kwargs["model_info"] = model_info if function_name == "_ageneric_api_call_with_fallbacks": @@ -8663,7 +8686,13 @@ class Router: original_exception=exception, deployment=deployment_id, time_to_cooldown=_time_to_cooldown, - requested_model_group=(get_litellm_metadata_from_kwargs(kwargs) or {}).get("model_group"), + requested_model_group=( + get_litellm_metadata_from_kwargs( + cast("dict[str, object]", kwargs) # cast-ok: untyped kwargs dict + ) + or {} + ).get("model_group"), + request_kwargs=cast("dict[str, object]", kwargs), # cast-ok: kwargs is the untyped request dict ) # setting deployment_id in cooldown deployments return result @@ -10005,6 +10034,7 @@ class Router: model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.get("custom_llm_provider", None), api_base=deployment.litellm_params.api_base, + litellm_params=deployment.litellm_params, ) # done reading model["litellm_params"] # Check if provider is supported: either in enum or JSON-configured @@ -10116,10 +10146,21 @@ class Router: credential_values: Final = ( CredentialAccessor.get_credential_values(credential_name) if credential_name is not None else {} ) - vertex_project: Final = credential_values.get("vertex_project") or deployment.litellm_params.vertex_project - vertex_location: Final = credential_values.get("vertex_location") or deployment.litellm_params.vertex_location + from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES + + vertex_project: Final = ( + cast("str | None", credential_values.get("vertex_project")) # cast-ok: vertex credential values are str + or deployment.litellm_params.vertex_project + ) + vertex_location: Final = ( + cast("str | None", credential_values.get("vertex_location")) # cast-ok: vertex credential values are str + or deployment.litellm_params.vertex_location + ) vertex_credentials: Final = ( - credential_values.get("vertex_credentials") or deployment.litellm_params.vertex_credentials + cast( # cast-ok: vertex_credentials holds the typed credential union + "VERTEX_CREDENTIALS_TYPES | None", credential_values.get("vertex_credentials") + ) + or deployment.litellm_params.vertex_credentials ) if vertex_project is None or vertex_location is None: @@ -11415,6 +11456,7 @@ class Router: litellm_model, llm_provider, _, _ = litellm.get_llm_provider( model=litellm_params.model, custom_llm_provider=litellm_params.custom_llm_provider, + litellm_params=litellm_params, ) except litellm.exceptions.BadRequestError as e: verbose_router_logger.error("litellm.router.py::get_model_group_info() - %s", e) @@ -12390,6 +12432,7 @@ class Router: model, custom_llm_provider, _, _ = litellm.get_llm_provider( model=deployment_params.get("model") or group, custom_llm_provider=deployment_params.get("custom_llm_provider"), + litellm_params=LiteLLM_Params.model_validate(deployment_params), ) supported: Final = litellm.get_supported_openai_params( model=model, diff --git a/litellm/router_utils/clientside_credential_handler.py b/litellm/router_utils/clientside_credential_handler.py index e65bd777eca..65980449937 100644 --- a/litellm/router_utils/clientside_credential_handler.py +++ b/litellm/router_utils/clientside_credential_handler.py @@ -11,7 +11,8 @@ If given, generate a unique model_id for the deployment. Ensures cooldowns are applied correctly. """ -from typing import Final +from collections.abc import MutableMapping +from typing import Final, cast # noqa: TID251 # narrows the untyped request dict for mutation from litellm.types.utils import server_owned_wif_litellm_params @@ -112,8 +113,29 @@ def get_dynamic_litellm_params(litellm_params: dict, request_kwargs: dict) -> di # admin's value in ``litellm_params`` and have it forwarded to the # redirected upstream. if "api_base" in request_kwargs or "base_url" in request_kwargs: + from litellm.llms.github_copilot.per_user_auth import github_copilot_auth_mode + + # A per-user OAuth credential must survive the clear: dropping + # litellm_credential_name / github_copilot_auth_type would flip the call to + # shared mode and send the admin's Copilot token to the redirected host. + # Per-user mode ignores the caller's api_base anyway (the session's + # validated host wins), so keeping the mode is fail-closed. + credential_name_value: Final[object] = litellm_params.get( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # litellm_params is the untyped request dict + "litellm_credential_name" + ) + auth_type: Final[object] = litellm_params.get( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # litellm_params is the untyped request dict + "github_copilot_auth_type" + ) + per_user_credential_locked: Final = github_copilot_auth_mode( + credential_name_value, # pyright: ignore[reportUnknownArgumentType] # value comes from the untyped request dict + auth_type, # pyright: ignore[reportUnknownArgumentType] # value comes from the untyped request dict + ) for field in _ADMIN_CONFIG_FIELDS_TO_CLEAR_ON_BASE_OVERRIDE: - litellm_params.pop(field, None) + if per_user_credential_locked and field in ("litellm_credential_name", "github_copilot_auth_type"): + continue + cast( # cast-ok: litellm_params is a mutable request dict at runtime + "MutableMapping[str, object]", litellm_params + ).pop(field, None) if field in request_kwargs: litellm_params[field] = request_kwargs[field] litellm_params[DISABLE_WORKLOAD_IDENTITY_PARAM] = True diff --git a/litellm/router_utils/cooldown_handlers.py b/litellm/router_utils/cooldown_handlers.py index 84693589873..a5f6277d2d7 100644 --- a/litellm/router_utils/cooldown_handlers.py +++ b/litellm/router_utils/cooldown_handlers.py @@ -271,12 +271,25 @@ def _is_cooldown_required( return True +def _is_per_user_provider_request(request_kwargs: Mapping[str, object]) -> bool: + from litellm.llms.github_copilot.per_user_auth import ( + github_copilot_user_session_from, + is_github_copilot_per_user_request, + ) + + return ( + is_github_copilot_per_user_request(request_kwargs) + or github_copilot_user_session_from(request_kwargs.get("litellm_params")) is not None + ) + + def _should_run_cooldown_logic( litellm_router_instance: LitellmRouter, deployment: str | None, exception_status: str | int, original_exception: Exception, time_to_cooldown: float | None = None, + request_kwargs: Mapping[str, object] | None = None, ) -> bool: """ Helper that decides if cooldown logic should be run @@ -295,6 +308,38 @@ def _should_run_cooldown_logic( ) return False + if isinstance( + original_exception, + ( + litellm.CallerCredentialAuthenticationError, + litellm.CallerCredentialRateLimitError, + ), + ): + verbose_router_logger.debug( + "Should Not Run Cooldown Logic: caller-credential errors are scoped to one user's connection" + ) + return False + + if request_kwargs is not None and _is_per_user_provider_request(request_kwargs): + verbose_router_logger.debug( + "Should Not Run Cooldown Logic: the failure came from one caller's per-user provider session" + ) + return False + + if deployment is not None: + failed_deployment: Final = litellm_router_instance.get_deployment(model_id=deployment) + if failed_deployment is not None: + from litellm.llms.github_copilot.per_user_auth import ( + github_copilot_per_user_credential_name, + ) + + failed_litellm_params: Final = getattr(failed_deployment, "litellm_params", None) + if github_copilot_per_user_credential_name(failed_litellm_params) is not None: + verbose_router_logger.debug( + "Should Not Run Cooldown Logic: the failed deployment authenticates per caller" + ) + return False + ######################################################### # If time_to_cooldown is 0 or 0.0000000, don't run cooldown logic ######################################################### @@ -442,6 +487,7 @@ def set_cooldown_deployments( deployment: str | None = None, time_to_cooldown: float | None = None, requested_model_group: str | None = None, + request_kwargs: Mapping[str, object] | None = None, ) -> bool: """ Add a model to the list of models being cooled down for that minute, if it exceeds the allowed fails / minute @@ -463,6 +509,7 @@ def set_cooldown_deployments( exception_status=exception_status, original_exception=original_exception, time_to_cooldown=time_to_cooldown, + request_kwargs=request_kwargs, ) is False or deployment is None diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index d95774d4d85..e1de57c91a0 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -156,6 +156,7 @@ def _trigger_cooldown_for_failed_deployment( original_exception=exception, deployment=deployment_id, time_to_cooldown=time_to_cooldown, + request_kwargs=kwargs, ) verbose_router_logger.debug("Triggered cooldown for fallback deployment %s", deployment_id) diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index 23162c8c156..ff31d7a65c4 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -87,6 +87,8 @@ class ProviderConnection: litellm_credential_name: str | None = None configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None use_xai_oauth: bool | None = None + github_copilot_auth_type: str | None = None + github_copilot_user_session: object | None = None token_exchange_endpoint: str | None = None token_exchange_profile: str | None = None token_exchange_scope: str | None = None diff --git a/litellm/types/proxy/litellm_pre_call_utils.py b/litellm/types/proxy/litellm_pre_call_utils.py index e0d3f3dac66..0b980d38320 100644 --- a/litellm/types/proxy/litellm_pre_call_utils.py +++ b/litellm/types/proxy/litellm_pre_call_utils.py @@ -1,4 +1,6 @@ -from typing_extensions import TypedDict +from collections.abc import Mapping + +from typing_extensions import NotRequired, ReadOnly, TypedDict class RedactedDict(dict): @@ -22,3 +24,5 @@ class SecretFields(TypedDict): """ raw_headers: dict + user_provider_credentials: NotRequired[ReadOnly[Mapping[str, str]]] + user_provider_credentials_user_id: NotRequired[ReadOnly[str]] diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 01ebf139a88..4a962e5b3bd 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -4060,8 +4060,12 @@ oauth_token_exchange_litellm_params: Final = ( "token_exchange_profile", "token_exchange_scope", ) +github_copilot_oauth_litellm_params: Final = ("github_copilot_auth_type",) server_owned_wif_litellm_params: Final = ( - anthropic_wif_litellm_params + openai_wif_litellm_params + oauth_token_exchange_litellm_params + anthropic_wif_litellm_params + + openai_wif_litellm_params + + oauth_token_exchange_litellm_params + + github_copilot_oauth_litellm_params ) secret_bearing_wif_litellm_params: Final = tuple(sorted(WIF_SECRET_BEARING_KEYS)) diff --git a/litellm/utils.py b/litellm/utils.py index b67a7294c6d..25e0b3aba8c 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1740,6 +1740,9 @@ def client(original_function): ## LOAD CREDENTIALS load_credentials_from_list(kwargs) + from litellm.llms.github_copilot.per_user_auth import attach_github_copilot_user_session + + attach_github_copilot_user_session(kwargs) # pyright: ignore[reportUnknownArgumentType] # kwargs is the untyped request dict kwargs["litellm_logging_obj"] = logging_obj LLMCachingHandler: Final = _get_cached_llm_caching_handler() _llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler( @@ -2028,6 +2031,9 @@ def client(original_function): kwargs["litellm_logging_obj"] = logging_obj ## LOAD CREDENTIALS load_credentials_from_list(kwargs) + from litellm.llms.github_copilot.per_user_auth import aattach_github_copilot_user_session + + await aattach_github_copilot_user_session(kwargs) # pyright: ignore[reportUnknownArgumentType] # kwargs is the untyped request dict logging_obj.llm_caching_handler = _llm_caching_handler # [OPTIONAL] CHECK BUDGET if litellm.max_budget: @@ -8301,13 +8307,18 @@ def convert_to_dict(message: BaseModel | dict) -> dict: raise TypeError(f"Invalid message type: {type(message)}. Expected dict or Pydantic model.") -def convert_list_message_to_dict(messages: Sequence): - new_messages: Final = [] - for message in messages: - convert_msg_to_dict = cast(AllMessageValues, convert_to_dict(message)) - cleaned_message = cleanup_none_field_in_message(message=convert_msg_to_dict) - new_messages.append(cleaned_message) - return new_messages +def convert_list_message_to_dict( + messages: Sequence[BaseModel | Mapping[str, object]], +) -> list[dict[str, object]]: # mutable-ok: callers mutate the returned message dicts + def _as_message_value(message: BaseModel | Mapping[str, object]) -> AllMessageValues: + return cast( # cast-ok: message dicts satisfy the TypedDict shape + AllMessageValues, + convert_to_dict( + cast("BaseModel | dict[str, object]", message) # cast-ok: messages are dicts or pydantic models at runt + ), + ) + + return [dict(cleanup_none_field_in_message(message=_as_message_value(message))) for message in messages] def validate_and_fix_openai_messages(messages: list): @@ -8323,7 +8334,12 @@ def validate_and_fix_openai_messages(messages: list): if message.get("tool_calls"): message["tool_calls"] = jsonify_tools(tools=message["tool_calls"]) - convert_msg_to_dict = cast(AllMessageValues, convert_to_dict(message)) + convert_msg_to_dict = cast( # cast-ok: message dicts satisfy the TypedDict shape + AllMessageValues, + convert_to_dict( + cast("BaseModel | dict[str, object]", message) # cast-ok: dicts or pydantic models at runtime + ), + ) cleaned_message = cleanup_none_field_in_message(message=convert_msg_to_dict) new_messages.append(cleaned_message) return validate_chat_completion_user_messages(messages=new_messages) diff --git a/schema.prisma b/schema.prisma index 35d026bb6f6..f3ebee52dcd 100644 --- a/schema.prisma +++ b/schema.prisma @@ -453,6 +453,20 @@ model LiteLLM_MCPUserCredentials { @@unique([user_id, server_id]) } +// Per-user provider connections (e.g. GitHub Copilot OAuth) keyed by credential name +model LiteLLM_UserProviderCredentials { + id String @id @default(uuid()) + user_id String + credential_name String + provider String + credential_b64 String + created_at DateTime @default(now()) @map("created_at") + updated_at DateTime @default(now()) @updatedAt @map("updated_at") + + @@unique([user_id, credential_name]) + @@index([credential_name]) +} + // Per-user environment variable values for MCP servers. // values_b64 is an encrypted JSON object: {VAR_NAME: "value", ...}. model LiteLLM_MCPUserEnvVars { diff --git a/terraform/provider/tools/endpointaudit/coverage_allowlist.txt b/terraform/provider/tools/endpointaudit/coverage_allowlist.txt index 9305d2e109d..9430ada8468 100644 --- a/terraform/provider/tools/endpointaudit/coverage_allowlist.txt +++ b/terraform/provider/tools/endpointaudit/coverage_allowlist.txt @@ -129,6 +129,13 @@ PATCH /team/{team_id} POST /team/model/add POST /team/model/delete +# Per-user OAuth connections managed by the calling user through the dashboard +# device flow; caller-scoped tokens, not Terraform-managed state +GET /credentials/user_connections +POST /credentials/{credential_name}/user_connection/start +POST /credentials/{credential_name}/user_connection/poll +DELETE /credentials/{credential_name}/user_connection + # Known gaps awaiting a resource or data source; remove the entry when it lands GET /credentials # known gap: plural credentials data source GET /cache/settings # known gap: cache settings resource diff --git a/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py b/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py index 83c0cb2fec2..88fc8ba6e6a 100644 --- a/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py +++ b/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py @@ -76,9 +76,7 @@ def test_github_copilot_embedding_config_get_complete_url(): assert url == "https://api.githubcopilot.com/embeddings" # Test with custom API base from authenticator - config.authenticator.get_api_base.return_value = ( - "https://api.enterprise.githubcopilot.com" - ) + config.authenticator.get_api_base.return_value = "https://api.enterprise.githubcopilot.com" url = config.get_complete_url( api_base=None, api_key=None, @@ -258,3 +256,56 @@ def test_transform_embedding_response_rejects_a_body_that_is_not_an_object(body: def test_transform_embedding_response_object_without_data_is_an_invalid_response_object(): with pytest.raises(Exception, match="Invalid response object"): _transform(httpx.Response(200, json={"model": "text-embedding-3-small"})) + + +def test_validate_environment_uses_per_user_session_and_skips_authenticator(): + """With a per-user session the Authorization header carries the caller's Copilot token + and the shared Authenticator is never consulted.""" + from litellm.llms.github_copilot.per_user_auth import GithubCopilotUserSession + + config = GithubCopilotEmbeddingConfig() + config.authenticator = MagicMock() + config.authenticator.get_api_key.side_effect = AssertionError("shared authenticator must not run") + + session = GithubCopilotUserSession(token="user-copilot-token", api_base="https://api.githubcopilot.com") + headers = config.validate_environment( + headers={}, + model="github_copilot/text-embedding-3-small", + messages=[], + optional_params={}, + litellm_params={"github_copilot_user_session": session}, + ) + assert headers["Authorization"] == "Bearer user-copilot-token" + config.authenticator.get_api_key.assert_not_called() + + +def test_per_user_session_token_wins_over_caller_authorization(): + from litellm.llms.github_copilot.per_user_auth import GithubCopilotUserSession + + config = GithubCopilotEmbeddingConfig() + config.authenticator = MagicMock() + session = GithubCopilotUserSession(token="user-copilot-token", api_base="https://api.githubcopilot.com") + headers = config.validate_environment( + headers={"Authorization": "Bearer caller-token"}, + model="github_copilot/text-embedding-3-small", + messages=[], + optional_params={}, + litellm_params={"github_copilot_user_session": session}, + ) + assert headers["Authorization"] == "Bearer user-copilot-token" + + +def test_get_complete_url_prefers_per_user_session_api_base(): + from litellm.llms.github_copilot.per_user_auth import GithubCopilotUserSession + + config = GithubCopilotEmbeddingConfig() + config.authenticator = MagicMock() + session = GithubCopilotUserSession(token="t", api_base="https://tenant.githubcopilot.com") + url = config.get_complete_url( + api_base="https://attacker.example", + api_key=None, + model="m", + optional_params={}, + litellm_params={"github_copilot_user_session": session}, + ) + assert url == "https://tenant.githubcopilot.com/embeddings" diff --git a/tests/unit/llms/github_copilot/messages/test_github_copilot_messages_transformation.py b/tests/unit/llms/github_copilot/messages/test_github_copilot_messages_transformation.py index 9d8673ed3d8..b029e5aba0b 100644 --- a/tests/unit/llms/github_copilot/messages/test_github_copilot_messages_transformation.py +++ b/tests/unit/llms/github_copilot/messages/test_github_copilot_messages_transformation.py @@ -1,5 +1,6 @@ from unittest.mock import MagicMock +import httpx import pytest @@ -322,3 +323,64 @@ def test_github_copilot_messages_config_probes_capabilities_under_copilot_namesp ``anthropic`` namespace and ignored the exact ``github_copilot/claude-*`` cost-map entries.""" assert GithubCopilotAnthropicMessagesConfig().custom_llm_provider == "github_copilot" + + +def test_validate_environment_uses_per_user_session_and_skips_authenticator(): + from litellm.llms.github_copilot.per_user_auth import GithubCopilotUserSession + + config = GithubCopilotAnthropicMessagesConfig() + config.authenticator = MagicMock() + config.authenticator.get_api_key.side_effect = AssertionError("shared authenticator must not run") + config.authenticator.get_api_base.side_effect = AssertionError("shared authenticator must not run") + + session = GithubCopilotUserSession(token="user-copilot-token", api_base="https://tenant.githubcopilot.com/") + headers, api_base = config.validate_anthropic_messages_environment( + headers={}, + model="github_copilot/claude-sonnet-4.5", + messages=[], + optional_params={}, + litellm_params={"github_copilot_user_session": session}, + ) + assert headers["Authorization"] == "Bearer user-copilot-token" + assert api_base == "https://tenant.githubcopilot.com" + config.authenticator.get_api_key.assert_not_called() + + +def test_per_user_session_token_wins_over_caller_authorization(): + from litellm.llms.github_copilot.per_user_auth import GithubCopilotUserSession + + config = GithubCopilotAnthropicMessagesConfig() + config.authenticator = MagicMock() + config.authenticator.get_api_base.return_value = "https://api.githubcopilot.com" + session = GithubCopilotUserSession(token="user-copilot-token", api_base="https://api.githubcopilot.com") + headers, _ = config.validate_anthropic_messages_environment( + headers={"authorization": "Bearer caller-token"}, + model="github_copilot/claude-sonnet-4.5", + messages=[], + optional_params={}, + litellm_params={"github_copilot_user_session": session}, + ) + assert headers["Authorization"] == "Bearer user-copilot-token" + + +def test_transform_response_carries_upstream_usage(): + config = GithubCopilotAnthropicMessagesConfig() + raw = httpx.Response( + 200, + json={ + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": "github_copilot/claude-sonnet-4.5", + "content": [{"type": "text", "text": "hi"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 14, "output_tokens": 3}, + }, + ) + result = config.transform_anthropic_messages_response( + model="github_copilot/claude-sonnet-4.5", + raw_response=raw, + logging_obj=MagicMock(), + ) + assert result["usage"]["input_tokens"] == 14 + assert result["usage"]["output_tokens"] == 3 diff --git a/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py b/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py index 4d5990d8d18..12fdc374123 100644 --- a/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py +++ b/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py @@ -9,7 +9,7 @@ Source: litellm/llms/github_copilot/responses/transformation.py from unittest.mock import patch, MagicMock - +import httpx import pytest import litellm from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map @@ -753,3 +753,82 @@ class TestGithubCopilotReasoningStreamItemIdNormalization: }, ) assert event.item_id == "stable_rs_id" + + +def test_validate_environment_uses_per_user_session_and_skips_authenticator(): + from litellm.llms.github_copilot.per_user_auth import GithubCopilotUserSession + + config = GithubCopilotResponsesAPIConfig() + config.authenticator = MagicMock() + config.authenticator.get_api_key.side_effect = AssertionError("shared authenticator must not run") + + session = GithubCopilotUserSession(token="user-copilot-token", api_base="https://api.githubcopilot.com") + headers = config.validate_environment( + headers={}, + model="github_copilot/gpt-5.1", + litellm_params={"github_copilot_user_session": session}, + ) + assert headers["Authorization"] == "Bearer user-copilot-token" + config.authenticator.get_api_key.assert_not_called() + + +def test_per_user_session_token_wins_over_caller_authorization(): + """extra_headers["Authorization"] (any casing) must never displace the per-user + session token on the outgoing request.""" + from litellm.llms.github_copilot.per_user_auth import GithubCopilotUserSession + + config = GithubCopilotResponsesAPIConfig() + config.authenticator = MagicMock() + session = GithubCopilotUserSession(token="user-copilot-token", api_base="https://api.githubcopilot.com") + headers = config.validate_environment( + headers={"Authorization": "Bearer caller-token", "authorization": "Bearer caller-token-lower"}, + model="github_copilot/gpt-5.1", + litellm_params={"github_copilot_user_session": session}, + ) + assert headers["Authorization"] == "Bearer user-copilot-token" + assert "authorization" not in headers + + +def test_shared_mode_keeps_caller_authorization(): + """Pin: without a session, caller extra_headers still override the shared token.""" + config = GithubCopilotResponsesAPIConfig() + config.authenticator = MagicMock() + config.authenticator.get_api_key.return_value = "shared-api-key" + headers = config.validate_environment( + headers={"Authorization": "Bearer caller-token"}, + model="github_copilot/gpt-5.1", + litellm_params={}, + ) + assert headers["Authorization"] == "Bearer caller-token" + + +def test_transform_response_carries_upstream_usage(): + config = GithubCopilotResponsesAPIConfig() + raw = httpx.Response( + 200, + json={ + "id": "resp_1", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "github_copilot/gpt-5.1", + "output": [ + { + "type": "message", + "id": "m1", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "hi", "annotations": []}], + } + ], + "usage": {"input_tokens": 9, "output_tokens": 4, "total_tokens": 13}, + }, + ) + result = config.transform_response_api_response( + model="github_copilot/gpt-5.1", + raw_response=raw, + logging_obj=MagicMock(), + ) + assert result.usage is not None + assert result.usage.input_tokens == 9 + assert result.usage.output_tokens == 4 diff --git a/tests/unit/llms/github_copilot/test_github_copilot_transformation.py b/tests/unit/llms/github_copilot/test_github_copilot_transformation.py index 1715d106e0a..c775beeae24 100644 --- a/tests/unit/llms/github_copilot/test_github_copilot_transformation.py +++ b/tests/unit/llms/github_copilot/test_github_copilot_transformation.py @@ -918,3 +918,51 @@ def test_openai_handler_repairs_github_copilot_empty_choices( assert result.choices[0].message.content == "Hi there" assert result.choices[0].finish_reason == "stop" mock_request.assert_called_once() + + +def test_validate_environment_uses_per_user_session_and_skips_authenticator(): + from litellm.llms.github_copilot.per_user_auth import GithubCopilotUserSession + + config = GithubCopilotConfig() + config.authenticator = MagicMock() + config.authenticator.get_api_key.side_effect = AssertionError("shared authenticator must not run") + + session = GithubCopilotUserSession(token="user-copilot-token", api_base="https://api.githubcopilot.com") + headers = config.validate_environment( + headers={}, model="github_copilot/gpt-4o", messages=[], optional_params={}, + litellm_params={"github_copilot_user_session": session}, + ) + assert headers["Authorization"] == "Bearer user-copilot-token" + config.authenticator.get_api_key.assert_not_called() + + +def test_transform_response_carries_upstream_usage(): + """Usage from the upstream Copilot payload must survive so TPM limits and spend work.""" + config = GithubCopilotConfig() + raw_response = httpx.Response( + 200, + json={ + "id": "chatcmpl-u", + "object": "chat.completion", + "created": 1, + "model": "github_copilot/gpt-4o", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 11, "completion_tokens": 7, "total_tokens": 18}, + }, + ) + result = config.transform_response( + model="github_copilot/gpt-4o", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=MagicMock(model_call_details={}), + request_data={}, + messages=[], + optional_params={}, + litellm_params={}, + encoding=None, + ) + assert result.usage.prompt_tokens == 11 + assert result.usage.completion_tokens == 7 + assert result.usage.total_tokens == 18 diff --git a/tests/unit/llms/github_copilot/test_per_user_auth.py b/tests/unit/llms/github_copilot/test_per_user_auth.py new file mode 100644 index 00000000000..9b16d3c0aa7 --- /dev/null +++ b/tests/unit/llms/github_copilot/test_per_user_auth.py @@ -0,0 +1,1084 @@ +"""Unit tests for per-user GitHub Copilot OAuth connections (per_user_auth).""" + +import asyncio +import datetime +import time +from collections.abc import Callable +from contextlib import contextmanager +from typing import Final +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +import litellm +from litellm.exceptions import ( + AuthenticationError, + BadRequestError, + CallerCredentialAuthenticationError, + CallerCredentialRateLimitError, + ServiceUnavailableError, +) +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.github_copilot.per_user_auth import ( + _SESSION_CACHE, + GITHUB_COPILOT_USER_SESSION_KWARG_KEY, + GithubCopilotUserSession, + _session_cache_key, + aattach_github_copilot_user_session, + aexchange_github_token, + apoll_device_flow, + astart_device_flow, + attach_github_copilot_user_session, + evict_copilot_user_session, + exchange_github_token, + github_copilot_auth_mode, + github_copilot_user_session_from, + validated_copilot_api_base, +) +from litellm.types.utils import CredentialItem, ModelResponse + +EXCHANGE_URL = "https://api.github.com/copilot_internal/v2/token" +DEVICE_CODE_URL = "https://github.com/login/device/code" +ACCESS_TOKEN_URL = "https://github.com/login/oauth/access_token" + + +def _credential(auth_type="per_user_oauth"): + values = {} if auth_type is None else {"github_copilot_auth_type": auth_type} + return CredentialItem(credential_name="copilot-cred", credential_values=values, credential_info={}) + + +def _exchange_payload(token="copilot-token", expires_at=None, api_base="https://api.githubcopilot.com"): + return { + "token": token, + "expires_at": expires_at if expires_at is not None else int(time.time()) + 3600, + "endpoints": {"api": api_base}, + } + + +@contextmanager +def github_http(respond: Callable[[httpx.Request], httpx.Response]): + """Swap both module-level HTTP clients for httpx MockTransport versions and + collect every request, so tests assert on the real HTTP edge (URL, headers) + while no socket is ever opened.""" + requests: list[httpx.Request] = [] + + def capture(request: httpx.Request) -> httpx.Response: + requests.append(request) + return respond(request) + + async_client = AsyncHTTPHandler(transport=httpx.MockTransport(capture)) + sync_client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(capture))) + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=async_client), + patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client", return_value=async_client), + patch("litellm.llms.custom_httpx.llm_http_handler.get_httpx_client", return_value=sync_client), + patch.object(litellm, "module_level_client", sync_client), + ): + yield requests + + +def _exchange_responder(payload=None, status=200): + def respond(request: httpx.Request) -> httpx.Response: + if str(request.url) == EXCHANGE_URL: + return httpx.Response(status, json=payload or _exchange_payload(), request=request) + pytest.fail(f"unexpected request to {request.url}") + + return respond + + +@pytest.fixture(autouse=True) +def _clear_session_cache(): + _SESSION_CACHE.flush_cache() + yield + _SESSION_CACHE.flush_cache() + + +@pytest.fixture +def per_user_credential(): + with patch.object(litellm, "credential_list", [_credential()]): + yield + + +def test_named_per_user_credential_selects_per_user_mode(per_user_credential): + assert github_copilot_auth_mode("copilot-cred", None) is True + + +def test_named_shared_credential_stays_shared_even_if_kwargs_ask_for_per_user(): + with patch.object(litellm, "credential_list", [_credential(auth_type="shared")]): + assert github_copilot_auth_mode("copilot-cred", "per_user_oauth") is False + + +def test_named_credential_without_flag_stays_shared(): + with patch.object(litellm, "credential_list", [_credential(auth_type=None)]): + assert github_copilot_auth_mode("copilot-cred", None) is False + + +def test_per_user_kwargs_without_credential_name_raise_400(): + with pytest.raises(BadRequestError): + github_copilot_auth_mode(None, "per_user_oauth") + + +def test_unknown_named_credential_falls_back_to_kwargs(): + with patch.object(litellm, "credential_list", []): + assert github_copilot_auth_mode("missing-cred", "per_user_oauth") is True + assert github_copilot_auth_mode("missing-cred", None) is False + + +@pytest.mark.parametrize( + "raw,expected", + [ + ("https://api.githubcopilot.com", "https://api.githubcopilot.com"), + ("https://api.githubcopilot.com/", "https://api.githubcopilot.com"), + ("https://tenant.githubcopilot.com:443", "https://tenant.githubcopilot.com:443"), + ("githubcopilot.com", "https://api.githubcopilot.com"), # no scheme -> default + ("http://api.githubcopilot.com", "https://api.githubcopilot.com"), # http rejected + ("https://evil.com", "https://api.githubcopilot.com"), + ("https://githubcopilot.com.evil.com", "https://api.githubcopilot.com"), + ("https://user:pw@api.githubcopilot.com", "https://api.githubcopilot.com"), + ("https://api.githubcopilot.com:8080", "https://api.githubcopilot.com"), + (None, "https://api.githubcopilot.com"), + (1234, "https://api.githubcopilot.com"), + ("https://API.GITHUBCOPILOT.COM", "https://API.GITHUBCOPILOT.COM"), + ], +) +def test_validated_copilot_api_base(raw, expected): + assert validated_copilot_api_base(raw) == expected + + +def test_session_repr_and_str_hide_the_token(caplog): + session = GithubCopilotUserSession(token="secret-copilot-token", api_base="https://api.githubcopilot.com") + assert "secret-copilot-token" not in repr(session) + assert "secret-copilot-token" not in str(session) + with caplog.at_level("DEBUG"): + import logging + + logging.getLogger("litellm").debug("session: %s", session) + assert "secret-copilot-token" not in caplog.text + + +def test_exchange_returns_session_and_caches_per_user(): + with github_http(_exchange_responder()) as requests: + s1 = exchange_github_token("user-a", "gh-token", "copilot-cred") + assert len(requests) == 1 + assert requests[0].headers["Authorization"] == "token gh-token" + s2 = exchange_github_token("user-a", "gh-token", "copilot-cred") + assert len(requests) == 1 + assert s2.token == s1.token == "copilot-token" + # different user -> separate exchange + exchange_github_token("user-b", "gh-token", "copilot-cred") + assert len(requests) == 2 + # different github token -> separate exchange + exchange_github_token("user-a", "gh-token-2", "copilot-cred") + assert len(requests) == 3 + + +@pytest.mark.parametrize("status", [401, 403, 404]) +def test_exchange_auth_failures_raise_caller_auth_error_and_evict(status): + with github_http(_exchange_responder(status=status)): + with pytest.raises(CallerCredentialAuthenticationError) as exc: + exchange_github_token("user-a", "gh-token", "copilot-cred") + assert "Reconnect GitHub Copilot for credential 'copilot-cred' in the LiteLLM UI (LLM Credentials)" in str( + exc.value + ) + assert _SESSION_CACHE.get_cache(_session_cache_key("user-a", "gh-token")) is None + + +def test_exchange_429_raises_caller_rate_limit(): + with github_http(_exchange_responder(status=429)): + with pytest.raises(CallerCredentialRateLimitError): + exchange_github_token("user-a", "gh-token", "copilot-cred") + + +def test_exchange_other_failure_raises_service_unavailable_without_body(): + with github_http(_exchange_responder(status=500, payload={"secret": "body"})): + with pytest.raises(ServiceUnavailableError) as exc: + exchange_github_token("user-a", "gh-token", "copilot-cred") + assert "body" not in str(exc.value) + + +def test_exchange_expiry_margin_refetches_near_expiry(): + with github_http(_exchange_responder(_exchange_payload(expires_at=int(time.time()) + 30))) as requests: + exchange_github_token("user-a", "gh-token", "copilot-cred") + # ttl under the 60s safety margin -> not cached, second call refetches + exchange_github_token("user-a", "gh-token", "copilot-cred") + assert len(requests) == 2 + + +@pytest.mark.asyncio +async def test_aexchange_single_flight_one_http_call(): + def respond(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json=_exchange_payload(), request=request) + + with github_http(respond) as requests: + s1, s2 = await asyncio.gather( + aexchange_github_token("user-a", "gh-token", "copilot-cred"), + aexchange_github_token("user-a", "gh-token", "copilot-cred"), + ) + assert len(requests) == 1 + assert s1.token == s2.token == "copilot-token" + + +@pytest.mark.asyncio +async def test_aexchange_401_raises_and_evicts(): + with github_http(_exchange_responder(status=401)): + with pytest.raises(CallerCredentialAuthenticationError): + await aexchange_github_token("user-a", "gh-token", "copilot-cred") + assert _SESSION_CACHE.get_cache(_session_cache_key("user-a", "gh-token")) is None + + +def _kwargs_with_connection(token="gho_user_token"): + return { + "litellm_credential_name": "copilot-cred", + "secret_fields": { + "user_provider_credentials_user_id": "user-a", + "user_provider_credentials": {"copilot-cred": token}, + }, + } + + +@pytest.mark.asyncio +async def test_aattach_injects_session_into_kwargs(per_user_credential): + with github_http(_exchange_responder()): + kwargs = _kwargs_with_connection() + await aattach_github_copilot_user_session(kwargs) + session = github_copilot_user_session_from(kwargs) + assert session is not None + assert session.token == "copilot-token" + + +def test_attach_sync_injects_session(per_user_credential): + with github_http(_exchange_responder()): + kwargs = _kwargs_with_connection() + attach_github_copilot_user_session(kwargs) + assert isinstance(kwargs[GITHUB_COPILOT_USER_SESSION_KWARG_KEY], GithubCopilotUserSession) + + +def test_attach_per_user_without_connection_raises_connect_401(per_user_credential): + kwargs = {"litellm_credential_name": "copilot-cred", "secret_fields": {}} + with pytest.raises(CallerCredentialAuthenticationError) as exc: + attach_github_copilot_user_session(kwargs) + assert "Connect GitHub Copilot for credential 'copilot-cred' in the LiteLLM UI (LLM Credentials)" in str(exc.value) + + +def test_attach_shared_mode_does_nothing_without_connection(): + with patch.object(litellm, "credential_list", [_credential(auth_type=None)]): + kwargs = {"litellm_credential_name": "copilot-cred"} + attach_github_copilot_user_session(kwargs) + assert GITHUB_COPILOT_USER_SESSION_KWARG_KEY not in kwargs + + +def test_evict_drops_cached_session(): + with github_http(_exchange_responder()) as requests: + exchange_github_token("user-a", "gh-token", "copilot-cred") + evict_copilot_user_session("user-a", "gh-token") + exchange_github_token("user-a", "gh-token", "copilot-cred") + assert len(requests) == 2 + + +@pytest.mark.asyncio +async def test_astart_device_flow_shape(): + def respond(request: httpx.Request) -> httpx.Response: + assert str(request.url) == DEVICE_CODE_URL + return httpx.Response( + 200, + json={ + "device_code": "dc", + "user_code": "UC-123", + "verification_uri": "https://github.com/login/device", + "expires_in": 900, + "interval": 5, + }, + request=request, + ) + + with github_http(respond): + start = await astart_device_flow() + assert (start.device_code, start.user_code, start.verification_uri, start.expires_in, start.interval) == ( + "dc", + "UC-123", + "https://github.com/login/device", + 900, + 5, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "github_error,expected_status", + [ + ("authorization_pending", "pending"), + ("slow_down", "slow_down"), + ("expired_token", "expired"), + ("access_denied", "denied"), + ], +) +async def test_apoll_device_flow_error_mapping(github_error, expected_status): + def respond(request: httpx.Request) -> httpx.Response: + assert str(request.url) == ACCESS_TOKEN_URL + return httpx.Response(200, json={"error": github_error}, request=request) + + with github_http(respond): + poll = await apoll_device_flow("dc") + assert poll.status == expected_status + + +@pytest.mark.asyncio +async def test_apoll_device_flow_connected(): + def respond(request: httpx.Request) -> httpx.Response: + assert str(request.url) == ACCESS_TOKEN_URL + return httpx.Response(200, json={"access_token": "gho_x"}, request=request) + + with github_http(respond): + poll = await apoll_device_flow("dc") + assert poll.status == "connected" and poll.access_token == "gho_x" + + +def test_is_github_copilot_per_user_request(): + from litellm.llms.github_copilot.per_user_auth import is_github_copilot_per_user_request + + session = GithubCopilotUserSession(token="t", api_base="https://api.githubcopilot.com") + assert is_github_copilot_per_user_request({GITHUB_COPILOT_USER_SESSION_KWARG_KEY: session}) is True + assert is_github_copilot_per_user_request({}) is False + assert is_github_copilot_per_user_request({GITHUB_COPILOT_USER_SESSION_KWARG_KEY: "not-a-session"}) is False + assert is_github_copilot_per_user_request({"litellm_credential_name": "copilot-cred"}) is False + + +def _chat_completion_response(): + from litellm.utils import convert_to_model_response_object + + return convert_to_model_response_object( + response_object={ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hello"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, + }, + model_response_object=ModelResponse(), + ) + + +def _secret_fields(user_id: str, github_token: str): + return { + "user_provider_credentials_user_id": user_id, + "user_provider_credentials": {"copilot-cred": github_token}, + } + + +def test_two_per_user_calls_with_identical_bodies_both_reach_upstream( # test-quality-ok: upstream dispatch count is the observable signal of the cache exclusion + per_user_credential, monkeypatch +): + """Response-cache keys ignore caller identity: without the exclusion the second user's call + would be served from the first's cache entry and neither run on their own seat.""" + from litellm.caching import Cache + + monkeypatch.setattr(litellm, "cache", Cache()) + with ( + github_http(_exchange_responder()), + patch( + "litellm.main.openai_chat_completions.completion", + return_value=_chat_completion_response(), + ) as mock_completion, + ): + base = { + "model": "github_copilot/gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "litellm_credential_name": "copilot-cred", + } + litellm.completion(**base, secret_fields=_secret_fields("user-a", "gho_a")) + litellm.completion(**base, secret_fields=_secret_fields("user-b", "gho_b")) + assert mock_completion.call_count == 2 + + +@pytest.mark.asyncio +async def test_two_per_user_async_calls_with_identical_bodies_both_reach_upstream( # test-quality-ok: upstream dispatch count is the observable signal of the cache exclusion + per_user_credential, monkeypatch +): + """Async path: per-user calls must bypass async_get_cache read and write the same way.""" + from litellm.caching import Cache + + monkeypatch.setattr(litellm, "cache", Cache()) + with ( + github_http(_exchange_responder()), + patch( + "litellm.main.openai_chat_completions.acompletion", + new=AsyncMock(return_value=_chat_completion_response()), + ) as mock_completion, + ): + base = { + "model": "github_copilot/gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "litellm_credential_name": "copilot-cred", + } + await litellm.acompletion(**base, secret_fields=_secret_fields("user-a", "gho_a")) + await litellm.acompletion(**base, secret_fields=_secret_fields("user-b", "gho_b")) + assert mock_completion.call_count == 2 + + +def test_shared_mode_second_identical_call_hits_response_cache( # test-quality-ok: the cache hit is only observable as the upstream not being dispatched + tmp_path, monkeypatch +): + """Control: shared device-login github_copilot keeps caching exactly as before, so a second + identical call does not touch the upstream.""" + import json + + from litellm.caching import Cache + + # a mounted shared login: api-key.json with a far-future token, no device flow needed + (tmp_path / "api-key.json").write_text( + json.dumps( + { + "token": "shared-copilot-token", + "expires_at": 4102444800, + "endpoints": {"api": "https://api.githubcopilot.com"}, + } + ) + ) + monkeypatch.setenv("GITHUB_COPILOT_TOKEN_DIR", str(tmp_path)) + + monkeypatch.setattr(litellm, "cache", Cache()) + with patch( + "litellm.main.openai_chat_completions.completion", + return_value=_chat_completion_response(), + ) as mock_completion: + call = { + "model": "github_copilot/gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + } + litellm.completion(**call) + litellm.completion(**call) + assert mock_completion.call_count == 1, "shared mode must still serve the second call from cache" + + +def _embedding_response_payload(): + return { + "object": "list", + "data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 4, "total_tokens": 4}, + } + + +def _responses_payload(): + return { + "id": "resp_1", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-5.3-codex", + "output": [ + { + "type": "message", + "id": "m1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "hi", "annotations": []}], + } + ], + "usage": {"input_tokens": 4, "output_tokens": 2, "total_tokens": 6}, + } + + +def _anthropic_messages_payload(): + return { + "id": "msg_1", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "hi"}], + "model": "claude-haiku-4.5", + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 4, "output_tokens": 2}, + } + + +def _copilot_surface_responder(model_payload): + def respond(request: httpx.Request) -> httpx.Response: + if str(request.url) == EXCHANGE_URL: + return httpx.Response(200, json=_exchange_payload(), request=request) + if request.url.host.endswith("githubcopilot.com"): + return httpx.Response(200, json=model_payload, request=request) + pytest.fail(f"unexpected request to {request.url}") + + return respond + + +def _no_authenticator(): + from litellm.llms.github_copilot.authenticator import Authenticator + + return patch.object( # test-quality-ok: patching the shared-login edge asserts it is never consulted + Authenticator, "get_api_key", side_effect=AssertionError("shared Authenticator must not run") + ) + + +def _assert_exchange_then_model_call(requests, github_token, model_host="https://api.githubcopilot.com"): + exchange = [r for r in requests if str(r.url) == EXCHANGE_URL] + upstream = [r for r in requests if r.url.host.endswith("githubcopilot.com")] + assert exchange, "expected the GitHub token exchange request" + assert upstream, "expected the model request to the Copilot host" + assert exchange[0].headers["Authorization"] == f"token {github_token}" + assert upstream[0].headers["Authorization"].startswith("Bearer ") + assert upstream[0].url.host == httpx.URL(model_host).host + + +def test_e2e_completion_per_user_credential(per_user_credential): + with ( + github_http(_exchange_responder()) as requests, + _no_authenticator(), + patch( + "litellm.main.openai_chat_completions.completion", + return_value=_chat_completion_response(), + ) as mock_completion, + ): + response = litellm.completion( + model="github_copilot/gpt-4o", + messages=[{"role": "user", "content": "hi"}], + litellm_credential_name="copilot-cred", + secret_fields=_secret_fields("user-a", "gho_a"), + ) + assert len(requests) == 1 and str(requests[0].url) == EXCHANGE_URL + assert requests[0].headers["Authorization"] == "token gho_a" + call_kwargs = mock_completion.call_args.kwargs + assert call_kwargs["api_key"] == "copilot-token" + assert call_kwargs["api_base"] == "https://api.githubcopilot.com" + extra_headers = call_kwargs["optional_params"]["extra_headers"] + assert extra_headers["Authorization"] == "Bearer copilot-token" + assert response.usage.total_tokens > 0 + + +def test_e2e_completion_per_user_session_token_wins_over_caller_authorization(per_user_credential): + """The session token reaches the wire even when the caller passes their own + Authorization in extra_headers (chat dispatches through _complete_custom_openai).""" + with ( + github_http(_exchange_responder()), + _no_authenticator(), + patch( + "litellm.main.openai_chat_completions.completion", + return_value=_chat_completion_response(), + ) as mock_completion, + ): + litellm.completion( + model="github_copilot/gpt-4o", + messages=[{"role": "user", "content": "hi"}], + litellm_credential_name="copilot-cred", + secret_fields=_secret_fields("user-a", "gho_a"), + extra_headers={"Authorization": "Bearer caller-token", "x-custom": "keep"}, + ) + extra_headers = mock_completion.call_args.kwargs["optional_params"]["extra_headers"] + assert extra_headers["Authorization"] == "Bearer copilot-token" + assert extra_headers["x-custom"] == "keep" + + +_EVIL_AUTH_HEADERS = {"authorization": "Bearer caller-evil", "x-keep": "1"} + + +def _assert_session_token_on_the_wire(requests): + upstream = [r for r in requests if r.url.host.endswith("githubcopilot.com")] + assert upstream, "expected the model request to the Copilot host" + assert upstream[0].headers["authorization"] == "Bearer copilot-token" + + +def test_e2e_completion_per_user_caller_authorization_cannot_override_session(per_user_credential): + """Caller-supplied Authorization must be stripped at the wrapper: chat's SDK + client owns its own httpx transport, so the capture point is the headers + handed to openai_chat_completions.completion.""" + caller_headers: dict[str, str] = dict(_EVIL_AUTH_HEADERS) + with ( + github_http(_exchange_responder()), + _no_authenticator(), + patch( + "litellm.main.openai_chat_completions.completion", + return_value=_chat_completion_response(), + ) as mock_completion, + ): + litellm.completion( + model="github_copilot/gpt-4o", + messages=[{"role": "user", "content": "hi"}], + litellm_credential_name="copilot-cred", + secret_fields=_secret_fields("user-a", "gho_a"), + extra_headers=caller_headers, + ) + extra_headers = mock_completion.call_args.kwargs["optional_params"]["extra_headers"] + assert extra_headers["Authorization"] == "Bearer copilot-token" + assert "authorization" not in extra_headers + assert extra_headers["x-keep"] == "1" + assert caller_headers == dict(_EVIL_AUTH_HEADERS), "the caller's dict must not be mutated" + + +def test_e2e_embedding_per_user_caller_authorization_cannot_override_session(per_user_credential): + with ( + github_http(_copilot_surface_responder(_embedding_response_payload())) as requests, + _no_authenticator(), + ): + litellm.embedding( + model="github_copilot/text-embedding-3-small", + input=["hi"], + litellm_credential_name="copilot-cred", + secret_fields=_secret_fields("user-a", "gho_a"), + extra_headers=dict(_EVIL_AUTH_HEADERS), + ) + _assert_session_token_on_the_wire(requests) + + +def test_e2e_responses_per_user_caller_authorization_cannot_override_session(per_user_credential): + with ( + github_http(_copilot_surface_responder(_responses_payload())) as requests, + _no_authenticator(), + ): + litellm.responses( + model="github_copilot/gpt-5.3-codex", + input=[{"role": "user", "content": "hi"}], + litellm_credential_name="copilot-cred", + secret_fields=_secret_fields("user-a", "gho_a"), + extra_headers=dict(_EVIL_AUTH_HEADERS), + ) + _assert_session_token_on_the_wire(requests) + + +@pytest.mark.asyncio +async def test_e2e_anthropic_messages_per_user_caller_authorization_cannot_override_session( + per_user_credential, +): + with ( + github_http(_copilot_surface_responder(_anthropic_messages_payload())) as requests, + _no_authenticator(), + ): + await litellm.anthropic.messages.acreate( + model="github_copilot/claude-haiku-4.5", + max_tokens=16, + messages=[{"role": "user", "content": "hi"}], + litellm_credential_name="copilot-cred", + secret_fields=_secret_fields("user-a", "gho_a"), + extra_headers=dict(_EVIL_AUTH_HEADERS), + ) + _assert_session_token_on_the_wire(requests) + + +def test_e2e_embedding_per_user_credential(per_user_credential): + with ( + github_http(_copilot_surface_responder(_embedding_response_payload())) as requests, + _no_authenticator(), + ): + response = litellm.embedding( + model="github_copilot/text-embedding-3-small", + input=["hi"], + litellm_credential_name="copilot-cred", + secret_fields=_secret_fields("user-a", "gho_a"), + ) + _assert_exchange_then_model_call(requests, "gho_a") + assert requests[-1].url.path == "/embeddings" + assert response.usage.prompt_tokens == 4 + + +def test_e2e_responses_per_user_credential(per_user_credential): + with ( + github_http(_copilot_surface_responder(_responses_payload())) as requests, + _no_authenticator(), + ): + response = litellm.responses( + model="github_copilot/gpt-5.3-codex", + input=[{"role": "user", "content": "hi"}], + litellm_credential_name="copilot-cred", + secret_fields=_secret_fields("user-a", "gho_a"), + ) + _assert_exchange_then_model_call(requests, "gho_a") + assert requests[-1].url.path == "/responses" + assert response.usage.total_tokens == 6 + + +@pytest.mark.asyncio +async def test_e2e_anthropic_messages_per_user_credential(per_user_credential): + with ( + github_http(_copilot_surface_responder(_anthropic_messages_payload())) as requests, + _no_authenticator(), + ): + response = await litellm.anthropic.messages.acreate( + model="github_copilot/claude-haiku-4.5", + max_tokens=16, + messages=[{"role": "user", "content": "hi"}], + litellm_credential_name="copilot-cred", + secret_fields=_secret_fields("user-a", "gho_a"), + ) + _assert_exchange_then_model_call(requests, "gho_a") + assert requests[-1].url.path == "/v1/messages" + assert response["usage"]["input_tokens"] == 4 + + +_CONNECT_MESSAGE = "Connect GitHub Copilot for credential 'copilot-cred' in the LiteLLM UI (LLM Credentials)" + + +def test_e2e_completion_not_connected_raises_connect_401(per_user_credential): + with pytest.raises(AuthenticationError, match="Connect GitHub Copilot"): + litellm.completion( + model="github_copilot/gpt-4o", + messages=[{"role": "user", "content": "hi"}], + litellm_credential_name="copilot-cred", + secret_fields={"user_provider_credentials_user_id": "user-a", "user_provider_credentials": {}}, + ) + + +def test_e2e_embedding_not_connected_raises_connect_401(per_user_credential): + with pytest.raises(AuthenticationError, match="Connect GitHub Copilot"): + litellm.embedding( + model="github_copilot/text-embedding-3-small", + input=["hi"], + litellm_credential_name="copilot-cred", + secret_fields={"user_provider_credentials_user_id": "user-a", "user_provider_credentials": {}}, + ) + + +def test_e2e_responses_not_connected_raises_connect_401(per_user_credential): + with pytest.raises(AuthenticationError, match="Connect GitHub Copilot"): + litellm.responses( + model="github_copilot/gpt-5.3-codex", + input=[{"role": "user", "content": "hi"}], + litellm_credential_name="copilot-cred", + secret_fields={"user_provider_credentials_user_id": "user-a", "user_provider_credentials": {}}, + ) + + +@pytest.mark.asyncio +async def test_e2e_anthropic_messages_not_connected_raises_connect_401(per_user_credential): + with pytest.raises(AuthenticationError, match="Connect GitHub Copilot"): + await litellm.anthropic.messages.acreate( + model="github_copilot/claude-haiku-4.5", + max_tokens=16, + messages=[{"role": "user", "content": "hi"}], + litellm_credential_name="copilot-cred", + secret_fields={"user_provider_credentials_user_id": "user-a", "user_provider_credentials": {}}, + ) + + +def _shared_login(tmp_path, monkeypatch): + import json + + (tmp_path / "api-key.json").write_text( + json.dumps( + { + "token": "shared-copilot-token", + "expires_at": 4102444800, + "endpoints": {"api": "https://api.githubcopilot.com"}, + } + ) + ) + monkeypatch.setenv("GITHUB_COPILOT_TOKEN_DIR", str(tmp_path)) + + +def test_e2e_completion_shared_mode_still_uses_authenticator(tmp_path, monkeypatch): + _shared_login(tmp_path, monkeypatch) + with patch( + "litellm.main.openai_chat_completions.completion", + return_value=_chat_completion_response(), + ) as mock_completion: + response = litellm.completion( + model="github_copilot/gpt-4o", + messages=[{"role": "user", "content": "hi"}], + ) + assert mock_completion.call_args.kwargs["api_key"] == "shared-copilot-token" + assert response.usage.total_tokens > 0 + + +def test_e2e_embedding_shared_mode_still_uses_authenticator(tmp_path, monkeypatch): + _shared_login(tmp_path, monkeypatch) + with github_http(_copilot_surface_responder(_embedding_response_payload())) as requests: + response = litellm.embedding( + model="github_copilot/text-embedding-3-small", + input=["hi"], + ) + upstream = [r for r in requests if r.url.host.endswith("githubcopilot.com")] + assert upstream and upstream[0].headers["Authorization"] == "Bearer shared-copilot-token" + assert response.usage.prompt_tokens == 4 + + +def test_e2e_responses_shared_mode_still_uses_authenticator(tmp_path, monkeypatch): + _shared_login(tmp_path, monkeypatch) + with github_http(_copilot_surface_responder(_responses_payload())) as requests: + response = litellm.responses( + model="github_copilot/gpt-5.3-codex", + input=[{"role": "user", "content": "hi"}], + ) + upstream = [r for r in requests if r.url.host.endswith("githubcopilot.com")] + assert upstream and upstream[0].headers["Authorization"] == "Bearer shared-copilot-token" + assert response.usage.total_tokens == 6 + + +@pytest.mark.asyncio +async def test_e2e_anthropic_messages_shared_mode_still_uses_authenticator(tmp_path, monkeypatch): + # The shared Authenticator's access-token refresh refuses to run inside an event + # loop, so on this async surface the mounted-login path can't reach it; patch the + # shared-login edge to emulate a mounted api-key.json. + from litellm.llms.github_copilot.authenticator import Authenticator + + with ( + github_http(_copilot_surface_responder(_anthropic_messages_payload())) as requests, + patch.object( # test-quality-ok: patching the shared-login edge is the fixture under test + Authenticator, "get_api_key", return_value="shared-copilot-token" + ) as mock_api_key, + patch.object( # test-quality-ok: same edge; api_key read is mocked so base must be too + Authenticator, "get_api_base", return_value="https://api.githubcopilot.com" + ), + ): + response = await litellm.anthropic.messages.acreate( + model="github_copilot/claude-haiku-4.5", + max_tokens=16, + messages=[{"role": "user", "content": "hi"}], + ) + mock_api_key.assert_called() + upstream = [r for r in requests if r.url.host.endswith("githubcopilot.com")] + assert upstream and upstream[0].headers["Authorization"] == "Bearer shared-copilot-token" + assert response["usage"]["input_tokens"] == 4 + + +def test_sync_exchange_single_flight_one_http_call(): + import threading + + def slow_respond(request: httpx.Request) -> httpx.Response: + time.sleep(0.05) + return httpx.Response(200, json=_exchange_payload(), request=request) + + sessions: list[GithubCopilotUserSession] = [] + with github_http(slow_respond) as requests: + threads = [ + threading.Thread( + target=lambda: sessions.append(exchange_github_token("user-a", "gh-token", "copilot-cred")) + ) + for _ in range(6) + ] + for t in threads: + t.start() + for t in threads: + t.join() + assert len(requests) == 1 + assert len(sessions) == 6 and all(s.token == "copilot-token" for s in sessions) + + +def test_two_users_same_github_token_do_not_share_session(): + with github_http(_exchange_responder()): + s1 = exchange_github_token("user-a", "gh-same", "copilot-cred") + s2 = exchange_github_token("user-b", "gh-same", "copilot-cred") + assert _session_cache_key("user-a", "gh-same") != _session_cache_key("user-b", "gh-same") + assert s1.token == s2.token == "copilot-token" + + +def test_router_constructs_with_per_user_credential_deployment(per_user_credential): + from litellm.llms.github_copilot.authenticator import Authenticator + + with patch.object(Authenticator, "get_api_key", side_effect=AssertionError("shared login")) as get_key: + router = litellm.Router( + model_list=[ + { + "model_name": "copilot", + "litellm_params": { + "model": "github_copilot/gpt-4.1", + "litellm_credential_name": "copilot-cred", + }, + } + ] + ) + get_key.assert_not_called() + deployments: Final = router.get_model_list() + assert deployments is not None + assert any(d["model_name"] == "copilot" for d in deployments) + + +def test_provider_info_per_user_without_session_returns_default_base_and_no_key(per_user_credential): + from litellm.llms.github_copilot.authenticator import Authenticator + from litellm.types.router import LiteLLM_Params + + params: Final = LiteLLM_Params( + model="github_copilot/gpt-4.1", + litellm_credential_name="copilot-cred", + ) + with patch.object(Authenticator, "get_api_key", side_effect=AssertionError("shared login")) as get_key: + _, _, dynamic_api_key, api_base = litellm.get_llm_provider( + model="github_copilot/gpt-4.1", litellm_params=params + ) + get_key.assert_not_called() + assert api_base == "https://api.githubcopilot.com" + assert dynamic_api_key is None + + +def test_provider_info_shared_github_copilot_still_uses_authenticator(): + from litellm.llms.github_copilot.authenticator import Authenticator + from litellm.types.router import LiteLLM_Params + + params: Final = LiteLLM_Params(model="github_copilot/gpt-4.1") + with ( + patch.object(Authenticator, "get_api_key", return_value="shared-token") as get_key, + patch.object(Authenticator, "get_api_base", return_value="https://api.githubcopilot.com"), + ): + _, _, dynamic_api_key, _ = litellm.get_llm_provider(model="github_copilot/gpt-4.1", litellm_params=params) + get_key.assert_called_once() + assert dynamic_api_key == "shared-token" + + +@pytest.mark.asyncio +async def test_aexchange_cancelled_waiter_does_not_cancel_the_shared_exchange(per_user_credential): + """Two waiters share one in-flight exchange; cancelling one must not cancel the + underlying future the other is waiting on.""" + gate = asyncio.Event() + + class _GatedClient: + async def get(self, url, headers=None): + await gate.wait() + return httpx.Response(200, json=_exchange_payload(), request=httpx.Request("GET", url)) + + with patch( # test-quality-ok: the gate holds the HTTP edge open so two waiters overlap + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=_GatedClient() + ): + first = asyncio.ensure_future(aexchange_github_token("user-a", "gh-shared", "copilot-cred")) + await asyncio.sleep(0) + second = asyncio.ensure_future(aexchange_github_token("user-a", "gh-shared", "copilot-cred")) + await asyncio.sleep(0) + second.cancel() + gate.set() + session = await first + assert session.token == "copilot-token" and session.api_base == "https://api.githubcopilot.com" + with pytest.raises(asyncio.CancelledError): + await second + + +@pytest.mark.asyncio +async def test_aexchange_failure_with_a_single_caller_emits_no_unretrieved_warning(per_user_credential): + """The leader removes the failed future before anyone awaits it; the stored + exception must be consumed so the loop does not log 'never retrieved'.""" + import litellm.llms.github_copilot.per_user_auth as pua + + captured: list[tuple[object, object]] = [] + loop = asyncio.get_running_loop() + previous_handler = loop.get_exception_handler() + loop.set_exception_handler(lambda _loop, context: captured.append((_loop, context))) + try: + with ( + github_http(_exchange_responder(status=401)), + pytest.raises(CallerCredentialAuthenticationError), + ): + await aexchange_github_token("user-a", "gh-fail", "copilot-cred") + _IN_FLIGHT_KEYS = list(pua._IN_FLIGHT) + import gc + + gc.collect() + await asyncio.sleep(0) + finally: + loop.set_exception_handler(previous_handler) + assert _IN_FLIGHT_KEYS == [] + assert not any( + "never retrieved" in str(context.get("message", "")) for _, context in captured if isinstance(context, dict) + ) + + +@pytest.mark.asyncio +async def test_per_user_streaming_call_writes_nothing_to_the_response_cache(per_user_credential, monkeypatch): + """Per-user streamed completions must never be written to the shared response + cache: the handler built before the session attach would otherwise use a stale + request_kwargs copy without the session marker.""" + from litellm.caching import Cache + from litellm.caching.caching_handler import LLMCachingHandler + + monkeypatch.setattr(litellm, "cache", Cache()) + handler = LLMCachingHandler( + original_function=litellm.acompletion, + request_kwargs={"model": "github_copilot/gpt-4o"}, + start_time=datetime.datetime.now(), + ) + kwargs = { + "model": "github_copilot/gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + GITHUB_COPILOT_USER_SESSION_KWARG_KEY: GithubCopilotUserSession( + token="copilot-token", api_base="https://api.githubcopilot.com" + ), + } + await handler.async_get_cache( + model="github_copilot/gpt-4o", + original_function=litellm.acompletion, + logging_obj=MagicMock(), + start_time=datetime.datetime.now(), + call_type="acompletion", + kwargs=kwargs, + ) + assert ( + handler.should_store_result_in_cache(original_function=litellm.acompletion, kwargs=handler.request_kwargs) + is False + ) + + +@pytest.mark.asyncio +async def test_shared_streaming_call_still_writes_response_cache(tmp_path, monkeypatch): + """Control: shared device-login streaming keeps caching normally.""" + _shared_login(tmp_path, monkeypatch) + from litellm.caching import Cache + from litellm.caching.caching_handler import LLMCachingHandler + + monkeypatch.setattr(litellm, "cache", Cache()) + handler = LLMCachingHandler( + original_function=litellm.acompletion, + request_kwargs={"model": "github_copilot/gpt-4o"}, + start_time=datetime.datetime.now(), + ) + kwargs = { + "model": "github_copilot/gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + } + await handler.async_get_cache( + model="github_copilot/gpt-4o", + original_function=litellm.acompletion, + logging_obj=MagicMock(), + start_time=datetime.datetime.now(), + call_type="acompletion", + kwargs=kwargs, + ) + assert ( + handler.should_store_result_in_cache(original_function=litellm.acompletion, kwargs=handler.request_kwargs) + is True + ) + + +def _per_user_router(): + return litellm.Router( + model_list=[ + { + "model_name": "copilot", + "litellm_params": { + "model": "github_copilot/gpt-4.1", + "litellm_credential_name": "copilot-cred", + }, + } + ] + ) + + +def test_per_user_deployment_rejects_a_caller_credential_override(per_user_credential): + """Caller kwargs win the litellm_params merge downstream, so the deployment + layer is the last place that still sees both names: a different name on a + per-user deployment must be refused there even if discovery missed the hop.""" + router: Final = _per_user_router() + deployment: Final = router.get_model_list()[0] + kwargs: Final = {"litellm_credential_name": "shared-cred", "metadata": {}} + with pytest.raises(litellm.BadRequestError): + router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) + + +def test_per_user_deployment_allows_the_configured_credential_name(per_user_credential): + router: Final = _per_user_router() + deployment: Final = router.get_model_list()[0] + kwargs: Final = {"litellm_credential_name": "copilot-cred", "metadata": {}} + router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) + assert kwargs["litellm_credential_name"] == "copilot-cred" + assert "model_info" in kwargs + + +def test_shared_deployment_keeps_caller_credential_override(): + from litellm.llms.github_copilot.authenticator import Authenticator + + with ( + patch.object(litellm, "credential_list", [_credential(auth_type="shared")]), + patch.object(Authenticator, "get_api_key", return_value="shared-token"), + ): + router: Final = _per_user_router() + deployment: Final = router.get_model_list()[0] + kwargs: Final = {"litellm_credential_name": "other-cred", "metadata": {}} + router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) + assert kwargs["litellm_credential_name"] == "other-cred" + assert "model_info" in kwargs diff --git a/tests/unit/proxy/auth/test_auth_utils.py b/tests/unit/proxy/auth/test_auth_utils.py index 412a104f252..8be1cec0a0c 100644 --- a/tests/unit/proxy/auth/test_auth_utils.py +++ b/tests/unit/proxy/auth/test_auth_utils.py @@ -2498,6 +2498,67 @@ class TestGetDynamicLitellmParamsClearsAdminConfigOnBaseOverride: assert out["organization"] == "org-attacker" assert out["extra_body"] == {"attacker": "value"} + def test_per_user_oauth_credential_survives_base_override(self): + # Fail-closed: a caller-redirected api_base must not drop + # litellm_credential_name / github_copilot_auth_type off a per-user + # deployment, which would flip the call to shared mode and send the + # admin's Copilot token to the attacker's host. Keeping them is safe + # because per-user mode ignores the caller api_base anyway. + from litellm.router_utils.clientside_credential_handler import ( + get_dynamic_litellm_params, + ) + from litellm.types.utils import CredentialItem + + with patch.object( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="copilot-cred", + credential_values={"github_copilot_auth_type": "per_user_oauth"}, + credential_info={}, + ) + ], + ): + out = get_dynamic_litellm_params( + litellm_params={ + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + "github_copilot_auth_type": "per_user_oauth", + }, + request_kwargs={"api_base": "https://attacker.example"}, + ) + assert out["litellm_credential_name"] == "copilot-cred" + assert out["github_copilot_auth_type"] == "per_user_oauth" + assert out["api_base"] == "https://attacker.example" + + def test_shared_oauth_credential_name_is_cleared_on_base_override(self): + from litellm.router_utils.clientside_credential_handler import ( + get_dynamic_litellm_params, + ) + from litellm.types.utils import CredentialItem + + with patch.object( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="copilot-cred", + credential_values={"github_copilot_auth_type": "shared"}, + credential_info={}, + ) + ], + ): + out = get_dynamic_litellm_params( + litellm_params={ + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + "github_copilot_auth_type": "shared", + }, + request_kwargs={"api_base": "https://attacker.example"}, + ) + assert "github_copilot_auth_type" not in out + def test_field_echo_does_not_preserve_admin_value(self): # Regression: a caller that echoes an admin-config field name with # an *empty* value (or any value) must not be able to keep the @@ -4386,6 +4447,24 @@ def test_get_model_from_request_vertex_ai_passthrough( assert model == expected_model +@pytest.mark.parametrize( + "field", + ["github_copilot_auth_type", "user_provider_credentials", "github_copilot_user_session"], +) +def test_per_user_credential_slots_in_request_body_are_rejected(field: str) -> None: + """The per-user Copilot mode is decided by the stored credential's values and the caller's + connection by proxy-injected secret_fields; a body-supplied value could only spoof either.""" + with pytest.raises(ValueError, match="Rejected Request") as error: + is_request_body_safe( + request_body={"model": "github_copilot/gpt-4o", field: "attacker-chosen"}, + general_settings={}, + llm_router=None, + model="github_copilot/gpt-4o", + ) + + assert field in str(error.value) + + def test_get_end_user_id_from_request_body_always_returns_str(): mock_request: Final = MagicMock(spec=Request) mock_request.headers = {} diff --git a/tests/unit/proxy/auth/test_route_checks.py b/tests/unit/proxy/auth/test_route_checks.py index 540b66ddaba..de5b2f4b112 100644 --- a/tests/unit/proxy/auth/test_route_checks.py +++ b/tests/unit/proxy/auth/test_route_checks.py @@ -4702,9 +4702,7 @@ def test_is_llm_api_route(): all_llm_api_routes = llm_passthrough_router.routes for route in all_llm_api_routes: - print("route", route) route_path = str(route.path) - print("route_path", route_path) assert RouteChecks.is_llm_api_route(route_path) is True @@ -4720,3 +4718,91 @@ def test_route_matches_pattern(): is False ) assert RouteChecks._route_matches_pattern("/v1/{thread_id}/messages", "/v1/messages/thread_2345") is False + + +def _scope_request(method: str, path: str) -> Request: + return Request({"type": "http", "method": method, "path": path, "query_string": b""}) + + +def _internal_user_check(route: str, method: str, role: LitellmUserRoles = LitellmUserRoles.INTERNAL_USER) -> None: + user_obj = LiteLLM_UserTable(user_id="u", user_email="u@x", user_role=role.value) + valid_token = UserAPIKeyAuth(user_id="u", user_role=role) + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=role.value, + route=route, + request=_scope_request(method, route), + valid_token=valid_token, + request_data={}, + ) + + +@pytest.mark.parametrize( + ("method", "route"), + [ + ("GET", "/credentials/user_connections"), + ("POST", "/credentials/x/user_connection/start"), + ("POST", "/credentials/x/user_connection/poll"), + ("DELETE", "/credentials/x/user_connection"), + ("DELETE", "/credentials/a/b/user_connection"), + ], +) +def test_internal_user_credential_connection_routes_allowed(method: str, route: str) -> None: + assert _internal_user_check(route, method) is None + + +@pytest.mark.parametrize( + ("method", "route"), + [ + ("PATCH", "/credentials/x/user_connection"), + ("GET", "/credentials/by_name/x/user_connection"), + ("PATCH", "/credentials/x/user_connection/start"), + ("GET", "/credentials/x/user_connection/poll"), + ("DELETE", "/credentials/user_connections"), + ("PATCH", "/credentials/user_connections"), + ], +) +def test_internal_user_credential_connection_route_collisions_denied(method: str, route: str) -> None: + """A name colliding with the connection paths would fall through to the generic + credential CRUD handlers, so only the exact method/route pairs are allowed.""" + with pytest.raises(Exception, match="Only proxy admin"): + _internal_user_check(route, method) + + +@pytest.mark.parametrize( + ("method", "route"), + [ + ("GET", "/credentials/user_connections"), + ("POST", "/credentials/x/user_connection/start"), + ("POST", "/credentials/x/user_connection/poll"), + ("DELETE", "/credentials/x/user_connection"), + ], +) +def test_view_only_user_denied_credential_connection_routes(method: str, route: str) -> None: + with pytest.raises(Exception, match="Only proxy admin"): + _internal_user_check(route, method, role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY) + + +def test_user_connection_route_registration_order_wins_over_generic_credential_routes() -> None: + """The :path placeholders in the generic CRUD routes could swallow the + user_connection paths; the concrete routes must register (and match) first.""" + from starlette.routing import Match + + from litellm.proxy.credential_endpoints import endpoints as credential_endpoints + from litellm.proxy.proxy_server import app + + for method, path, expected in ( + ("DELETE", "/credentials/foo/user_connection", credential_endpoints.delete_user_connection), + ("POST", "/credentials/foo/user_connection/start", credential_endpoints.start_user_connection), + ): + scope = {"type": "http", "method": method, "path": path} + match = next( + ( + route + for route in app.router.routes + if hasattr(route, "matches") and route.matches(scope)[0] == Match.FULL + ), + None, + ) + assert match is not None, f"no route matched {method} {path}" + assert getattr(match, "endpoint", None) is expected diff --git a/tests/unit/proxy/credential_endpoints/test_endpoints.py b/tests/unit/proxy/credential_endpoints/test_endpoints.py index b69150a95ec..3c4e8b20f72 100644 --- a/tests/unit/proxy/credential_endpoints/test_endpoints.py +++ b/tests/unit/proxy/credential_endpoints/test_endpoints.py @@ -2,19 +2,22 @@ import json from contextlib import contextmanager +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import ec +from fastapi import HTTPException, Request from fastapi.testclient import TestClient import litellm -from litellm.proxy._types import UserAPIKeyAuth +from litellm.models.credentials import CredentialSource +from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.credential_endpoints.endpoints import get_llm_router from litellm.proxy.proxy_server import app -from litellm.models.credentials import CredentialSource from litellm.types.utils import CredentialItem client = TestClient(app) @@ -76,7 +79,9 @@ def credential_store(): llm_router: object | None = None, **repository_calls: AsyncMock, ) -> None: - patch("litellm.proxy.proxy_server.prisma_client", _prisma_without_credential_rows() if connected else None).start() + patch( + "litellm.proxy.proxy_server.prisma_client", _prisma_without_credential_rows() if connected else None + ).start() patch("litellm.proxy.proxy_server.master_key", "sk-test-master").start() patch.object(litellm, "credential_list", list(in_memory)).start() app.dependency_overrides[get_llm_router] = lambda: llm_router @@ -753,7 +758,9 @@ class TestNonAdminCannotPersistWifFieldsOnCredential: update_by_name = AsyncMock(return_value=None) router = MagicMock() router.get_deployment.return_value = {"model_name": "claude-opus-5-5"} - router.get_deployment_credentials.return_value = {"anthropic_keycloak_token_url": "https://keycloak.internal/token"} + router.get_deployment_credentials.return_value = { + "anthropic_keycloak_token_url": "https://keycloak.internal/token" + } credential_store(find_by_name=AsyncMock(return_value=stored), update_by_name=update_by_name, llm_router=router) response = _patch_credential( @@ -1170,7 +1177,9 @@ class TestManagementReadsTheStoredCredential: prisma = MagicMock() prisma.db.litellm_credentialstable.find_unique = AsyncMock(return_value=None) - with patch.object(litellm, "credential_list", [only_in_memory]): # test-quality-ok: the config.yaml fallback under test + with patch.object( + litellm, "credential_list", [only_in_memory] + ): # test-quality-ok: the config.yaml fallback under test resolved = await hydrate_named_credential_authoritative("config-yaml-credential", prisma) assert resolved is not None @@ -1393,6 +1402,767 @@ def test_update_credential_still_accepts_a_body_without_credential_values(creden assert set(json.loads(written["credential_values"])) == {"api_key"}, "stored values survive an info-only patch" +from contextlib import contextmanager as _ctx + +from litellm.caching.dual_cache import DualCache +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.proxy.credential_endpoints import user_provider_credentials as upc + +_DEVICE_CODE_URL = "https://github.com/login/device/code" +_ACCESS_TOKEN_URL = "https://github.com/login/oauth/access_token" +_COPILOT_TOKEN_URL = "https://api.github.com/copilot_internal/v2/token" +_GITHUB_USER_URL = "https://api.github.com/user" + + +def _as_user(user_id="user-a"): + return lambda: UserAPIKeyAuth(api_key="test-key", user_role="internal_user", user_id=user_id) + + +def _per_user_credential(name="copilot-cred"): + return CredentialItem( + credential_name=name, + credential_values={"github_copilot_auth_type": "per_user_oauth"}, + credential_info={}, + ) + + +def _connection_row(user_id="user-a", credential_name="copilot-cred", login="octo"): + from datetime import datetime, timezone + from types import SimpleNamespace + + return SimpleNamespace( + user_id=user_id, + credential_name=credential_name, + provider="github_copilot", + credential_b64=upc._encode( + upc.GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login=login) + ), + created_at=datetime(2026, 1, 1, tzinfo=timezone.utc), + updated_at=datetime(2026, 1, 1, tzinfo=timezone.utc), + ) + + +@pytest.fixture(autouse=True) +def _connection_master_key(monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-master") + + +@_ctx +def _github_http(mapping): + """Patch the shared async HTTP client to an httpx MockTransport that answers + GitHub endpoints from {url: response-json dict or callable}; anything else fails.""" + + def respond(request): + entry = mapping.get(str(request.url)) + if entry is None: + raise AssertionError(f"unexpected request to {request.url}") + payload = entry(request) if callable(entry) else entry + return httpx.Response(200, json=payload, request=request) + + client = AsyncHTTPHandler(transport=httpx.MockTransport(respond)) + with patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client): + yield + + +_DEVICE_FLOW_START = { + "device_code": "dc-1", + "user_code": "UC-1", + "verification_uri": "https://github.com/login/device", + "expires_in": 900, + "interval": 5, +} + + +def _patch_user_connection_env(monkeypatch, credentials, rows=()): + """Per-user env: credential in memory, prisma table answering find_many with rows, + and a real in-memory cache for device-flow + credential caching.""" + monkeypatch.setattr(litellm, "credential_list", list(credentials)) + prisma_client = MagicMock() + table = MagicMock() + table.find_many = AsyncMock(return_value=list(rows)) + table.find_unique = AsyncMock(return_value=rows[0] if rows else None) + table.upsert = AsyncMock(return_value=None) + table.delete_many = AsyncMock(return_value=None) + prisma_client.db.litellm_userprovidercredentials = table + prisma_client.writer_db = prisma_client.db + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + cache = DualCache() + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache) + return table + + +def test_user_connections_lists_per_user_credentials_with_connection_state(monkeypatch): + _patch_user_connection_env( + monkeypatch, + [_per_user_credential(), CredentialItem(credential_name="shared", credential_values={}, credential_info={})], + rows=[_connection_row()], + ) + response = _call_as("GET", "/credentials/user_connections", auth=_as_user()) + assert response.status_code == 200, response.text + assert response.json() == { + "connections": [ + { + "credential_name": "copilot-cred", + "provider": "github_copilot", + "connected": True, + "github_login": "octo", + "connected_at": "2026-01-01T00:00:00+00:00", + } + ] + } + + +def test_user_connections_requires_a_user_id(monkeypatch): + _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + response = _call_as("GET", "/credentials/user_connections", auth=_as_admin) + assert response.status_code == 401 + + +def test_user_connection_start_returns_device_flow_fields(monkeypatch): + _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + with _github_http({_DEVICE_CODE_URL: _DEVICE_FLOW_START}): + response = _call_as("POST", "/credentials/copilot-cred/user_connection/start", auth=_as_user()) + assert response.status_code == 200, response.text + assert {key: response.json()[key] for key in ("user_code", "verification_uri", "expires_in", "interval")} == { + "user_code": "UC-1", + "verification_uri": "https://github.com/login/device", + "expires_in": 900, + "interval": 5, + } + handle = response.json()["flow_handle"] + assert isinstance(handle, str) and "device_code" not in handle and _DEVICE_FLOW_START["device_code"] not in handle + + +def test_user_connection_start_404s_for_shared_credential(monkeypatch): + _patch_user_connection_env( + monkeypatch, [CredentialItem(credential_name="shared", credential_values={}, credential_info={})] + ) + response = _call_as("POST", "/credentials/shared/user_connection/start", auth=_as_user()) + assert response.status_code == 404 + + +def test_user_connection_poll_rejects_an_undecryptable_flow_handle(monkeypatch): + _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + response = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": "not-a-real-handle"}, + auth=_as_user(), + ) + assert response.status_code == 400 + + +def test_user_connection_poll_connected_persists_and_reports_login(monkeypatch): + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + with _github_http( + { + _DEVICE_CODE_URL: _DEVICE_FLOW_START, + _ACCESS_TOKEN_URL: {"access_token": "gho_1"}, + _COPILOT_TOKEN_URL: {"token": "copilot-tok", "expires_at": 4102444800}, + _GITHUB_USER_URL: {"login": "octo"}, + } + ): + start = _call_as("POST", "/credentials/copilot-cred/user_connection/start", auth=_as_user()) + assert start.status_code == 200, start.text + poll = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": start.json()["flow_handle"]}, + auth=_as_user(), + ) + assert poll.status_code == 200, poll.text + assert poll.json() == {"status": "connected", "interval": None, "github_login": "octo"} + table.upsert.assert_awaited_once() + + +def test_user_connection_poll_pending(monkeypatch): + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + with _github_http( + { + _DEVICE_CODE_URL: _DEVICE_FLOW_START, + _ACCESS_TOKEN_URL: {"error": "authorization_pending"}, + } + ): + start = _call_as("POST", "/credentials/copilot-cred/user_connection/start", auth=_as_user()) + poll = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": start.json()["flow_handle"]}, + auth=_as_user(), + ) + assert poll.json()["status"] == "pending" + table.upsert.assert_not_awaited() + + +def test_delete_user_connection_disconnects_and_is_idempotent(monkeypatch): + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()], rows=[_connection_row()]) + response = _call_as("DELETE", "/credentials/copilot-cred/user_connection", auth=_as_user()) + assert response.status_code == 200 + assert response.json() == {"status": "disconnected"} + table.delete_many.assert_awaited_once_with(where={"user_id": "user-a", "credential_name": "copilot-cred"}) + + # idempotent: row now gone + table.find_unique = AsyncMock(return_value=None) + again = _call_as("DELETE", "/credentials/copilot-cred/user_connection", auth=_as_user()) + assert again.status_code == 200 + assert again.json() == {"status": "disconnected"} + + +def test_deleting_the_credential_purges_user_connections(monkeypatch, credential_store): + rows = [_connection_row(user_id="u1"), _connection_row(user_id="u2")] + credential_store( + in_memory=[_per_user_credential()], + delete_by_name=AsyncMock(return_value=_per_user_credential()), + ) + prisma_client = MagicMock() + table = MagicMock() + table.find_many = AsyncMock(return_value=rows) + table.delete_many = AsyncMock(return_value=None) + prisma_client.db.litellm_userprovidercredentials = table + prisma_client.writer_db = prisma_client.db + prisma_client.db.litellm_credentialstable.find_unique = AsyncMock(return_value=None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", DualCache()) + + response = _delete_credential("copilot-cred") + assert response.status_code == 200, response.text + table.delete_many.assert_awaited_with(where={"credential_name": "copilot-cred"}) + + +def test_user_connections_isolated_per_user(monkeypatch): + """Two users connected to the same credential: A's listing shows only A's login, + A's delete removes only A's row, and a poll only reads A's device-flow entry.""" + + rows = [_connection_row(user_id="user-a", login="octo"), _connection_row(user_id="user-b", login="hubot")] + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()], rows=rows) + + async def _find_many_matching_user(*, where): + return [row for row in rows if row.user_id == where["user_id"]] + + table.find_many = AsyncMock(side_effect=_find_many_matching_user) + + listing = _call_as("GET", "/credentials/user_connections", auth=_as_user("user-a")) + assert listing.status_code == 200 + body = listing.json()["connections"] + assert [c["github_login"] for c in body] == ["octo"] + # find_many must be scoped to the calling user + where = table.find_many.await_args.kwargs["where"] + assert where["user_id"] == "user-a" + + # B's in-flight device flow handle is bound to B: A polling with it is rejected + poll = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": _flow_handle_for(user_id="user-b")}, + auth=_as_user("user-a"), + ) + assert poll.status_code == 400 + table.upsert.assert_not_awaited() + + delete = _call_as("DELETE", "/credentials/copilot-cred/user_connection", auth=_as_user("user-a")) + assert delete.json() == {"status": "disconnected"} + table.delete_many.assert_awaited_once_with(where={"user_id": "user-a", "credential_name": "copilot-cred"}) + + +@pytest.mark.parametrize("seat_status", [403, 404]) +def test_user_connection_poll_reports_no_copilot_seat(monkeypatch, seat_status): + """A GitHub seat check that rejects the user's token maps to no_copilot_seat, stores + nothing, and clears the in-flight device code.""" + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + + def respond(request): + url = str(request.url) + if url == _DEVICE_CODE_URL: + return httpx.Response(200, json=_DEVICE_FLOW_START, request=request) + if url == _ACCESS_TOKEN_URL: + return httpx.Response(200, json={"access_token": "gho_seat"}, request=request) + if url == _COPILOT_TOKEN_URL: + return httpx.Response(seat_status, json={}, request=request) + pytest.fail(f"unexpected request to {url}") + + client_obj = AsyncHTTPHandler(transport=httpx.MockTransport(respond)) + with patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client_obj): + start = _call_as("POST", "/credentials/copilot-cred/user_connection/start", auth=_as_user()) + assert start.status_code == 200 + poll = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": start.json()["flow_handle"]}, + auth=_as_user(), + ) + assert poll.status_code == 200, poll.text + assert poll.json()["status"] == "no_copilot_seat" + table.upsert.assert_not_awaited() + + +def test_user_connection_poll_rate_limit_returns_429(monkeypatch): + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + + def respond(request): + url = str(request.url) + if url == _DEVICE_CODE_URL: + return httpx.Response(200, json=_DEVICE_FLOW_START, request=request) + if url == _ACCESS_TOKEN_URL: + return httpx.Response(200, json={"access_token": "gho_limited"}, request=request) + if url == _COPILOT_TOKEN_URL: + return httpx.Response(429, json={}, request=request) + pytest.fail(f"unexpected request to {url}") + + client_obj = AsyncHTTPHandler(transport=httpx.MockTransport(respond)) + with patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client_obj): + start = _call_as("POST", "/credentials/copilot-cred/user_connection/start", auth=_as_user()) + poll = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": start.json()["flow_handle"]}, + auth=_as_user(), + ) + assert poll.status_code == 429 + table.upsert.assert_not_awaited() + + +@pytest.mark.parametrize( + "method,path", + [ + ("POST", "/credentials/missing/user_connection/start"), + ("POST", "/credentials/missing/user_connection/poll"), + ("DELETE", "/credentials/missing/user_connection"), + ("POST", "/credentials/shared/user_connection/poll"), + ("DELETE", "/credentials/shared/user_connection"), + ], +) +def test_user_connection_routes_404_for_missing_or_shared_credentials(monkeypatch, method, path): + _patch_user_connection_env( + monkeypatch, [CredentialItem(credential_name="shared", credential_values={}, credential_info={})] + ) + response = _call_as(method, path, json_body={"flow_handle": "x"}, auth=_as_user()) + assert response.status_code == 404 + + +@pytest.mark.parametrize( + "method,path", + [ + ("GET", "/credentials/user_connections"), + ("POST", "/credentials/copilot-cred/user_connection/start"), + ("POST", "/credentials/copilot-cred/user_connection/poll"), + ("DELETE", "/credentials/copilot-cred/user_connection"), + ], +) +def test_user_connection_routes_401_without_user_id(monkeypatch, method, path): + _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + response = _call_as(method, path, json_body={"flow_handle": "x"}, auth=_as_admin) # admin token has no user_id + assert response.status_code == 401 + + +def test_user_connection_routes_roles(): + """internal_user is allowed (self-scoped), internal_user_view_only is denied, proxy_admin + is allowed; and POST /credentials stays admin-only for an internal user.""" + from litellm.proxy.auth.route_checks import RouteChecks + + routes = [ + ("/credentials/user_connections", "GET"), + ("/credentials/copilot-cred/user_connection/start", "POST"), + ("/credentials/copilot-cred/user_connection/poll", "POST"), + ("/credentials/copilot-cred/user_connection", "DELETE"), + ] + + def outcome(role, route, method="GET"): + if role == LitellmUserRoles.PROXY_ADMIN: + return "allowed" + user_obj = LiteLLM_UserTable(user_id="u", user_email="u@x", user_role=role.value) + valid_token = UserAPIKeyAuth(user_id="u", user_role=role) + request = MagicMock(spec=Request) + request.method = method + request.query_params = {} + try: + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=role.value, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + except HTTPException as exc: + return f"denied:{exc.status_code}" + except Exception: + return "denied" + return "allowed" + + for route, method in routes: + assert outcome(LitellmUserRoles.INTERNAL_USER, route, method) == "allowed" + assert outcome(LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, route, method) == "denied" + assert outcome(LitellmUserRoles.INTERNAL_USER, "/credentials", method="POST") != "allowed" + + +class TestNonAdminCannotSetPerUserOauthOnCredential: + """github_copilot_auth_type selects the GitHub OAuth path; it is server-owned WIF, so a + team admin must not set it on a credential and a proxy admin can.""" + + def test_non_admin_cannot_create_a_per_user_oauth_credential(self): + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ): + response = _post_credential( + { + "credential_name": "attacker-cred", + "credential_values": {"github_copilot_auth_type": "per_user_oauth"}, + "credential_info": {"custom_llm_provider": "github_copilot"}, + }, + auth=_as_non_admin, + ) + assert response.status_code == 403, response.text + assert "github_copilot_auth_type" in response.json()["error"]["message"] + + def test_non_admin_cannot_patch_a_credential_to_per_user_oauth(self, restore_credential_list): + stored = CredentialItem( + credential_name="existing", + credential_values={"api_key": "sk"}, + credential_info={"custom_llm_provider": "github_copilot"}, + ) + with _repository_holding(stored) as repository: + repository.find_unique_by_name = AsyncMock(return_value=stored) + response = _patch_credential( + "existing", + { + "credential_name": "existing", + "credential_values": {"github_copilot_auth_type": "per_user_oauth"}, + "credential_info": {}, + }, + auth=_as_non_admin, + ) + assert response.status_code == 403, response.text + + def test_proxy_admin_can_create_a_per_user_oauth_credential(self, restore_credential_list): + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "admin-cred", + "credential_values": {"github_copilot_auth_type": "per_user_oauth"}, + "credential_info": {"custom_llm_provider": "github_copilot"}, + }, + ) + assert response.status_code == 200, response.text + repository.create.assert_awaited_once() + + def test_proxy_admin_can_patch_a_credential_to_per_user_oauth(self, restore_credential_list): + stored = CredentialItem( + credential_name="existing", + credential_values={"api_key": "sk"}, + credential_info={"custom_llm_provider": "github_copilot"}, + ) + with _repository_holding(stored) as repository: + repository.find_unique_by_name = AsyncMock(return_value=stored) + repository.update_by_name = AsyncMock(return_value=stored) + response = _patch_credential( + "existing", + { + "credential_name": "existing", + "credential_values": {"github_copilot_auth_type": "per_user_oauth"}, + "credential_info": {}, + }, + ) + assert response.status_code == 200, response.text + + +def test_label_only_patch_does_not_purge_user_connections(): + """A PATCH that only changes display_name sends no credential_name in the + body, and that must not read as a rename that purges every user's stored + connection.""" + stored: Final = CredentialItem( + credential_name="copilot-cred", + credential_values={"github_copilot_auth_type": "per_user_oauth"}, + credential_info={"custom_llm_provider": "github_copilot"}, + ) + with _repository_holding(stored): + from litellm.proxy import proxy_server + + table: Final = MagicMock() + table.find_many = AsyncMock(return_value=[_connection_row()]) + table.delete_many = AsyncMock(return_value=None) + proxy_server.prisma_client.db.litellm_userprovidercredentials = table + proxy_server.prisma_client.writer_db = proxy_server.prisma_client.db + + response: Final = _patch_credential( + "copilot-cred", + {"display_name": "Renamed Copilot", "credential_info": {"custom_llm_provider": "github_copilot"}}, + ) + + assert response.status_code == 200, response.text + table.find_many.assert_not_awaited() + table.delete_many.assert_not_awaited() + + +def test_switching_away_from_per_user_oauth_purges_user_connections(): + stored: Final = CredentialItem( + credential_name="copilot-cred", + credential_values={"github_copilot_auth_type": "per_user_oauth"}, + credential_info={"custom_llm_provider": "github_copilot"}, + ) + with _repository_holding(stored): + from litellm.proxy import proxy_server + + table: Final = MagicMock() + table.find_many = AsyncMock(return_value=[_connection_row()]) + table.delete_many = AsyncMock(return_value=None) + proxy_server.prisma_client.db.litellm_userprovidercredentials = table + proxy_server.prisma_client.writer_db = proxy_server.prisma_client.db + + response: Final = _patch_credential( + "copilot-cred", + { + "credential_name": "copilot-cred", + "credential_values": {}, + "credential_values_to_delete": ["github_copilot_auth_type"], + "credential_info": {"custom_llm_provider": "github_copilot"}, + }, + ) + + assert response.status_code == 200, response.text + table.delete_many.assert_awaited_with(where={"credential_name": "copilot-cred"}) + + +def test_pending_device_flow_survives_disconnect_via_stateless_handle(monkeypatch): + """The device flow is stateless: no worker-local entry exists to clear on + disconnect, so a still-valid issued handle polled on any worker completes.""" + _patch_user_connection_env(monkeypatch, [_per_user_credential()], rows=[_connection_row()]) + + with _github_http({_DEVICE_CODE_URL: _DEVICE_FLOW_START}): + start = _call_as("POST", "/credentials/copilot-cred/user_connection/start", auth=_as_user()) + assert start.status_code == 200 + delete = _call_as("DELETE", "/credentials/copilot-cred/user_connection", auth=_as_user()) + assert delete.json() == {"status": "disconnected"} + + +def test_user_connection_cache_keys_cannot_collide_on_colons(): + """user_id/credential_name are joined losslessly, so ('a:b','c') and ('a','b:c') + can never share a cache or device-flow entry.""" + assert upc._cache_key("a:b", "c") != upc._cache_key("a", "b:c") + + +def test_user_connection_slashed_credential_name_routes(monkeypatch): + """A credential named 'team/copilot' must still reach start/poll/delete: the + routes take :path parameters the same way the admin CRUD routes do.""" + table = _patch_user_connection_env(monkeypatch, [_per_user_credential(name="team/copilot")]) + with _github_http( + { + _DEVICE_CODE_URL: _DEVICE_FLOW_START, + _ACCESS_TOKEN_URL: {"error": "authorization_pending"}, + } + ): + start = _call_as("POST", "/credentials/team/copilot/user_connection/start", auth=_as_user()) + assert start.status_code == 200, start.text + assert start.json()["user_code"] == "UC-1" + + poll = _call_as( + "POST", + "/credentials/team/copilot/user_connection/poll", + json_body={"flow_handle": start.json()["flow_handle"]}, + auth=_as_user(), + ) + assert poll.status_code == 200 + assert poll.json()["status"] == "pending" + + delete = _call_as("DELETE", "/credentials/team/copilot/user_connection", auth=_as_user()) + assert delete.status_code == 200 + assert delete.json() == {"status": "disconnected"} + table.delete_many.assert_awaited_once_with(where={"user_id": "user-a", "credential_name": "team/copilot"}) + + +def test_user_connection_route_check_allows_slashed_names_for_internal_users(): + """The RouteChecks allowlist entries use :path placeholders too, so an internal + user's request to /credentials/team/copilot/user_connection/... still matches.""" + from litellm.proxy.auth.route_checks import RouteChecks + + for route in ( + "/credentials/team/copilot/user_connection/start", + "/credentials/team/copilot/user_connection/poll", + "/credentials/team/copilot/user_connection", + ): + assert RouteChecks._route_matches_pattern( + route=route, pattern=route.replace("team/copilot", "{credential_name:path}") + ) + + user_obj = LiteLLM_UserTable(user_id="u", user_email="u@x", user_role=LitellmUserRoles.INTERNAL_USER.value) + request = MagicMock(spec=Request) + request.method = "POST" + request.query_params = {} + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.INTERNAL_USER.value, + route="/credentials/team/copilot/user_connection/poll", + request=request, + valid_token=UserAPIKeyAuth(user_id="u", user_role=LitellmUserRoles.INTERNAL_USER), + request_data={}, + ) + + +def _flow_handle_for(user_id="user-a", credential_name="copilot-cred", expires_at=None): + import time as _time + + from litellm.models.credentials import UserConnectionFlowHandle + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + + return encrypt_value_helper( + UserConnectionFlowHandle( + user_id=user_id, + credential_name=credential_name, + device_code="DC-1", + interval=5, + expires_at=expires_at if expires_at is not None else _time.time() + 900, + ).model_dump_json() + ) + + +def test_user_connection_poll_rejects_a_handle_bound_to_another_user(monkeypatch): + _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + response = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": _flow_handle_for(user_id="user-b")}, + auth=_as_user("user-a"), + ) + assert response.status_code == 400 + + +def test_user_connection_poll_rejects_a_handle_bound_to_another_credential(monkeypatch): + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + with _github_http({_DEVICE_CODE_URL: _DEVICE_FLOW_START}): + start = _call_as("POST", "/credentials/copilot-cred/user_connection/start", auth=_as_user()) + assert start.status_code == 200 + forged = _flow_handle_for(credential_name="other-cred") + response = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": forged}, + auth=_as_user(), + ) + assert response.status_code == 400 + assert start.json()["flow_handle"] != forged + table.upsert.assert_not_awaited() + + +def test_user_connection_poll_rejects_an_expired_handle(monkeypatch): + import time as _time + + _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + response = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": _flow_handle_for(expires_at=_time.time() - 1)}, + auth=_as_user(), + ) + assert response.status_code == 400 + + +def test_user_connection_poll_valid_handle_works_with_a_fresh_cache(monkeypatch): + """The handle carries the device code: a poll on a worker whose cache is empty + still completes the flow, no cross-worker state required.""" + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + with _github_http( + { + _DEVICE_CODE_URL: _DEVICE_FLOW_START, + _ACCESS_TOKEN_URL: {"access_token": "gho_1"}, + _COPILOT_TOKEN_URL: {"token": "copilot-tok", "expires_at": 4102444800}, + _GITHUB_USER_URL: {"login": "octo"}, + } + ): + start = _call_as("POST", "/credentials/copilot-cred/user_connection/start", auth=_as_user()) + assert start.status_code == 200, start.text + # simulate a different worker: swap in a brand new empty cache + from litellm.caching.dual_cache import DualCache + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache()) + poll = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": start.json()["flow_handle"]}, + auth=_as_user(), + ) + assert poll.status_code == 200, poll.text + assert poll.json()["status"] == "connected" + table.upsert.assert_awaited_once() + + +def test_delete_user_connection_fails_closed_when_the_tombstone_write_fails(monkeypatch): + """Redis attached but its set() raises: the disconnect must not delete the row, + or the cached token would keep working until TTL with nothing left to revoke.""" + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()], rows=[_connection_row()]) + from types import SimpleNamespace + + from litellm.proxy import proxy_server + + failing_redis = SimpleNamespace( + async_set_cache=AsyncMock(side_effect=ConnectionError("redis down")), + async_get_cache=AsyncMock(return_value=None), + async_delete_cache=AsyncMock(return_value=None), + ) + monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache(redis_cache=failing_redis)) + + response = _call_as("DELETE", "/credentials/copilot-cred/user_connection", auth=_as_user()) + assert response.status_code == 503, response.text + table.delete_many.assert_not_awaited() + + +def test_delete_user_connection_fails_when_the_tombstone_write_silently_noops(monkeypatch): + """RedisCache.async_set_cache swallows client errors, so a set() that returns + without writing still looks successful. The disconnect must verify the + tombstone by reading the key back and refuse to delete the row.""" + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()], rows=[_connection_row()]) + from types import SimpleNamespace + + from litellm.proxy import proxy_server + + store: dict = {} + silent_redis = SimpleNamespace( + async_set_cache=AsyncMock(return_value=None), # reports success, writes nothing + async_get_cache=AsyncMock(side_effect=lambda key: store.get(key)), + async_delete_cache=AsyncMock(return_value=None), + ) + monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache(redis_cache=silent_redis)) + + response = _call_as("DELETE", "/credentials/copilot-cred/user_connection", auth=_as_user()) + assert response.status_code == 503, response.text + table.delete_many.assert_not_awaited() + + +def test_user_connection_poll_reports_503_when_the_cache_overwrite_cannot_land(monkeypatch): + """The connect save is durable, so a stale not-connected marker that survives + both the overwrite and the delete must surface as a retryable 503 rather + than a false "connected".""" + table = _patch_user_connection_env(monkeypatch, [_per_user_credential()]) + from types import SimpleNamespace + + from litellm.proxy import proxy_server + + stale_redis = SimpleNamespace( + async_set_cache=AsyncMock(return_value=None), + async_get_cache=AsyncMock(return_value="stale-value"), + async_delete_cache=AsyncMock(return_value=None), + ) + monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache(redis_cache=stale_redis)) + + with _github_http( + { + _DEVICE_CODE_URL: _DEVICE_FLOW_START, + _ACCESS_TOKEN_URL: {"access_token": "gho_1"}, + _COPILOT_TOKEN_URL: {"token": "copilot-tok", "expires_at": 4102444800}, + _GITHUB_USER_URL: {"login": "octo"}, + } + ): + start = _call_as("POST", "/credentials/copilot-cred/user_connection/start", auth=_as_user()) + poll = _call_as( + "POST", + "/credentials/copilot-cred/user_connection/poll", + json_body={"flow_handle": start.json()["flow_handle"]}, + auth=_as_user(), + ) + assert poll.status_code == 503, poll.text + assert "cache could not be refreshed" in poll.text + table.upsert.assert_awaited_once() + + def _labeled_credential(name: str = "openai-prod", display_name: str | None = "Prod OpenAI") -> CredentialItem: return CredentialItem( credential_name=name, diff --git a/tests/unit/proxy/credential_endpoints/test_user_provider_credentials.py b/tests/unit/proxy/credential_endpoints/test_user_provider_credentials.py new file mode 100644 index 00000000000..b549ecc67a9 --- /dev/null +++ b/tests/unit/proxy/credential_endpoints/test_user_provider_credentials.py @@ -0,0 +1,397 @@ +"""Tests for the per-user provider credentials DB/cache module.""" + +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.caching.dual_cache import DualCache +from litellm.proxy.credential_endpoints.user_provider_credentials import ( + GithubCopilotUserConnectionPayload, + aget_user_provider_tokens, + decode_user_provider_credential, + delete_user_provider_credential, + delete_user_provider_credentials_for_credential, + drop_user_provider_credential_cache, + get_user_provider_credential, + invalidate_user_provider_credential_cache, + set_user_provider_credential_cache, + upsert_user_provider_credential, +) + + +def _prisma(table=None): + prisma_client = MagicMock() + prisma_client.db.litellm_userprovidercredentials = table or MagicMock() + prisma_client.writer_db = prisma_client.db + return prisma_client + + +def _row(user_id="user-a", credential_name="copilot-cred", credential_b64="cipher", provider="github_copilot"): + return SimpleNamespace( + user_id=user_id, + credential_name=credential_name, + credential_b64=credential_b64, + provider=provider, + ) + + +def _fake_redis(monkeypatch): + """A RedisCache backed by a fakeredis client; two DualCaches can share it.""" + import fakeredis + + from litellm.caching.redis_cache import RedisCache + + fake = fakeredis.FakeAsyncRedis() + redis = RedisCache(host="localhost", port=6379) + monkeypatch.setattr(redis, "init_async_client", lambda: fake) + return redis + + +@pytest.fixture(autouse=True) +def _master_key(monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-master") + + +@pytest.mark.asyncio +async def test_upsert_then_get_roundtrips_payload(): + stored = {} + + async def fake_upsert(*, where, data): + stored["row"] = _row( + user_id=where["user_id_credential_name"]["user_id"], + credential_name=where["user_id_credential_name"]["credential_name"], + credential_b64=data["create"]["credential_b64"], + ) + + async def fake_find_unique(*, where): + return stored.get("row") + + table = MagicMock() + table.upsert = AsyncMock(side_effect=fake_upsert) + table.find_unique = AsyncMock(side_effect=fake_find_unique) + prisma_client = _prisma(table) + + await upsert_user_provider_credential( + prisma_client, + "user-a", + "copilot-cred", + "github_copilot", + GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo"), + ) + payload = await get_user_provider_credential(prisma_client, "user-a", "copilot-cred") + assert payload is not None + assert payload.access_token == "gho_secret" + assert payload.github_login == "octo" + # ciphertext at rest, not plaintext + assert "gho_secret" not in stored["row"].credential_b64 + + +@pytest.mark.asyncio +async def test_delete_returns_prior_payload_and_is_idempotent(): + table = MagicMock() + table.delete_many = AsyncMock(return_value=None) + table.find_unique = AsyncMock(return_value=None) + prisma_client = _prisma(table) + assert await delete_user_provider_credential(prisma_client, "user-a", "copilot-cred") is None + table.delete_many.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_delete_for_credential_returns_affected_user_ids(): + table = MagicMock() + table.find_many = AsyncMock(return_value=[_row(user_id="u1"), _row(user_id="u2"), _row(user_id="u1")]) + table.delete_many = AsyncMock(return_value=None) + prisma_client = _prisma(table) + user_ids = await delete_user_provider_credentials_for_credential(prisma_client, "copilot-cred") + assert sorted(user_ids) == ["u1", "u2"] + table.delete_many.assert_awaited_once_with(where={"credential_name": "copilot-cred"}) + + +@pytest.mark.asyncio +async def test_aget_tokens_returns_plaintext_and_caches_ciphertext(monkeypatch): + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + + ciphertext = upc._encode(payload) + table = MagicMock() + table.find_many = AsyncMock(return_value=[_row(credential_b64=ciphertext)]) + prisma_client = _prisma(table) + redis = _fake_redis(monkeypatch) + cache = DualCache(redis_cache=redis) + + tokens = await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) + assert tokens == {"copilot-cred": "gho_secret"} + + # second call served from cache: no new DB read, and cache holds ciphertext not plaintext + tokens2 = await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) + assert tokens2 == {"copilot-cred": "gho_secret"} + assert table.find_many.await_count == 1 + cached = await redis.async_get_cache(upc._cache_key("user-a", "copilot-cred")) + assert cached == ciphertext + assert "gho_secret" not in cached + assert await cache.in_memory_cache.async_get_cache(upc._cache_key("user-a", "copilot-cred")) is None + + +@pytest.mark.asyncio +async def test_aget_tokens_negative_caches_missing_connection(monkeypatch): + table = MagicMock() + table.find_many = AsyncMock(return_value=[]) + prisma_client = _prisma(table) + cache = DualCache(redis_cache=_fake_redis(monkeypatch)) + + assert await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) == {} + assert await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) == {} + assert table.find_many.await_count == 1 + + +@pytest.mark.asyncio +async def test_invalidate_cache_forces_db_refetch(monkeypatch): + payload = GithubCopilotUserConnectionPayload(access_token="gho_new", github_login="octo") + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + + table = MagicMock() + table.find_many = AsyncMock(return_value=[_row(credential_b64=upc._encode(payload))]) + prisma_client = _prisma(table) + cache = DualCache(redis_cache=_fake_redis(monkeypatch)) + + await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) + await drop_user_provider_credential_cache(cache, "user-a", "copilot-cred") + tokens = await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) + assert tokens == {"copilot-cred": "gho_new"} + assert table.find_many.await_count == 2 + + +def test_decode_rejects_garbage(): + assert decode_user_provider_credential("not-a-cipher") is None + + +@pytest.mark.asyncio +async def test_disconnect_on_one_worker_invalidates_other_workers_via_redis(monkeypatch): + """Two DualCache instances (two workers) sharing one Redis: a disconnect on A must + be visible on B, not just in A's in-memory layer.""" + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + + redis = _fake_redis(monkeypatch) + cache_a = DualCache(redis_cache=redis) + cache_b = DualCache(redis_cache=redis) + + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + ciphertext = upc._encode(payload) + table = MagicMock() + table.find_many = AsyncMock(side_effect=[[_row(credential_b64=ciphertext)], []]) + prisma_client = _prisma(table) + + # worker A "connects" (DB read + cache write) + tokens_a = await aget_user_provider_tokens(prisma_client, cache_a, "user-a", ["copilot-cred"]) + assert tokens_a == {"copilot-cred": "gho_secret"} + + # worker B reads the same connection straight from Redis (its in-memory layer is empty) + tokens_b = await aget_user_provider_tokens(prisma_client, cache_b, "user-a", ["copilot-cred"]) + assert tokens_b == {"copilot-cred": "gho_secret"} + + # worker A disconnects: the key is overwritten with the not-connected tombstone + await invalidate_user_provider_credential_cache(cache_a, "user-a", "copilot-cred") + + # worker B sees the tombstone from Redis, no worker-local staleness, no DB read needed + tokens_b_after = await aget_user_provider_tokens(prisma_client, cache_b, "user-a", ["copilot-cred"]) + assert tokens_b_after == {} + assert table.find_many.await_count == 1 + + +@pytest.mark.asyncio +async def test_stale_fill_cannot_overwrite_a_disconnect_tombstone(monkeypatch): + """A token read racing a disconnect: the read's fill uses set-if-absent, so it + must lose to the tombstone the disconnect wrote.""" + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + + redis = _fake_redis(monkeypatch) + cache = DualCache(redis_cache=redis) + + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + ciphertext = upc._encode(payload) + + async def _read_while_disconnecting(*args, **kwargs): + # the row was fetched before the disconnect landed; the fill that follows must + # lose to the tombstone instead of resurrecting the token + await invalidate_user_provider_credential_cache(cache, "user-a", "copilot-cred") + return [_row(credential_b64=ciphertext)] + + table = MagicMock() + table.find_many = AsyncMock(side_effect=_read_while_disconnecting) + prisma_client = _prisma(table) + + await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) + assert await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) == {} + + +@pytest.mark.asyncio +async def test_redis_read_failure_falls_back_to_the_database(monkeypatch): + """A broken Redis must not 401 a connected user: reads degrade to a cache miss.""" + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + + redis = _fake_redis(monkeypatch) + monkeypatch.setattr(redis, "async_get_cache", AsyncMock(side_effect=RuntimeError("redis down"))) + cache = DualCache(redis_cache=redis) + + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + table = MagicMock() + table.find_many = AsyncMock(return_value=[_row(credential_b64=upc._encode(payload))]) + prisma_client = _prisma(table) + + assert await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) == { + "copilot-cred": "gho_secret" + } + + +@pytest.mark.asyncio +async def test_redis_write_and_delete_failures_do_not_fail_the_request(monkeypatch): + redis = _fake_redis(monkeypatch) + monkeypatch.setattr(redis, "async_set_cache", AsyncMock(side_effect=RuntimeError("redis down"))) + monkeypatch.setattr(redis, "async_delete_cache", AsyncMock(side_effect=RuntimeError("redis down"))) + cache = DualCache(redis_cache=redis) + + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + + table = MagicMock() + table.find_many = AsyncMock(return_value=[_row(credential_b64=upc._encode(payload))]) + prisma_client = _prisma(table) + + assert await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) == { + "copilot-cred": "gho_secret" + } + await invalidate_user_provider_credential_cache(cache, "user-a", "copilot-cred") + + +@pytest.mark.asyncio +async def test_without_redis_no_worker_local_entries_are_kept(): + """No Redis -> no cache at all: a disconnect on another worker cannot leave a + stale local entry, because none was ever written.""" + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + ciphertext = upc._encode(payload) + table = MagicMock() + table.find_many = AsyncMock(side_effect=[[_row(credential_b64=ciphertext)], [_row(credential_b64=ciphertext)], []]) + prisma_client = _prisma(table) + cache_a = DualCache() + cache_b = DualCache() + + assert await aget_user_provider_tokens(prisma_client, cache_a, "user-a", ["copilot-cred"]) == { + "copilot-cred": "gho_secret" + } + assert await aget_user_provider_tokens(prisma_client, cache_b, "user-a", ["copilot-cred"]) == { + "copilot-cred": "gho_secret" + } + await invalidate_user_provider_credential_cache(cache_b, "user-a", "copilot-cred") + assert await aget_user_provider_tokens(prisma_client, cache_a, "user-a", ["copilot-cred"]) == {} + assert table.find_many.await_count == 3 + + +@pytest.mark.asyncio +async def test_connect_overwrites_a_not_connected_tombstone(monkeypatch): + """Connect writes the new connection over any cached tombstone so the next + read sees the token immediately, no TTL wait.""" + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + + redis = _fake_redis(monkeypatch) + cache = DualCache(redis_cache=redis) + await invalidate_user_provider_credential_cache(cache, "user-a", "copilot-cred") + + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + assert await set_user_provider_credential_cache(cache, "user-a", "copilot-cred", payload) + + table = MagicMock() + table.find_many = AsyncMock(return_value=[]) + prisma_client = _prisma(table) + assert await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) == { + "copilot-cred": "gho_secret" + } + table.find_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_stale_not_connected_fill_cannot_overwrite_a_connect_write(monkeypatch): + """A read that fetched "not connected" before the poll saved the row fills + with set-if-absent; the connect's plain overwrite must win so the marker + never lingers for the TTL.""" + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + + redis = _fake_redis(monkeypatch) + cache = DualCache(redis_cache=redis) + + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + + async def _read_while_connecting(*args, **kwargs): + assert await set_user_provider_credential_cache(cache, "user-a", "copilot-cred", payload) + return [] + + table = MagicMock() + table.find_many = AsyncMock(side_effect=_read_while_connecting) + prisma_client = _prisma(table) + + await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) + assert await aget_user_provider_tokens(prisma_client, cache, "user-a", ["copilot-cred"]) == { + "copilot-cred": "gho_secret" + } + + +@pytest.mark.asyncio +async def test_set_reports_connected_without_redis(): + """No Redis configured: there is no stale entry to defend against, so the + cache refresh succeeds by definition.""" + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + assert await set_user_provider_credential_cache(DualCache(), "user-a", "copilot-cred", payload) + + +@pytest.mark.asyncio +async def test_set_reports_connected_when_delete_clears_a_dropped_write(): + """async_set_cache swallows errors: a set that reports success but writes + nothing is verified by read-back, and a working delete still resolves the + stale entry.""" + store: dict = {} + silent_redis = SimpleNamespace( + async_set_cache=AsyncMock(return_value=None), + async_get_cache=AsyncMock(side_effect=lambda key: store.get(key)), + async_delete_cache=AsyncMock(side_effect=lambda key: store.pop(key, None)), + ) + cache = DualCache(redis_cache=silent_redis) + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + assert await set_user_provider_credential_cache(cache, "user-a", "copilot-cred", payload) + + +@pytest.mark.asyncio +async def test_set_reports_failure_when_a_stale_entry_survives_both_attempts(): + """A write silently dropped while a not-connected marker stays cached means + the caller cannot trust the connect: the helper reports failure instead of + claiming success.""" + stale = "stale-value" + silent_redis = SimpleNamespace( + async_set_cache=AsyncMock(return_value=None), + async_get_cache=AsyncMock(return_value=stale), + async_delete_cache=AsyncMock(return_value=None), + ) + cache = DualCache(redis_cache=silent_redis) + payload = GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + assert not await set_user_provider_credential_cache(cache, "user-a", "copilot-cred", payload) + + +@pytest.mark.asyncio +async def test_reads_hit_the_writer_engine_not_a_stale_replica(): + """With DATABASE_URL_READ_REPLICA set, prisma_client.db is the reader. A row + saved by connect exists only on the writer, so every lookup must route to + writer_db or a fresh connection 401s as not-connected.""" + reader_table: Final = MagicMock() + reader_table.find_many = AsyncMock(return_value=[]) + writer_table: Final = MagicMock() + writer_table.find_many = AsyncMock(return_value=[]) + prisma_client: Final = MagicMock() + prisma_client.db.litellm_userprovidercredentials = reader_table + prisma_client.writer_db.litellm_userprovidercredentials = writer_table + + await aget_user_provider_tokens(prisma_client, DualCache(), "user-a", ["copilot-cred"]) + + writer_table.find_many.assert_awaited_once() + reader_table.find_many.assert_not_awaited() diff --git a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py index dbb2f16d056..beffe454b8d 100644 --- a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py @@ -9581,3 +9581,188 @@ class TestFederationGateScopesToWhatTheWriteTouches: assert "deleted successfully" in result["message"] mock_prisma.db.litellm_proxymodeltable.delete.assert_awaited_once() + + +class TestNonAdminCannotSetPerUserOauthOnModel: + """github_copilot_auth_type is a server-owned WIF field; a team admin must not write it + onto a deployment's litellm_params nor attach a per-user credential by name.""" + + @pytest.mark.asyncio + async def test_add_new_model_non_admin_cannot_set_github_copilot_auth_type(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) + mock_prisma = MagicMock() + mock_prisma.writer_db.litellm_uisettings.find_unique = AsyncMock(return_value=None) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", mock_prisma + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", True + ), + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params( + model="github_copilot/gpt-4o", + github_copilot_auth_type="per_user_oauth", + ), + model_info={"id": "wif-gate-copilot-1"}, + ), + user_api_key_dict=non_admin, + ) + assert "proxy admin" in str(exc_info.value.message).lower() + assert exc_info.value.param == "github_copilot_auth_type" + mock_prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_add_new_model_admin_can_set_github_copilot_auth_type(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + mock_prisma = MagicMock() + created_row = MagicMock() + created_row.model_id = "wif-gate-copilot-2" + created_row.model_dump_json.return_value = "{}" + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=created_row) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", mock_prisma + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ), + 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": ["wif-gate-copilot-2"]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.proxy_config", + MagicMock(add_deployment=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None))), + ), + ): + result = await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params( + model="github_copilot/gpt-4o", + github_copilot_auth_type="per_user_oauth", + ), + model_info={"id": "wif-gate-copilot-2"}, + ), + user_api_key_dict=admin, + ) + assert result is created_row + + @pytest.mark.asyncio + async def test_add_new_model_non_admin_cannot_attach_a_per_user_oauth_credential(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) + mock_prisma = MagicMock() + mock_prisma.writer_db.litellm_uisettings.find_unique = AsyncMock(return_value=None) + per_user_credential_row = MagicMock() + per_user_credential_row.credential_values = {"github_copilot_auth_type": "per_user_oauth"} + mock_prisma.db.litellm_credentialstable.find_unique = AsyncMock( + return_value=per_user_credential_row + ) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", mock_prisma + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", True + ), + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params( + model="github_copilot/gpt-4o", + litellm_credential_name="admin-copilot-cred", + ), + model_info={"id": "wif-gate-copilot-3"}, + ), + user_api_key_dict=non_admin, + ) + mock_prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_add_new_model_admin_can_attach_a_per_user_oauth_credential(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + mock_prisma = MagicMock() + created_row = MagicMock() + created_row.model_id = "wif-gate-copilot-4" + created_row.model_dump_json.return_value = "{}" + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=created_row) + per_user_credential_row = MagicMock() + per_user_credential_row.credential_values = {"github_copilot_auth_type": "per_user_oauth"} + mock_prisma.db.litellm_credentialstable.find_unique = AsyncMock( + return_value=per_user_credential_row + ) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", mock_prisma + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ), + 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": ["wif-gate-copilot-4"]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.proxy_config", + MagicMock(add_deployment=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None))), + ), + ): + result = await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params( + model="github_copilot/gpt-4o", + litellm_credential_name="admin-copilot-cred", + ), + model_info={"id": "wif-gate-copilot-4"}, + ), + user_api_key_dict=admin, + ) + assert result is created_row diff --git a/tests/unit/proxy/test_litellm_pre_call_utils.py b/tests/unit/proxy/test_litellm_pre_call_utils.py index cbf06319d13..2ba736d8ae0 100644 --- a/tests/unit/proxy/test_litellm_pre_call_utils.py +++ b/tests/unit/proxy/test_litellm_pre_call_utils.py @@ -16,16 +16,31 @@ from pydantic import ValidationError as PydanticValidationError from starlette.datastructures import Headers import litellm -from litellm.constants import SERVER_STREAMING_CLASSIFICATION_KEY, SERVER_STREAMING_CLASSIFICATION_MARKER +from litellm.constants import ( + ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY, + SERVER_STREAMING_CLASSIFICATION_KEY, + SERVER_STREAMING_CLASSIFICATION_MARKER, + SESSION_ID_GENERATED_METADATA_KEY, + SESSION_ID_OMITTED_METADATA_KEY, +) +from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs +from litellm.litellm_core_utils.get_provider_specific_headers import ( + ProviderSpecificHeaderUtils, +) +from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( + TRUSTED_CALLBACK_VARS_FIELD, +) +from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY +from litellm.litellm_core_utils.redact_messages import _get_turn_off_message_logging_from_dynamic_params +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.fireworks_ai.common_utils import get_fireworks_session_id from litellm.proxy._types import AddTeamCallback, ProxyException, TeamCallbackMetadata, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import ( KeyAndTeamLoggingSettings, LiteLLMProxyRequestSetup, _apply_credential_overrides_from_model_config, _extract_credential_from_entry, - get_dynamic_logging_metadata, _get_enforced_params, - get_metadata_variable_name, _match_and_track_policies, _promoted_trace_control_fields, _resolve_credential_from_model_config, @@ -36,25 +51,11 @@ from litellm.proxy.litellm_pre_call_utils import ( add_provider_specific_headers_to_request, check_if_token_is_service_account, clean_headers, + get_dynamic_logging_metadata, + get_metadata_variable_name, move_guardrails_to_metadata, ) -from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload -from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY -from litellm.litellm_core_utils.redact_messages import _get_turn_off_message_logging_from_dynamic_params -from litellm.litellm_core_utils.get_provider_specific_headers import ( - ProviderSpecificHeaderUtils, -) -from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( - TRUSTED_CALLBACK_VARS_FIELD, -) -from litellm.constants import ( - ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY, - SESSION_ID_GENERATED_METADATA_KEY, - SESSION_ID_OMITTED_METADATA_KEY, -) -from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM -from litellm.llms.fireworks_ai.common_utils import get_fireworks_session_id from litellm.types.utils import CredentialItem @@ -3029,8 +3030,6 @@ def test_add_headers_to_llm_call_by_model_group_existing_headers_in_data(): litellm.model_group_settings = original_model_group_settings -from typing import Optional - from fastapi.responses import Response from litellm.integrations.custom_logger import CustomLogger @@ -3041,7 +3040,7 @@ from litellm.types.utils import StandardLoggingPayload class TestCustomLogger(CustomLogger): def __init__(self): - self.standard_logging_object: Optional[StandardLoggingPayload] = None + self.standard_logging_object: StandardLoggingPayload | None = None super().__init__() async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -5599,9 +5598,9 @@ async def test_team_guardrail_merges_with_global_policy(): shadowed and non-default guardrails silently received an empty requested_guardrails list. """ + from litellm.proxy.litellm_pre_call_utils import move_guardrails_to_metadata from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry from litellm.proxy.policy_engine.policy_registry import get_policy_registry - from litellm.proxy.litellm_pre_call_utils import move_guardrails_to_metadata from litellm.types.proxy.policy_engine import ( Policy, PolicyAttachment, @@ -8992,6 +8991,8 @@ def test_body_snapshot_drops_only_the_marker_and_keeps_caller_value(marker: str) refresh_proxy_server_request_body_snapshot(caller_data) assert caller_request["body"][SERVER_STREAMING_CLASSIFICATION_KEY] == "caller-value", caller_request + + def test_arize_otlp_protocol_on_a_key_logging_entry_reaches_the_destination(monkeypatch): from litellm.integrations.otel.model.config import is_otel_v2_enabled from litellm.proxy.litellm_pre_call_utils import resolve_tenant_otel_destinations @@ -9025,3 +9026,1123 @@ def test_arize_otlp_protocol_on_a_key_logging_entry_reaches_the_destination(monk assert destinations[0].endpoint == "https://arize.internal.example/v1/traces" finally: is_otel_v2_enabled.cache_clear() + + +class TestResolveUserProviderCredentials: + """``_resolve_user_provider_credentials_for_request`` injects the calling user's stored + provider tokens into secret_fields for deployments carrying a per_user_oauth credential.""" + + def _router(self): + from litellm import Router + + return Router( + model_list=[ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "openai/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + }, + { + "model_name": "shared-chat", + "litellm_params": { + "model": "openai/gpt-4o", + "litellm_credential_name": "shared-cred", + }, + }, + ], + fallbacks=[{"copilot-chat": ["shared-chat"]}, {"*": ["copilot-chat"]}], + ) + + def _env(self, monkeypatch): + from litellm.caching.dual_cache import DualCache + + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="copilot-cred", + credential_values={"github_copilot_auth_type": "per_user_oauth"}, + credential_info={}, + ), + CredentialItem(credential_name="shared-cred", credential_values={"api_key": "k"}, credential_info={}), + ], + ) + prisma_client = MagicMock() + table = MagicMock() + table.find_many = AsyncMock(return_value=[]) + prisma_client.db.litellm_userprovidercredentials = table + prisma_client.writer_db = prisma_client.db + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", DualCache()) + return table + + @pytest.mark.asyncio + async def test_injects_tokens_for_per_user_deployments_only(self, monkeypatch): + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-master") + table = self._env(monkeypatch) + ciphertext = upc._encode(upc.GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo")) + row = SimpleNamespace(user_id="user-a", credential_name="copilot-cred", credential_b64=ciphertext) + table.find_many = AsyncMock(return_value=[row]) + + data = {"model": "copilot-chat", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=self._router(), + ) + + assert data["secret_fields"]["user_provider_credentials"] == {"copilot-cred": "gho_secret"} + assert data["secret_fields"]["user_provider_credentials_user_id"] == "user-a" + # the shared fallback credential is never looked up + call_where = table.find_many.await_args.kwargs["where"] + assert call_where == {"AND": ({"user_id": "user-a"}, {"credential_name": {"in": ["copilot-cred"]}})} + + @pytest.mark.asyncio + async def test_service_key_cannot_override_a_per_user_credential(self, monkeypatch): + """Keys with no user_id must still hit the body check: discovery ran before the + missing-user return, so a litellm_credential_name override is a 400.""" + import pytest + from fastapi import HTTPException + + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + self._env(monkeypatch) + router = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock( + return_value=[ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + ) + router.get_deployment = MagicMock(return_value=None) + data = {"model": "copilot-chat", "litellm_credential_name": "shared-cred", "secret_fields": {}} + with pytest.raises(HTTPException) as exc: + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id=None, + team_id=None, + llm_router=router, + ) + assert getattr(exc.value, "status_code", None) == 400 + + @pytest.mark.asyncio + async def test_service_key_without_override_leaves_data_unchanged(self, monkeypatch): + """Pin: no user_id and no override -> the function is a no-op.""" + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + self._env(monkeypatch) + router = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock( + return_value=[ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + ) + router.get_deployment = MagicMock(return_value=None) + data = {"model": "copilot-chat", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id=None, + team_id=None, + llm_router=router, + ) + assert data == {"model": "copilot-chat", "secret_fields": {}} + + @pytest.mark.asyncio + async def test_request_level_fallback_loads_per_user_tokens(self, monkeypatch): + """A request-level fallback naming a per-user group must surface its credential.""" + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table = self._env(monkeypatch) + router = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + + def _get_model_list(model_name=None, team_id=None): + if model_name != "copilot-chat": + return [] + return [ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + + router.get_model_list = MagicMock(side_effect=_get_model_list) + router.get_deployment = MagicMock(return_value=None) + data = { + "model": "gpt-4o", + "fallbacks": [{"gpt-4o": ["copilot-chat"]}], + "secret_fields": {}, + } + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + table.find_many.assert_awaited() + assert data["secret_fields"]["user_provider_credentials_user_id"] == "user-a" + + @pytest.mark.asyncio + async def test_request_fallback_chain_loads_transitive_per_user_tokens(self, monkeypatch): + """A -> B -> C through request-level fallback dicts: the transitive target's + per-user credential must be discovered too.""" + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table = self._env(monkeypatch) + router = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + + def _get_model_list(model_name=None, team_id=None): + if model_name != "copilot-chat": + return [] + return [ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + + router.get_model_list = MagicMock(side_effect=_get_model_list) + router.get_deployment = MagicMock(return_value=None) + data = { + "model": "gpt-4o", + "fallbacks": [{"gpt-4o": ["backup-chat"]}, {"backup-chat": ["copilot-chat"]}], + "secret_fields": {}, + } + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + table.find_many.assert_awaited() + assert data["secret_fields"]["user_provider_credentials_user_id"] == "user-a" + + @pytest.mark.asyncio + async def test_key_router_settings_fallback_loads_per_user_tokens(self, monkeypatch): + """The proxy prefers key/team router_settings.fallbacks over router-level + fallbacks (_configured_fallbacks), so discovery must read them too.""" + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table = self._env(monkeypatch) + router = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock( + side_effect=lambda model_name=None, team_id=None: ( + [ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + if model_name == "copilot-chat" + else [] + ) + ) + router.get_deployment = MagicMock(return_value=None) + data = {"model": "gpt-4o", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + router_settings={"fallbacks": [{"gpt-4o": [{"model": "copilot-chat"}]}]}, + ) + table.find_many.assert_awaited() + assert data["secret_fields"]["user_provider_credentials_user_id"] == "user-a" + + @pytest.mark.asyncio + async def test_fallback_discovery_bounds_matcher_calls_on_huge_lists(self, monkeypatch): + """30k admin-configured unmatched entries go through the exact index: the + router matcher only ever sees the small wildcard subset, once per group + per list, instead of being scored against every entry.""" + import litellm.router_utils.fallback_event_handlers as feh + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + self._env(monkeypatch) + entries_seen: Final[list[int]] = [] + real_matcher: Final = feh.get_fallback_model_group + + def _spy(fallbacks, model_group): + entries_seen.append(len(fallbacks)) + return real_matcher(fallbacks=fallbacks, model_group=model_group) + + monkeypatch.setattr(feh, "get_fallback_model_group", _spy) + router: Final = MagicMock() + router.fallbacks = [{"unmatched-%d" % i: ["x"]} for i in range(30_000)] + [{"*": ["generic-x"]}] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock(return_value=[]) + router.get_deployment = MagicMock(return_value=None) + data: Final = {"model": "gpt-4o", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + total: Final = sum(entries_seen) + assert total <= 8, ( + f"matcher must only see the 1 wildcard entry, got {total} entries across {len(entries_seen)} calls" + ) + + @pytest.mark.asyncio + async def test_duplicate_fallback_source_keys_union_targets(self, monkeypatch): + """Two entries with the same source key: the router picks the first, so + the per-user target on that first entry must be discovered even though + a later duplicate points elsewhere.""" + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table: Final = self._env(monkeypatch) + router: Final = MagicMock() + router.fallbacks = [{"gpt-4o": ["copilot-chat"]}, {"gpt-4o": ["other-chat"]}] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock( + side_effect=lambda model_name=None, team_id=None: ( + [ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + if model_name == "copilot-chat" + else [] + ) + ) + router.get_deployment = MagicMock(return_value=None) + data: Final = {"model": "gpt-4o", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + table.find_many.assert_awaited() + assert data["secret_fields"]["user_provider_credentials_user_id"] == "user-a" + + @pytest.mark.asyncio + async def test_fallback_discovery_scans_past_the_wildcard_cap_position(self, monkeypatch): + """A wildcard match at position 290 of an admin list is still discovered: + admin lists are scanned in full, only the request body's lists are + capped.""" + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table: Final = self._env(monkeypatch) + router: Final = MagicMock() + entries: Final = [{"azure/unmatched-%d" % i: ["x"]} for i in range(300)] + entries.insert(290, {"*": ["copilot-chat"]}) + router.fallbacks = entries + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock( + side_effect=lambda model_name=None, team_id=None: ( + [ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + if model_name == "copilot-chat" + else [] + ) + ) + router.get_deployment = MagicMock(return_value=None) + data: Final = {"model": "gpt-4o", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + table.find_many.assert_awaited() + assert data["secret_fields"]["user_provider_credentials_user_id"] == "user-a" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "entries", + [ + [{"gpt-4o": ["target-%d" % i for i in range(257)]}], + [{"azure/unmatched-%d" % i: []} for i in range(257)], + [{("k%d" % i): [] for i in range(257)}], + [{"m%d" % i: "t"} for i in range(257)], + ], + ids=["nested targets", "empty-list mappings", "one dict of empty lists", "scalar mappings"], + ) + async def test_request_fallback_entries_over_the_limit_get_a_400(self, monkeypatch, entries): + """The request body's fallback lists are the caller-controlled input to + discovery, so past 256 counted work units (a key counts even when its + value list is empty or a scalar) the request is refused instead of scanned.""" + from fastapi import HTTPException + + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + self._env(monkeypatch) + router: Final = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + data: Final = {"model": "gpt-4o", "fallbacks": entries, "secret_fields": {}} + with pytest.raises(HTTPException) as exc: + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + assert getattr(exc.value, "status_code", None) == 400 + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "entries", + [ + [{"gpt-4o": ["target-%d" % i for i in range(256)]}], + [{"azure/unmatched-%d" % i: []} for i in range(256)], + [{("k%d" % i): [] for i in range(256)}], + [{"m%d" % i: "t"} for i in range(256)], + ], + ids=["nested targets", "empty-list mappings", "one dict of empty lists", "scalar mappings"], + ) + async def test_request_fallback_entries_at_the_limit_are_accepted(self, monkeypatch, entries): + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + self._env(monkeypatch) + router: Final = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock(return_value=[]) + router.get_deployment = MagicMock(return_value=None) + data: Final = {"model": "gpt-4o", "fallbacks": entries, "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + assert data["secret_fields"] == {} + + @pytest.mark.asyncio + async def test_shared_only_credentials_skip_discovery_and_the_limit(self, monkeypatch): + """With no per-user credential configured the request's fallback lists + are never scanned and the size limit never fires: a proxy that never + opted into per-user OAuth keeps byte-identical behavior.""" + import litellm.router_utils.fallback_event_handlers as feh + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + self._env(monkeypatch) + monkeypatch.setattr( + litellm, + "credential_list", + [CredentialItem(credential_name="shared-cred", credential_values={"api_key": "k"}, credential_info={})], + ) + matcher_calls: Final[list[object]] = [] + real_matcher: Final = feh.get_fallback_model_group + + def _spy(fallbacks, model_group): + matcher_calls.append(model_group) + return real_matcher(fallbacks=fallbacks, model_group=model_group) + + monkeypatch.setattr(feh, "get_fallback_model_group", _spy) + router: Final = MagicMock() + router.fallbacks = [{"gpt-4o": ["copilot-chat"]}] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock(return_value=[]) + router.get_deployment = MagicMock(return_value=None) + data: Final = { + "model": "gpt-4o", + "fallbacks": [{"unmatched-%d" % i: ["x"]} for i in range(1000)], + "secret_fields": {}, + } + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + assert matcher_calls == [] + router.get_model_list.assert_not_called() + assert data["secret_fields"] == {} + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "auto_router_model", + [ + "auto_router/copilot-picker", + "auto_router/complexity_router", + "auto_router/adaptive_router", + "auto_router/quality_router", + ], + ids=["semantic", "complexity", "adaptive", "quality"], + ) + async def test_auto_router_deployment_loads_every_per_user_connection(self, monkeypatch, auto_router_model: str): + """An auto-router deployment can pick a Copilot group at request time, so + discovery must fall back to every per-user credential the caller has.""" + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-master") + table: Final = self._env(monkeypatch) + ciphertext: Final = upc._encode( + upc.GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo") + ) + table.find_many = AsyncMock( + return_value=[SimpleNamespace(user_id="user-a", credential_name="copilot-cred", credential_b64=ciphertext)] + ) + router: Final = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock( + return_value=[ + { + "model_name": "auto-chat", + "litellm_params": {"model": auto_router_model}, + } + ] + ) + router.get_deployment = MagicMock(return_value=None) + data: Final = {"model": "auto-chat", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + assert data["secret_fields"]["user_provider_credentials"] == {"copilot-cred": "gho_secret"} + call_where: Final = table.find_many.await_args.kwargs["where"] + assert call_where == {"AND": ({"user_id": "user-a"}, {"credential_name": {"in": ["copilot-cred"]}})} + + @pytest.mark.asyncio + async def test_non_per_user_group_loads_no_credentials(self, monkeypatch): + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table: Final = self._env(monkeypatch) + router: Final = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock( + return_value=[ + { + "model_name": "shared-chat", + "litellm_params": {"model": "github_copilot/gpt-4o", "litellm_credential_name": "shared-cred"}, + } + ] + ) + router.get_deployment = MagicMock(return_value=None) + data: Final = {"model": "shared-chat", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + table.find_many.assert_not_awaited() + assert data["secret_fields"] == {} + + @pytest.mark.asyncio + async def test_wildcard_targets_expand_once_not_per_group(self, monkeypatch): + """An admin ``{"*": [5000 targets]}`` chain must expand the entry once: + without fired-entry tracking each discovered group would re-expand the + same 5000 targets.""" + from litellm.proxy import litellm_pre_call_utils as lpc + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + self._env(monkeypatch) + target_visits: Final[list[int]] = [] + real_edges: Final = lpc._fallback_edges + + def _counting_edges(indexed_lists, model_group, fired_fuzzy): + result: Final = real_edges(indexed_lists, model_group, fired_fuzzy) + target_visits.append(len(result)) + return result + + monkeypatch.setattr(lpc, "_fallback_edges", _counting_edges) + router: Final = MagicMock() + router.fallbacks = [{"*": ["target-%d" % i for i in range(4999)] + ["copilot-chat"]}] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock( + side_effect=lambda model_name=None, team_id=None: ( + [ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + if model_name == "copilot-chat" + else [] + ) + ) + router.get_deployment = MagicMock(return_value=None) + data: Final = {"model": "gpt-4o", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + total: Final = sum(target_visits) + assert total <= 5100, f"targets must expand once, got {total} visits" + + @pytest.mark.asyncio + async def test_model_alias_map_target_loads_per_user_tokens(self, monkeypatch): + """litellm.model_alias_map rewrites data['model'] after pre-call discovery: + the alias target's per-user credential must be discovered on the rewritten + name the router actually receives.""" + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table = self._env(monkeypatch) + monkeypatch.setattr(litellm, "model_alias_map", {"chat-alias": "copilot-chat"}) + router = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock( + side_effect=lambda model_name=None, team_id=None: ( + [ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + if model_name == "copilot-chat" + else [] + ) + ) + router.get_deployment = MagicMock(return_value=None) + data = {"model": "chat-alias", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + table.find_many.assert_awaited() + assert data["secret_fields"]["user_provider_credentials_user_id"] == "user-a" + + @pytest.mark.asyncio + async def test_team_router_settings_fallback_loads_per_user_tokens(self, monkeypatch): + """Team router_settings fallbacks (hierarchical Key > Team source in + _get_hierarchical_router_settings) must be part of discovery.""" + from types import SimpleNamespace + + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + import litellm.proxy.auth.auth_checks as auth_checks + + table = self._env(monkeypatch) + monkeypatch.setattr( + auth_checks, + "get_team_object", + AsyncMock(return_value=SimpleNamespace(router_settings={"fallbacks": [{"gpt-4o": ["copilot-chat"]}]})), + ) + router = MagicMock() + router.fallbacks = [] + router.context_window_fallbacks = [] + router.content_policy_fallbacks = [] + router.get_model_list = MagicMock( + side_effect=lambda model_name=None, team_id=None: ( + [ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + if model_name == "copilot-chat" + else [] + ) + ) + router.get_deployment = MagicMock(return_value=None) + data = {"model": "gpt-4o", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id="team-1", + llm_router=router, + ) + table.find_many.assert_awaited() + assert data["secret_fields"]["user_provider_credentials_user_id"] == "user-a" + + @pytest.mark.asyncio + async def test_team_public_model_name_resolves_per_user_deployments(self, monkeypatch): + """A caller hitting a team's public model name must get the same per-user + credentials routing resolves: get_model_list needs the caller's team_id.""" + from litellm import Router + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table = self._env(monkeypatch) + router = Router( + model_list=[ + { + "model_name": "copilot-chat_team-1_dep", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + "model_info": { + "id": "dep-abc", + "team_id": "team-1", + "team_public_model_name": "copilot-chat", + }, + } + ] + ) + data = {"model": "copilot-chat", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id="team-1", + llm_router=router, + ) + table.find_many.assert_awaited() + call_where = table.find_many.await_args.kwargs["where"] + assert call_where == {"AND": ({"user_id": "user-a"}, {"credential_name": {"in": ["copilot-cred"]}})} + + @pytest.mark.asyncio + async def test_deployment_id_resolves_per_user_deployment(self, monkeypatch): + """A request addressed by deployment id still discovers the per-user credential.""" + from litellm import Router + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table = self._env(monkeypatch) + router = Router( + model_list=[ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + "model_info": {"id": "dep-xyz"}, + } + ] + ) + data = {"model": "dep-xyz", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + table.find_many.assert_awaited() + + @pytest.mark.asyncio + async def test_group_with_no_per_user_deployment_leaves_secret_fields_untouched(self, monkeypatch): + from litellm import Router + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + self._env(monkeypatch) + data = {"model": "shared-chat", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=Router( + model_list=[ + { + "model_name": "shared-chat", + "litellm_params": { + "model": "openai/gpt-4o", + "litellm_credential_name": "shared-cred", + }, + } + ] + ), + ) + assert "user_provider_credentials" not in data["secret_fields"] + + @pytest.mark.asyncio + async def test_unconnected_user_gets_empty_map_not_an_error(self, monkeypatch): + """A wildcard fallback to a per-user group is legitimate coverage; a user with no + connection gets an empty token map and the provider-side helper fails the call later.""" + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + self._env(monkeypatch) + data = {"model": "shared-chat", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=self._router(), + ) + assert data["secret_fields"]["user_provider_credentials"] == {} + + @pytest.mark.asyncio + async def test_missing_user_id_does_nothing(self, monkeypatch): + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + self._env(monkeypatch) + data = {"model": "copilot-chat", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id=None, + team_id=None, + llm_router=self._router(), + ) + assert "user_provider_credentials" not in data["secret_fields"] + + @pytest.mark.asyncio + async def test_unrelated_model_makes_zero_db_and_cache_calls(self, monkeypatch): + from litellm import Router + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table = self._env(monkeypatch) + data = {"model": "unrelated", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=Router( + model_list=[ + { + "model_name": "unrelated", + "litellm_params": {"model": "openai/gpt-4o"}, + } + ] + ), + ) + table.find_many.assert_not_awaited() + assert "user_provider_credentials" not in data["secret_fields"] + + @pytest.mark.asyncio + async def test_per_user_deployment_with_pydantic_litellm_params(self, monkeypatch): + """Router deployments may carry litellm_params as a LiteLLM_Params model rather + than a dict; the per-user credential must still be collected.""" + from litellm.proxy.litellm_pre_call_utils import _per_user_credential_names_for_groups + from litellm.types.router import LiteLLM_Params + + self._env(monkeypatch) + router = MagicMock() + router.get_model_list = MagicMock( + return_value=[ + { + "model_name": "copilot-chat", + "litellm_params": LiteLLM_Params( + model="github_copilot/gpt-4o", + litellm_credential_name="copilot-cred", + ), + } + ] + ) + router.get_deployment = MagicMock(return_value=None) + assert _per_user_credential_names_for_groups(router, frozenset({"copilot-chat"}), team_id=None) == ( + "copilot-cred", + ) + + +def _per_user_env(monkeypatch, credential_names=("copilot-cred",)): + from litellm.caching.dual_cache import DualCache + + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name=name, + credential_values={"github_copilot_auth_type": "per_user_oauth"}, + credential_info={}, + ) + for name in credential_names + ], + ) + prisma_client = MagicMock() + table = MagicMock() + table.find_many = AsyncMock(return_value=[]) + prisma_client.db.litellm_userprovidercredentials = table + prisma_client.writer_db = prisma_client.db + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", DualCache()) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-master") + return table + + +@pytest.mark.asyncio +@pytest.mark.parametrize("body_value", ["", None, "shared-cred"]) +async def test_request_body_litellm_credential_name_is_rejected_for_per_user_models(monkeypatch, body_value): + """A caller-supplied litellm_credential_name (empty, null, or another credential) + would downgrade a per-user deployment to the shared device login; it is a 400.""" + from fastapi import HTTPException + + from litellm import Router + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + _per_user_env(monkeypatch) + router = Router( + model_list=[ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + ) + data = {"model": "copilot-chat", "secret_fields": {}, "litellm_credential_name": body_value} + with pytest.raises(HTTPException) as exc: + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + assert exc.value.status_code == 400 + assert "litellm_credential_name cannot be set in the request body" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_request_body_litellm_credential_name_untouched_for_shared_models(monkeypatch): + """Regression: models with no per-user credential keep accepting the param.""" + from litellm import Router + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + _per_user_env(monkeypatch) + router = Router( + model_list=[ + { + "model_name": "shared-chat", + "litellm_params": {"model": "openai/gpt-4o", "litellm_credential_name": "shared-cred"}, + } + ] + ) + data = {"model": "shared-chat", "secret_fields": {}, "litellm_credential_name": "shared-cred"} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + assert "user_provider_credentials" not in data["secret_fields"] + + +@pytest.mark.asyncio +async def test_transitive_fallback_chain_loads_per_user_credential(monkeypatch): + """A -> B -> C where only C is per-user: the request to A must still load the + caller's connection for C's credential.""" + from litellm import Router + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + + table = _per_user_env(monkeypatch) + router = Router( + model_list=[ + {"model_name": "a", "litellm_params": {"model": "openai/gpt-4o"}}, + {"model_name": "b", "litellm_params": {"model": "openai/gpt-4o"}}, + { + "model_name": "c", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + }, + ], + fallbacks=[{"a": ["b"]}, {"b": ["c"]}], + ) + data = {"model": "a", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + where = table.find_many.await_args.kwargs["where"] + assert where == {"AND": ({"user_id": "user-a"}, {"credential_name": {"in": ["copilot-cred"]}})} + assert data["secret_fields"]["user_provider_credentials_user_id"] == "user-a" + + +@pytest.mark.asyncio +async def test_secret_fields_user_provider_credentials_repr_hides_tokens(monkeypatch): + """Debug logging prints the request body; the credential map must be a + RedactedDict so a GitHub token can never appear in str()/repr() output.""" + from litellm import Router + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + from litellm.proxy.litellm_pre_call_utils import _resolve_user_provider_credentials_for_request + from litellm.types.proxy.litellm_pre_call_utils import RedactedDict + + table = _per_user_env(monkeypatch) + ciphertext = upc._encode(upc.GithubCopilotUserConnectionPayload(access_token="gho_secret", github_login="octo")) + table.find_many = AsyncMock( + return_value=[SimpleNamespace(user_id="user-a", credential_name="copilot-cred", credential_b64=ciphertext)] + ) + router = Router( + model_list=[ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + ) + data = {"model": "copilot-chat", "secret_fields": {}} + await _resolve_user_provider_credentials_for_request( + data=data, + authenticated_user_id="user-a", + team_id=None, + llm_router=router, + ) + credentials = data["secret_fields"]["user_provider_credentials"] + assert isinstance(credentials, RedactedDict) + assert "gho_secret" not in str(data) + assert "gho_secret" not in repr(credentials) + assert credentials["copilot-cred"] == "gho_secret", "provider reads still get the token" + + +@pytest.mark.asyncio +async def test_header_mapped_user_id_cannot_load_another_users_connection(monkeypatch): + """general_settings.user_header_mappings lets a caller rename their user_id through a + request header. Per-user token resolution is bound to the key-authenticated identity, + so a header naming user B still loads only user A's connections.""" + from litellm import Router + from litellm.caching.dual_cache import DualCache + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + _per_user_env(monkeypatch) + load = AsyncMock(return_value={"copilot-cred": "gho_user_a_token"}) + monkeypatch.setattr(upc, "aget_user_provider_tokens", load) + router = Router( + model_list=[ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", DualCache()) + + request_mock = MagicMock(spec=Request) + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json", "X-OpenWebUI-User-Id": "user-b"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = {"model": "copilot-chat"} + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", user_id="user-a"), + proxy_config=MagicMock(), + general_settings={ + "user_header_mappings": [{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}] + }, + version="test-version", + ) + + called_user_ids = [call.kwargs["user_id"] for call in load.await_args_list] + assert called_user_ids == ["user-a"], f"header-mapped user-b must never reach the token loader: {called_user_ids}" + credentials = updated["secret_fields"]["user_provider_credentials"] + assert credentials["copilot-cred"] == "gho_user_a_token" + assert updated["secret_fields"]["user_provider_credentials_user_id"] == "user-a" + + +@pytest.mark.asyncio +async def test_per_user_token_resolution_uses_key_identity_without_header_mapping(monkeypatch): + """Control: with no user_header_mappings configured the key's user_id is used as before.""" + from litellm import Router + from litellm.caching.dual_cache import DualCache + from litellm.proxy.credential_endpoints import user_provider_credentials as upc + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + _per_user_env(monkeypatch) + load = AsyncMock(return_value={"copilot-cred": "gho_user_a_token"}) + monkeypatch.setattr(upc, "aget_user_provider_tokens", load) + router = Router( + model_list=[ + { + "model_name": "copilot-chat", + "litellm_params": { + "model": "github_copilot/gpt-4o", + "litellm_credential_name": "copilot-cred", + }, + } + ] + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", DualCache()) + + request_mock = MagicMock(spec=Request) + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json", "X-OpenWebUI-User-Id": "user-b"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = {"model": "copilot-chat"} + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", user_id="user-a"), + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + called_user_ids = [call.kwargs["user_id"] for call in load.await_args_list] + assert called_user_ids == ["user-a"] + assert updated["secret_fields"]["user_provider_credentials"]["copilot-cred"] == "gho_user_a_token" diff --git a/tests/unit/router_utils/test_cooldown_handlers.py b/tests/unit/router_utils/test_cooldown_handlers.py index d219abb7d3d..3a5097d5b1f 100644 --- a/tests/unit/router_utils/test_cooldown_handlers.py +++ b/tests/unit/router_utils/test_cooldown_handlers.py @@ -2,6 +2,7 @@ import asyncio import importlib import random import time +from collections.abc import Mapping from typing import Final from unittest.mock import MagicMock, patch @@ -577,6 +578,7 @@ def _vcr_outcome_gate(request, vcr): yield record_vcr_outcome(request, vcr) + @pytest.fixture(scope="function") def setup_and_teardown(): """ @@ -597,9 +599,11 @@ def setup_and_teardown(): loop.close() asyncio.set_event_loop(None) + def _make_router(model_list: list, **kwargs) -> Router: return Router(model_list=model_list, **kwargs) + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestDeploymentLevelAllowedFails: def test_deployment_level_allowed_fails_overrides_router_level(self): @@ -672,6 +676,7 @@ class TestDeploymentLevelAllowedFails: "secondary has no deployment-level policy; with allowed_fails=10 it should not cool down on first failure" ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestDeploymentLevelAllowedFailsPolicyByExceptionType: def test_rate_limit_error_triggers_cooldown_with_zero_threshold(self): @@ -746,6 +751,7 @@ class TestDeploymentLevelAllowedFailsPolicyByExceptionType: ) assert should_cooldown is True, "Should cooldown after exceeding InternalServerErrorAllowedFails=5" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestExceptionTypeCountersTrackedIndependently: def test_cache_key_suffix_separates_exception_type_counters(self): @@ -796,6 +802,7 @@ class TestExceptionTypeCountersTrackedIndependently: assert generic_counter_after == 1, "generic counter should now be 1" assert rl_counter_after == 3, "RateLimitError counter must remain unchanged after InternalServerError" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestCooldownCacheTTLCorrection: def _make_cooldown_cache(self) -> CooldownCache: @@ -897,6 +904,7 @@ class TestCooldownCacheTTLCorrection: assert active == [], "Expired entry must not appear in async active cooldowns" assert cc.in_memory_cache.get_cache(key) is None + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestFallbackDeploymentCooldown: def test_trigger_cooldown_for_failed_deployment_calls_set_cooldown(self): @@ -1087,6 +1095,7 @@ class TestFallbackDeploymentCooldown: call_kwargs = mock_set_cooldown.call_args[1] assert call_kwargs["time_to_cooldown"] == 15.0, "model_info.cooldown_time must take priority" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestSingleDeploymentModelGroupProtection: def test_generic_allowed_fails_does_not_bypass_single_deployment_protection(self): @@ -1146,6 +1155,7 @@ class TestSingleDeploymentModelGroupProtection: ) assert should_cooldown is True, "explicit per-exception-type policy must still cool down a solo deployment" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestShouldCooldownBasedOnAllowedFailsPolicyFalsyZero: def test_router_level_policy_of_zero_is_not_swallowed_by_allowed_fails(self): @@ -1174,6 +1184,7 @@ class TestShouldCooldownBasedOnAllowedFailsPolicyFalsyZero: ) assert should_cooldown is True, "RateLimitErrorAllowedFails=0 must cool down after the first failure" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestResolveAllowedFailsFromPolicyFallsThrough: def test_none_value_on_first_match_falls_through_to_next_type(self): @@ -1191,6 +1202,7 @@ class TestResolveAllowedFailsFromPolicyFallsThrough: result = _resolve_allowed_fails_from_policy(policy=policy, exception=exc) assert result == 3, "must fall through to BadRequestErrorAllowedFails when the more specific field is unset" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestDeploymentCallbackOnFailureCooldownTimePrecedence: def test_model_info_cooldown_time_used_in_primary_sync_path(self): @@ -1283,11 +1295,7 @@ class TestCallerScopedOAuthAuthFailureCooldown: }, "model_info": { "id": model_id, - **( - {"allowed_fails_policy": allowed_fails_policy} - if allowed_fails_policy is not None - else {} - ), + **({"allowed_fails_policy": allowed_fails_policy} if allowed_fails_policy is not None else {}), }, } ] @@ -1441,6 +1449,191 @@ class TestCallerScopedOAuthAuthFailureCooldown: 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 TestPerUserCopilotAuthFailureCooldown: + @staticmethod + def _router( + model_id: str, + litellm_params: Mapping[str, object], + allowed_fails_policy: Mapping[str, int] | None = None, + ) -> Router: + return _make_router( + model_list=[ + { + "model_name": "copilot", + "litellm_params": { + "model": "github_copilot/gpt-4o", + **litellm_params, + }, + "model_info": { + "id": model_id, + "allowed_fails_policy": dict(allowed_fails_policy) + if allowed_fails_policy is not None + else None, + }, + } + ] + ) + + @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 + async def test_per_user_oauth_caller_auth_failure_does_not_cooldown( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + from litellm.constants import GITHUB_COPILOT_AUTH_TYPE_KEY, GITHUB_COPILOT_PER_USER_AUTH_TYPE + + model_id: Final = "per-user-oauth" + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="copilot-per-user", + credential_values={GITHUB_COPILOT_AUTH_TYPE_KEY: GITHUB_COPILOT_PER_USER_AUTH_TYPE}, + credential_info={"custom_llm_provider": "github_copilot"}, + ) + ], + ) + router: Final = self._router(model_id, {"litellm_credential_name": "copilot-per-user"}) + exception: Final = litellm.CallerCredentialAuthenticationError( + message="reconnect", llm_provider="github_copilot", model="" + ) + result: Final = self._callback(router, model_id, exception) + assert result is False + assert get_cooldown_deployments(router, parent_otel_span=None) == [] + + @pytest.mark.asyncio + async def test_shared_mode_copilot_auth_failure_still_cools_down(self) -> None: + from litellm.llms.github_copilot.authenticator import Authenticator + + model_id: Final = "shared-mode" + with ( + patch.object(Authenticator, "get_api_key", return_value="shared-copilot-token"), + patch.object(Authenticator, "get_api_base", return_value="https://api.githubcopilot.com"), + ): + router: Final = self._router(model_id, {"api_key": "sk-shared"}) + self._callback( + router, + model_id, + litellm.AuthenticationError("upstream rejected the shared token", "github_copilot", "gpt-4o"), + ) + 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 + async def test_fallback_per_user_caller_auth_failure_does_not_cooldown( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + from litellm.constants import GITHUB_COPILOT_AUTH_TYPE_KEY, GITHUB_COPILOT_PER_USER_AUTH_TYPE + + model_id: Final = "fallback-per-user" + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="copilot-per-user", + credential_values={GITHUB_COPILOT_AUTH_TYPE_KEY: GITHUB_COPILOT_PER_USER_AUTH_TYPE}, + credential_info={"custom_llm_provider": "github_copilot"}, + ) + ], + ) + router: Final = self._router(model_id, {"litellm_credential_name": "copilot-per-user"}) + exception: Final = litellm.CallerCredentialAuthenticationError( + message="reconnect", llm_provider="github_copilot", model="" + ) + exception.failed_deployment_id = model_id + + _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exception) + + assert get_cooldown_deployments(router, parent_otel_span=None) == [] + + @pytest.mark.asyncio + async def test_fallback_shared_mode_auth_failure_still_cools_down(self) -> None: + from litellm.llms.github_copilot.authenticator import Authenticator + + model_id: Final = "fallback-shared" + with ( + patch.object(Authenticator, "get_api_key", return_value="shared-copilot-token"), + patch.object(Authenticator, "get_api_base", return_value="https://api.githubcopilot.com"), + ): + router: Final = self._router(model_id, {"api_key": "sk-shared"}) + exception: Final = litellm.AuthenticationError( + "upstream rejected the shared token", "github_copilot", "gpt-4o" + ) + exception.failed_deployment_id = model_id + + _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exception) + + assert get_cooldown_deployments(router, parent_otel_span=None) == [model_id] + + @pytest.mark.asyncio + async def test_fallback_rate_limit_on_per_user_deployment_does_not_cooldown( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + from litellm.constants import GITHUB_COPILOT_AUTH_TYPE_KEY, GITHUB_COPILOT_PER_USER_AUTH_TYPE + + model_id: Final = "fallback-per-user-429" + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="copilot-per-user", + credential_values={GITHUB_COPILOT_AUTH_TYPE_KEY: GITHUB_COPILOT_PER_USER_AUTH_TYPE}, + credential_info={"custom_llm_provider": "github_copilot"}, + ) + ], + ) + router: Final = self._router( + model_id, + {"litellm_credential_name": "copilot-per-user"}, + allowed_fails_policy={"RateLimitErrorAllowedFails": 0}, + ) + exception: Final = litellm.RateLimitError("copilot 429", "github_copilot", "gpt-4o") + exception.failed_deployment_id = model_id + + _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exception) + + assert get_cooldown_deployments(router, parent_otel_span=None) == [] + + @pytest.mark.asyncio + async def test_fallback_rate_limit_on_shared_deployment_still_cools_down(self) -> None: + from litellm.llms.github_copilot.authenticator import Authenticator + + model_id: Final = "fallback-shared-429" + with ( + patch.object(Authenticator, "get_api_key", return_value="shared-copilot-token"), + patch.object(Authenticator, "get_api_base", return_value="https://api.githubcopilot.com"), + ): + router: Final = self._router( + model_id, + {"api_key": "sk-shared"}, + allowed_fails_policy={"RateLimitErrorAllowedFails": 0}, + ) + exception: Final = litellm.RateLimitError("copilot 429", "github_copilot", "gpt-4o") + exception.failed_deployment_id = model_id + + _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exception) + + 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): @@ -1492,6 +1685,7 @@ class TestNewAllowedFailsPolicyFields: assert policy.BadGatewayErrorAllowedFails == 2 assert policy.NotFoundErrorAllowedFails == 1 + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestRouterLevelGetAllowedFailsFromPolicy: """Router.get_allowed_fails_from_policy must handle all AllowedFailsPolicy fields.""" @@ -1527,6 +1721,7 @@ class TestRouterLevelGetAllowedFailsFromPolicy: exc = litellm.RateLimitError("429", "openai", "gpt-4") assert router.get_allowed_fails_from_policy(exc) is None + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") @pytest.mark.asyncio async def test_router_cooldown_event_callback_no_deployment(): @@ -1551,6 +1746,7 @@ async def test_router_cooldown_event_callback_no_deployment(): # Assert that the router's get_deployment method was called mock_router.get_deployment.assert_called_once_with(model_id="test-deployment") + @pytest.fixture def testing_litellm_router(): return Router( @@ -1573,6 +1769,7 @@ def testing_litellm_router(): ] ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_should_run_cooldown_logic(testing_litellm_router): testing_litellm_router.disable_cooldowns = True @@ -1587,6 +1784,7 @@ def test_should_run_cooldown_logic(testing_litellm_router): testing_litellm_router.provider_default_deployment_ids = ["test_deployment"] assert _should_run_cooldown_logic(testing_litellm_router, "test_deployment", 500, Exception("Test")) is False + @pytest.fixture def single_deployment_router(): """A router with one deployment whose model_info.id is the lookup-able @@ -1603,6 +1801,7 @@ def single_deployment_router(): ] ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_should_run_cooldown_logic_generic_bad_request_excluded_by_default( single_deployment_router, @@ -1614,6 +1813,7 @@ def test_should_run_cooldown_logic_generic_bad_request_excluded_by_default( exc = litellm.BadRequestError("bad request", "openai", "gpt-5-mini") assert _should_run_cooldown_logic(single_deployment_router, "dep-1", 400, exc) is False + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_should_run_cooldown_logic_router_level_policy_does_not_override_bad_request_exclusion( single_deployment_router, @@ -1627,6 +1827,7 @@ def test_should_run_cooldown_logic_router_level_policy_does_not_override_bad_req single_deployment_router.allowed_fails_policy = AllowedFailsPolicy(BadRequestErrorAllowedFails=5) assert _should_run_cooldown_logic(single_deployment_router, "dep-1", 400, exc) is False + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_should_run_cooldown_logic_explicit_deployment_level_policy_overrides_content_policy_exclusion( single_deployment_router, @@ -1638,6 +1839,7 @@ def test_should_run_cooldown_logic_explicit_deployment_level_policy_overrides_co deployment_dict["model_info"]["allowed_fails_policy"] = {"ContentPolicyViolationErrorAllowedFails": 0} assert _should_run_cooldown_logic(single_deployment_router, "dep-1", 400, exc) is True + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") class TestHasExplicitAllowedFailsPolicyForException: def test_no_policy_anywhere_returns_false(self, single_deployment_router): @@ -1668,6 +1870,7 @@ class TestHasExplicitAllowedFailsPolicyForException: single_deployment_router.allowed_fails_policy = AllowedFailsPolicy(RateLimitErrorAllowedFails=3) assert _has_explicit_allowed_fails_policy_for_exception(single_deployment_router, None, exc) is False + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_should_cooldown_deployment_rate_limit_error(testing_litellm_router): """ @@ -1677,6 +1880,7 @@ def test_should_cooldown_deployment_rate_limit_error(testing_litellm_router): _exception = litellm.exceptions.RateLimitError("Rate limit", "openai", "gpt-5-mini") assert _should_cooldown_deployment(testing_litellm_router, "test_deployment", 429, _exception) is True + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_should_cooldown_deployment_auth_limit_error(testing_litellm_router): """ @@ -1686,6 +1890,7 @@ def test_should_cooldown_deployment_auth_limit_error(testing_litellm_router): _exception = litellm.exceptions.AuthenticationError("Unauthorized", "openai", "gpt-5-mini") assert _should_cooldown_deployment(testing_litellm_router, "test_deployment", 401, _exception) is True + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") @pytest.mark.parametrize("exception_status", (401, 402)) def test_is_cooldown_required_for_account_errors(testing_litellm_router, exception_status): @@ -1698,6 +1903,7 @@ def test_is_cooldown_required_for_account_errors(testing_litellm_router, excepti is True ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") @pytest.mark.parametrize("allowed_fails", (None, 0)) def test_single_deployment_402_does_not_cooldown( @@ -1726,6 +1932,7 @@ def test_single_deployment_402_does_not_cooldown( is False ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_single_deployment_402_respects_router_allowed_fails_policy() -> None: assert ( @@ -1751,6 +1958,7 @@ def test_single_deployment_402_respects_router_allowed_fails_policy() -> None: is True ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_single_deployment_402_respects_deployment_allowed_fails_policy() -> None: assert ( @@ -1778,6 +1986,7 @@ def test_single_deployment_402_respects_deployment_allowed_fails_policy() -> Non is True ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_multi_deployment_402_cools_down(testing_litellm_router: Router) -> None: assert ( @@ -1794,6 +2003,7 @@ def test_multi_deployment_402_cools_down(testing_litellm_router: Router) -> None is True ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") @pytest.mark.asyncio async def test_should_cooldown_deployment(testing_litellm_router): @@ -1844,6 +2054,7 @@ async def test_should_cooldown_deployment(testing_litellm_router): # expect this to fail since it's now 51% of requests are failing assert _should_cooldown_deployment(testing_litellm_router, deployment_id, 500, Exception("Test")) is True + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") @pytest.mark.asyncio async def test_should_cooldown_deployment_allowed_fails_set_on_router(): @@ -1870,6 +2081,7 @@ async def test_should_cooldown_deployment_allowed_fails_set_on_router(): assert _should_cooldown_deployment(router, "test_deployment", 500, Exception("Test")) is True + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_increment_deployment_successes_for_current_minute_does_not_write_to_redis( testing_litellm_router, @@ -1911,12 +2123,14 @@ def test_increment_deployment_successes_for_current_minute_does_not_write_to_red ) assert testing_litellm_router.cache.in_memory_cache.get_cache("test_deployment:successes") is not None + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_cast_exception_status_to_int(): assert cast_exception_status_to_int(200) == 200 assert cast_exception_status_to_int("404") == 404 assert cast_exception_status_to_int("invalid") == 500 + @pytest.fixture def router(): return Router( @@ -1931,6 +2145,7 @@ def router(): ] ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") @patch("litellm.router_utils.cooldown_handlers.get_deployment_successes_for_current_minute") @patch("litellm.router_utils.cooldown_handlers.get_deployment_failures_for_current_minute") @@ -1950,6 +2165,7 @@ def test_should_cooldown_high_traffic_all_fails(mock_failures, mock_successes, r assert should_cooldown is True, "Should cooldown when all requests fail with sufficient traffic" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") @patch("litellm.router_utils.cooldown_handlers.get_deployment_successes_for_current_minute") @patch("litellm.router_utils.cooldown_handlers.get_deployment_failures_for_current_minute") @@ -1967,6 +2183,7 @@ def test_no_cooldown_low_traffic(mock_failures, mock_successes, router): assert should_cooldown is False, "Should not cooldown when traffic is below threshold" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") @patch("litellm.router_utils.cooldown_handlers.get_deployment_successes_for_current_minute") @patch("litellm.router_utils.cooldown_handlers.get_deployment_failures_for_current_minute") @@ -1986,6 +2203,7 @@ def test_cooldown_rate_limit(mock_failures, mock_successes, router): assert should_cooldown is False, "Should not cooldown on rate limit error for single deployment models" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") @patch("litellm.router_utils.cooldown_handlers.get_deployment_successes_for_current_minute") @patch("litellm.router_utils.cooldown_handlers.get_deployment_failures_for_current_minute") @@ -2003,6 +2221,7 @@ def test_mixed_success_failure(mock_failures, mock_successes, router): assert should_cooldown is False, "Should not cooldown when failure rate is below threshold" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_is_cooldown_required_empty_string_exception_status(testing_litellm_router): """ @@ -2016,6 +2235,7 @@ def test_is_cooldown_required_empty_string_exception_status(testing_litellm_rout assert result is False, "Should not require cooldown when exception_status is empty string" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_should_cooldown_deployment_minimum_request_threshold(testing_litellm_router): """ @@ -2063,6 +2283,7 @@ def test_should_cooldown_deployment_minimum_request_threshold(testing_litellm_ro f"Should cooldown when we have {DEFAULT_FAILURE_THRESHOLD_MINIMUM_REQUESTS} failed requests (100% failure rate)" ) + @pytest.mark.asyncio async def test_dynamic_cooldowns(): """ @@ -2104,6 +2325,7 @@ async def test_dynamic_cooldowns(): assert "cooldown_time" in tmp_mock.call_args[0][0]["litellm_params"] assert tmp_mock.call_args[0][0]["litellm_params"]["cooldown_time"] == 0 + @pytest.mark.asyncio async def test_cooldown_time_zero_uses_zero_not_default(): """ @@ -2148,6 +2370,7 @@ async def test_cooldown_time_zero_uses_zero_not_default(): healthy_deployments, _ = await router._async_get_healthy_deployments(model="gpt-3.5-turbo", parent_otel_span=None) assert len(healthy_deployments) == 1 + def test_should_run_cooldown_logic_early_exit_on_zero_cooldown(): """ Unit test for _should_run_cooldown_logic to verify early exit when time_to_cooldown is 0 @@ -2203,6 +2426,7 @@ def test_should_run_cooldown_logic_early_exit_on_zero_cooldown(): ) assert result is True, "Should run cooldown logic when time_to_cooldown is positive" + @pytest.mark.parametrize("num_deployments", [1, 2]) def test_single_deployment_no_cooldowns(num_deployments: int): """ @@ -2268,9 +2492,7 @@ async def test_single_deployment_no_cooldowns_test_prod(): num_retries=0, ) - with patch.object( - router.cooldown_cache, "add_deployment_to_cooldown", new=MagicMock() - ) as mock_client: + with patch.object(router.cooldown_cache, "add_deployment_to_cooldown", new=MagicMock()) as mock_client: try: await router.acompletion( model="gpt-3.5-turbo", @@ -2371,11 +2593,10 @@ async def test_high_traffic_cooldowns_all_healthy_deployments(): raise e print("model_stats: ", model_stats) - cooldown_list = await async_get_cooldown_deployments( - litellm_router_instance=router, parent_otel_span=None - ) + cooldown_list = await async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None) assert len(cooldown_list) == 0 + @pytest.mark.asyncio() async def test_high_traffic_cooldowns_one_bad_deployment(): """ @@ -2439,7 +2660,6 @@ async def test_high_traffic_cooldowns_one_bad_deployment(): mock_response = "hi" elif bad_deployment_id == model_id: if num_failures / total_requests <= 0.6: - mock_response = "litellm.InternalServerError" elif num_failures / total_requests <= 0.25: @@ -2467,11 +2687,10 @@ async def test_high_traffic_cooldowns_one_bad_deployment(): raise e print("model_stats: ", model_stats) - cooldown_list = await async_get_cooldown_deployments( - litellm_router_instance=router, parent_otel_span=None - ) + cooldown_list = await async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None) assert len(cooldown_list) == 1 + @pytest.mark.asyncio() async def test_high_traffic_cooldowns_one_rate_limited_deployment(): """ @@ -2535,7 +2754,6 @@ async def test_high_traffic_cooldowns_one_rate_limited_deployment(): mock_response = "hi" elif bad_deployment_id == model_id: if num_failures / total_requests <= 0.6: - mock_response = "litellm.RateLimitError" elif num_failures / total_requests <= 0.25: @@ -2566,11 +2784,10 @@ async def test_high_traffic_cooldowns_one_rate_limited_deployment(): raise e print("model_stats: ", model_stats) - cooldown_list = await async_get_cooldown_deployments( - litellm_router_instance=router, parent_otel_span=None - ) + cooldown_list = await async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None) assert len(cooldown_list) == 1 + def test_router_fallbacks_with_cooldowns_and_model_id(): """ Test that after a RateLimitError, the router can still route subsequent @@ -2606,6 +2823,7 @@ def test_router_fallbacks_with_cooldowns_and_model_id(): ) assert response is not None + @pytest.mark.asyncio() async def test_router_fallbacks_with_cooldowns_and_dynamic_credentials(): """ @@ -2644,3 +2862,53 @@ async def test_router_fallbacks_with_cooldowns_and_dynamic_credentials(): await asyncio.sleep(1) cooled_down = await async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None) assert len(cooled_down) == 1 and cooled_down[0] in {"123", "456"} + + +@pytest.mark.parametrize( + "exception,status", + [ + ( + litellm.CallerCredentialAuthenticationError(message="reconnect", llm_provider="github_copilot", model=""), + 401, + ), + (litellm.CallerCredentialRateLimitError(message="slow down", llm_provider="github_copilot", model=""), 429), + ], +) +def test_caller_credential_errors_never_cool_down_the_shared_deployment(single_deployment_router, exception, status): + """A per-user credential failure is scoped to one caller's stored token; cooling down the + shared deployment would punish every other user on the group.""" + assert _should_run_cooldown_logic(single_deployment_router, "dep-1", status, exception) is False + + +def test_per_user_session_upstream_errors_never_cool_down_the_shared_deployment(single_deployment_router): + """A 429/401 from the caller's own Copilot seat is scoped to that user; it must + not cool down the shared deployment for every other caller.""" + from litellm.llms.github_copilot.per_user_auth import GithubCopilotUserSession + + session_kwargs = { + "github_copilot_user_session": GithubCopilotUserSession( + token="copilot-token", api_base="https://api.githubcopilot.com" + ) + } + for exc, status in ( + (litellm.RateLimitError("copilot 429", "github_copilot", "gpt-4o"), 429), + (litellm.AuthenticationError("copilot 401", "github_copilot", "gpt-4o"), 401), + (litellm.InternalServerError("copilot 500", "github_copilot", "gpt-4o"), 500), + ): + assert ( + _should_run_cooldown_logic(single_deployment_router, "dep-1", status, exc, request_kwargs=session_kwargs) + is False + ) + + +def test_shared_mode_upstream_429_still_cools_down_the_deployment(single_deployment_router): + """Regression: without the session marker the same 429 is a deployment-health + signal and cools down exactly as before.""" + exc = litellm.RateLimitError("copilot 429", "github_copilot", "gpt-4o") + assert ( + _should_run_cooldown_logic( + single_deployment_router, "dep-1", 429, exc, request_kwargs={"model": "github_copilot/gpt-4o"} + ) + is True + ) + assert _should_run_cooldown_logic(single_deployment_router, "dep-1", 429, exc) is True diff --git a/tests/unit/test_internal_context.py b/tests/unit/test_internal_context.py index 1f43f9c6312..ccbe1487fc4 100644 --- a/tests/unit/test_internal_context.py +++ b/tests/unit/test_internal_context.py @@ -54,6 +54,7 @@ _IN_MEMORY_ONLY_CALLERS: Final = frozenset( "litellm/llms/bedrock/base_aws_llm.py", "litellm/llms/custom_httpx/http_handler.py", "litellm/llms/gigachat/authenticator.py", + "litellm/llms/github_copilot/per_user_auth.py", "litellm/llms/litellm_proxy/skills/handler.py", "litellm/llms/openai/common_utils.py", "litellm/llms/openai_like/model_info.py", diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index f9bb957e777..16c9b875636 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -87,6 +87,7 @@ CONNECTION_NAMES: Final = ( "azure_scope", "azure_ad_token_provider", "litellm_credential_name", + "github_copilot_auth_type", "configurable_clientside_auth_params", "use_xai_oauth", "fireworks_forward_user_id", @@ -248,6 +249,7 @@ INTERNAL_STATE_NAMES: Final = ( "attempted_targets", "proxy_server_request", "secret_fields", + "github_copilot_user_session", "litellm_trusted_callback_vars", "_litellm_addressed_response_id", "_litellm_strip_stream_usage", diff --git a/tests/unit/types/test_router.py b/tests/unit/types/test_router.py index dce4225f70a..89238e04a5a 100644 --- a/tests/unit/types/test_router.py +++ b/tests/unit/types/test_router.py @@ -19,6 +19,7 @@ from litellm.types.utils import ( CustomPricingLiteLLMParams, MirroredPricingParams, anthropic_wif_litellm_params, + github_copilot_oauth_litellm_params, oauth_token_exchange_litellm_params, openai_wif_litellm_params, server_owned_wif_litellm_params, @@ -242,9 +243,7 @@ def test_model_info_rejects_offset_aware_access_window_times(): with pytest.raises(ValidationError): ModelInfo( id="x", - access_windows=[ - {"start": "22:00+05:00", "end": "06:00", "timezone": "UTC", "team_ids": ["t"]} - ], + access_windows=[{"start": "22:00+05:00", "end": "06:00", "timezone": "UTC", "team_ids": ["t"]}], ) @@ -333,8 +332,12 @@ def test_server_owned_registry_includes_anthropic_openai_and_oauth_token_exchang "token_exchange_profile", "token_exchange_scope", ) + assert github_copilot_oauth_litellm_params == ("github_copilot_auth_type",) assert server_owned_wif_litellm_params == ( - anthropic_wif_litellm_params + openai_wif_litellm_params + oauth_token_exchange_litellm_params + anthropic_wif_litellm_params + + openai_wif_litellm_params + + oauth_token_exchange_litellm_params + + github_copilot_oauth_litellm_params ) assert set(openai_wif_litellm_params) == { "openai_identity_provider_id", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useUserConnections.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useUserConnections.ts new file mode 100644 index 00000000000..2dac978348c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useUserConnections.ts @@ -0,0 +1,15 @@ +import { userConnectionsListCall, UserProviderConnectionsResponse } from "@/components/networking"; +import { useQuery } from "@tanstack/react-query"; +import { createQueryKeys } from "../common/queryKeysFactory"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; + +export const userConnectionsKeys = createQueryKeys("userConnections"); + +export const useUserConnections = () => { + const { accessToken } = useAuthorized(); + return useQuery({ + queryKey: userConnectionsKeys.list({}), + queryFn: async () => await userConnectionsListCall(accessToken!), + enabled: Boolean(accessToken), + }); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.test.tsx index 41f71a3bf12..fb6ed90c002 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.test.tsx @@ -128,13 +128,20 @@ describe("ModelsAndEndpointsPage", () => { expect(teamInfoProps).toHaveBeenLastCalledWith(expect.objectContaining({ is_proxy_admin: false })); }); - it("hides admin-only tabs for a non-admin user", () => { + it("hides admin-only tabs for a non-admin user but keeps LLM Credentials for their connections", () => { mockUseAuthorized.mockReturnValue(NON_ADMIN); renderPage(); - expect(screen.queryByRole("tab", { name: "LLM Credentials" })).not.toBeInTheDocument(); + expect(screen.getByRole("tab", { name: "LLM Credentials" })).toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "Pass-Through Endpoints" })).not.toBeInTheDocument(); expect(screen.queryByRole("tab", { name: "Health Status" })).not.toBeInTheDocument(); }); + it("hides LLM Credentials from an internal viewer", () => { + mockUseAuthorized.mockReturnValue({ ...NON_ADMIN, userRole: "Internal Viewer", isViewOnly: true }); + renderPage(); + expect(screen.queryByRole("tab", { name: "LLM Credentials" })).not.toBeInTheDocument(); + }); + it("keeps the full admin tab order for a real admin", () => { renderPage(); expect(screen.getAllByRole("tab").map((tab) => tab.textContent)).toEqual([ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx index a4e5afe0533..cd2922dff4b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx @@ -114,13 +114,14 @@ export default function ModelsAndEndpointsPage() { // effectiveSessionRole reports proxy_admin_viewer as "Admin", so isAdmin alone would show a // viewer these write-only panels; only the raw-role isViewOnly separates them. Health Status // stays: it is the bucket's one read view, and viewers keep read parity with admins. - ...(isAdmin && !isViewOnly ? (["llm-credentials", "pass-through"] as const) : []), + ...(!isViewOnly && (isAdmin || isInternalUser) ? (["llm-credentials"] as const) : []), + ...(isAdmin && !isViewOnly ? (["pass-through"] as const) : []), ...(isAdmin ? (["health"] as const) : []), ...(isAdmin && !isViewOnly ? (["retry-settings", "model-group-alias", "access-group-budgets", "price-data"] as const) : []), ], - [canCreate, canViewAutoRouters, isAdmin, isViewOnly], + [canCreate, canViewAutoRouters, isAdmin, isInternalUser, isViewOnly], ); const allModelsLabel = isAdmin ? "Deployed Models" : "Your Models"; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.integration.test.tsx new file mode 100644 index 00000000000..929b1435927 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.integration.test.tsx @@ -0,0 +1,91 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen, waitFor, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import LlmCredentialsPanel from "./LlmCredentialsPanel"; + +const networking = vi.hoisted(() => ({ + credentialListCall: vi.fn(), + userConnectionsListCall: vi.fn(), + userConnectionDeleteCall: vi.fn(), +})); + +vi.mock("@/components/networking", async () => ({ + ...(await vi.importActual("@/components/networking")), + ...networking, +})); + +const mockUseAuthorized = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => mockUseAuthorized() })); + +const session = (userRole: string) => ({ accessToken: "at", userRole, userId: "u1", isViewOnly: false }); + +const renderPanel = () => + render( + + + , + ); + +describe("LlmCredentialsPanel", () => { + beforeEach(() => { + networking.credentialListCall.mockReset().mockResolvedValue({ + credentials: [ + { credential_name: "shared-openai", credential_values: {}, credential_info: { custom_llm_provider: "openai" } }, + ], + }); + networking.userConnectionsListCall.mockReset().mockResolvedValue({ + connections: [ + { + credential_name: "copilot-per-user", + provider: "github_copilot", + connected: false, + github_login: null, + connected_at: null, + }, + { + credential_name: "copilot-team", + provider: "github_copilot", + connected: true, + github_login: "octocat", + connected_at: "2026-10-07T00:00:00Z", + }, + ], + }); + networking.userConnectionDeleteCall.mockReset().mockResolvedValue(undefined); + }); + + it("shows an internal user only their connections and never fetches the shared credential list", async () => { + mockUseAuthorized.mockReturnValue(session("Internal User")); + renderPanel(); + + expect(await screen.findByText("copilot-per-user")).toBeInTheDocument(); + expect(screen.getByText("Your connections")).toBeInTheDocument(); + expect(screen.getByText("Connected as @octocat")).toBeInTheDocument(); + expect(screen.queryByText("shared-openai")).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: /Add Credential/ })).not.toBeInTheDocument(); + expect(networking.credentialListCall).not.toHaveBeenCalled(); + }); + + it("shows an admin the credential list and their connections", async () => { + mockUseAuthorized.mockReturnValue(session("Admin")); + renderPanel(); + + expect(await screen.findByText("shared-openai")).toBeInTheDocument(); + expect(await screen.findByText("copilot-per-user")).toBeInTheDocument(); + expect(screen.getByRole("button", { name: /Add Credential/ })).toBeInTheDocument(); + }); + + it("disconnects the selected connection and refreshes the list", async () => { + const user = userEvent.setup(); + mockUseAuthorized.mockReturnValue(session("Internal User")); + renderPanel(); + + const row = await screen.findByRole("row", { name: /copilot-team/ }); + await user.click(within(row).getByRole("button", { name: "Disconnect" })); + + expect(networking.userConnectionDeleteCall).toHaveBeenCalledWith("at", "copilot-team"); + await waitFor(() => expect(networking.userConnectionsListCall).toHaveBeenCalledTimes(2)); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.tsx index c71ed177416..572d7df250f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.tsx @@ -1,7 +1,16 @@ "use client"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import CredentialsPanel from "@/components/model_add/CredentialsPanel"; +import UserConnectionsPanel from "@/components/model_add/UserConnectionsPanel"; +import { all_admin_roles } from "@/utils/roles"; export default function LlmCredentialsPanel() { - return ; + const { userRole } = useAuthorized(); + return ( +
+ {all_admin_roles.includes(userRole ?? "") && } + +
+ ); } diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx index 8d0880360cd..75297a508d2 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx @@ -75,6 +75,16 @@ vi.mock("../networking", async () => { { key: "api_key", label: "Delegated Access Token", field_type: "password" }, ], }, + { + provider: "GITHUB_COPILOT", + provider_display_name: "Github Copilot", + litellm_provider: "github_copilot", + default_model_placeholder: "github_copilot/chat", + credential_fields: [ + { key: "api_base", label: "API Base", field_type: "text" }, + { key: "api_key", label: "API Key", field_type: "password" }, + ], + }, ]), }; }); @@ -110,6 +120,16 @@ vi.mock("@/app/(dashboard)/hooks/providers/useProviderFields", () => ({ { key: "api_key", label: "Delegated Access Token", field_type: "password" }, ], }, + { + provider: "GITHUB_COPILOT", + provider_display_name: "Github Copilot", + litellm_provider: "github_copilot", + default_model_placeholder: "github_copilot/chat", + credential_fields: [ + { key: "api_base", label: "API Base", field_type: "text" }, + { key: "api_key", label: "API Key", field_type: "password" }, + ], + }, ], isLoading: false, error: null, @@ -555,11 +575,11 @@ describe("AddModelForm", () => { }); describe("credential-only provider auth types", () => { - const renderAsAdmin = async () => { + const renderAsAdmin = async (selectedProvider = Providers.MICROSOFT_365_COPILOT) => { 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; + props.selectedProvider = selectedProvider; renderWithProviders(); await screen.findByText("Existing Credentials"); return props; @@ -674,6 +694,94 @@ describe("AddModelForm", () => { api_key: "delegated-token", }); }); + + it("creates a per-user GitHub credential and attaches it without sending its auth type inline", async () => { + const user = userEvent.setup(); + vi.mocked(credentialCreateCall).mockClear(); + vi.mocked(modelCreateCall).mockClear(); + const props = await renderAsAdmin(Providers.GITHUB_COPILOT); + + await chooseSelectOption(user, screen.getByRole("combobox", { name: "Auth Type:" }), "Per-user GitHub OAuth"); + await user.click(screen.getByRole("button", { name: "Create credential" })); + const dialog = await screen.findByRole("dialog"); + const providerSelect = within(dialog).getByPlaceholderText("Select a provider"); + expect(providerSelect).toHaveValue(Providers.GITHUB_COPILOT); + expect(providerSelect).toBeDisabled(); + expect(within(dialog).getByRole("combobox", { name: "Auth Type:" })).toHaveTextContent("Per-user GitHub OAuth"); + fireEvent.change(within(dialog).getByLabelText("Credential Name:"), { + target: { value: "github-per-user" }, + }); + await user.click(within(dialog).getByRole("button", { name: "Add Credential" })); + + await waitFor(() => + expect(credentialCreateCall).toHaveBeenCalledWith("test-access-token", { + credential_name: "github-per-user", + credential_values: { github_copilot_auth_type: "per_user_oauth" }, + credential_info: { custom_llm_provider: Providers.GITHUB_COPILOT }, + }), + ); + expect(props.form.getValues("litellm_credential_name")).toBe("github-per-user"); + + props.handleOk.mockImplementation(async () => { + await handleAddModelSubmit( + { + litellm_credential_name: props.form.getValues("litellm_credential_name"), + custom_llm_provider: Providers.GITHUB_COPILOT, + model_mappings: [ + { + public_name: "copilot-model", + litellm_model: "github_copilot/chat", + }, + ], + }, + "test-access-token", + { resetFields: vi.fn() }, + ); + return true; + }); + await user.click(screen.getByRole("button", { name: "Add Model" })); + + await waitFor(() => expect(modelCreateCall).toHaveBeenCalledOnce()); + const modelParams = vi.mocked(modelCreateCall).mock.calls[0][1].litellm_params; + expect(modelParams).toMatchObject({ litellm_credential_name: "github-per-user" }); + expect(modelParams).not.toHaveProperty("github_copilot_auth_type"); + }); + + it("keeps shared GitHub device login inline when adding a model", async () => { + const user = userEvent.setup(); + vi.mocked(credentialCreateCall).mockClear(); + vi.mocked(modelCreateCall).mockClear(); + const props = await renderAsAdmin(Providers.GITHUB_COPILOT); + + expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent("Shared device login"); + expect(screen.queryByRole("button", { name: "Create credential" })).not.toBeInTheDocument(); + fireEvent.change(screen.getByLabelText("API Key"), { target: { value: "shared-device-token" } }); + props.handleOk.mockImplementation(async () => { + await handleAddModelSubmit( + { + api_key: "shared-device-token", + custom_llm_provider: Providers.GITHUB_COPILOT, + model_mappings: [ + { + public_name: "copilot-model", + litellm_model: "github_copilot/chat", + }, + ], + }, + "test-access-token", + { resetFields: vi.fn() }, + ); + return true; + }); + + await user.click(screen.getByRole("button", { name: "Add Model" })); + + await waitFor(() => expect(modelCreateCall).toHaveBeenCalledOnce()); + expect(credentialCreateCall).not.toHaveBeenCalled(); + expect(vi.mocked(modelCreateCall).mock.calls[0][1].litellm_params).toMatchObject({ + api_key: "shared-device-token", + }); + }); }); describe("cache control bindings reach the parent form store", () => { diff --git a/ui/litellm-dashboard/src/components/add_model/provider_auth_types.test.ts b/ui/litellm-dashboard/src/components/add_model/provider_auth_types.test.ts index 9670c6c6e27..5467ff0085a 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_auth_types.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/provider_auth_types.test.ts @@ -78,4 +78,13 @@ describe("provider auth types", () => { expect(hiddenAuthFieldKeys(authTypes, "first")).toEqual(["second_key", "third_key"]); }); + + it("defaults GitHub Copilot to shared device login and recognises a stored per-user credential", () => { + const authTypes = authTypesFor("GITHUB_COPILOT"); + + expect(inferAuthTypeId(authTypes, { api_key: "" })).toBe("shared_device_login"); + expect(inferAuthTypeId(authTypes, { github_copilot_auth_type: "per_user_oauth" })).toBe("per_user_oauth"); + expect(hiddenAuthFieldKeys(authTypes, "per_user_oauth")).toEqual(["api_base", "api_key"]); + expect(hiddenAuthFieldKeys(authTypes, "shared_device_login")).toEqual(["github_copilot_auth_type"]); + }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/provider_auth_types.ts b/ui/litellm-dashboard/src/components/add_model/provider_auth_types.ts index 37fbcb73e47..6fe54955133 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_auth_types.ts +++ b/ui/litellm-dashboard/src/components/add_model/provider_auth_types.ts @@ -12,7 +12,30 @@ export interface ProviderAuthType { const EMPTY_PROVIDER_AUTH_TYPES: readonly ProviderAuthType[] = []; +export const GITHUB_COPILOT_AUTH_TYPE_KEY = "github_copilot_auth_type"; +export const GITHUB_COPILOT_PER_USER_AUTH_TYPE = "per_user_oauth"; + export const PROVIDER_AUTH_TYPES: Partial> = { + GITHUB_COPILOT: [ + { + id: "shared_device_login", + label: "Shared device login", + description: + "One GitHub device login on the proxy host, stored in its token file, is used for every caller of models on this credential.", + fieldKeys: ["api_base", "api_key"], + requiredFieldKeys: [], + }, + { + id: GITHUB_COPILOT_PER_USER_AUTH_TYPE, + label: "Per-user GitHub OAuth", + credentialOnly: true, + description: + "Each LiteLLM user connects their own GitHub account from LLM Credentials, and their requests use their own GitHub Copilot access. Users who have not connected get a 401.", + fieldKeys: [GITHUB_COPILOT_AUTH_TYPE_KEY], + requiredFieldKeys: [GITHUB_COPILOT_AUTH_TYPE_KEY], + fixedValues: { [GITHUB_COPILOT_AUTH_TYPE_KEY]: GITHUB_COPILOT_PER_USER_AUTH_TYPE }, + }, + ], MICROSOFT_365_COPILOT: [ { id: "oauth_token_exchange", diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.integration.test.tsx new file mode 100644 index 00000000000..1c1a668a982 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.integration.test.tsx @@ -0,0 +1,70 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { describe, expect, it, vi } from "vitest"; +import { useFormContext } from "react-hook-form"; +import { Providers } from "../provider_info_helpers"; +import { GITHUB_COPILOT_AUTH_TYPE_KEY } from "./provider_auth_types"; +import type { MountedFormValues } from "../common_components/MountedFormField"; +import { MountedFormHost } from "../../../tests/mounted-form-host"; +import ProviderSpecificFields from "./provider_specific_fields"; + +vi.mock("../networking", async () => { + const actual = await vi.importActual("../networking"); + return { + ...actual, + getProviderCreateMetadata: vi.fn().mockResolvedValue([ + { + provider: "GITHUB_COPILOT", + provider_display_name: Providers.GITHUB_COPILOT, + litellm_provider: "github_copilot", + default_model_placeholder: "github_copilot/chat", + credential_fields: [ + { key: "api_base", label: "API Base", field_type: "text" }, + { key: "api_key", label: "API Key", field_type: "password" }, + ], + }, + ]), + }; +}); + +const createQueryClient = () => + new QueryClient({ + defaultOptions: { + queries: { + retry: false, + gcTime: 0, + }, + }, + }); + +const GitHubCopilotAuthTypeProbe = () => { + const { watch } = useFormContext(); + return {String(watch(GITHUB_COPILOT_AUTH_TYPE_KEY) ?? "")}; +}; + +describe("ProviderSpecificFields", () => { + it("offers per-user GitHub OAuth as a credential-only model auth type without mounting its fixed value", async () => { + const onCreateCredential = vi.fn(); + const queryClient = createQueryClient(); + render( + + + + + + , + ); + + expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent("Per-user GitHub OAuth"); + expect(screen.getByRole("button", { name: "Create credential" })).toBeInTheDocument(); + expect(screen.getByTestId("github-copilot-auth-type")).toBeEmptyDOMElement(); + await userEvent.click(screen.getByRole("button", { name: "Create credential" })); + expect(onCreateCredential).toHaveBeenCalledWith("per_user_oauth"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.integration.test.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.integration.test.tsx index 4b9f4e61a8f..c28260ca095 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialModal.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.integration.test.tsx @@ -60,6 +60,15 @@ vi.mock("../networking", async () => { { key: "api_key", label: "Delegated Access Token", field_type: "password" }, ], }, + { + provider: "GITHUB_COPILOT", + provider_display_name: Providers.GITHUB_COPILOT, + litellm_provider: "github_copilot", + credential_fields: [ + { key: "api_base", label: "API Base", field_type: "text" }, + { key: "api_key", label: "API Key", field_type: "password" }, + ], + }, ]), }; }); @@ -504,6 +513,63 @@ describe("CredentialModal with Anthropic workload identity federation", () => { }); }); +describe("CredentialModal with GitHub Copilot auth types", () => { + const perUserCopilotCredential: CredentialItem = { + credential_name: "copilot-per-user", + credential_values: { github_copilot_auth_type: "per_user_oauth" }, + credential_info: { custom_llm_provider: "GITHUB_COPILOT" }, + }; + + it("saves a per-user GitHub OAuth credential with only the auth type and no secret", async () => { + const user = userEvent.setup(); + const onSubmit = renderModal({ initialProvider: "GITHUB_COPILOT" }); + + expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent("Shared device login"); + expect(await screen.findByLabelText("API Key")).toBeInTheDocument(); + await chooseOption(user, /^Auth Type:/, "Per-user GitHub OAuth"); + expect(screen.queryByLabelText("API Key")).not.toBeInTheDocument(); + expect(screen.queryByLabelText("API Base")).not.toBeInTheDocument(); + fill("Credential Name:", "copilot-per-user"); + await user.click(screen.getByRole("button", { name: "Add Credential" })); + + expect(onSubmit).toHaveBeenCalledWith( + { + credential_name: "copilot-per-user", + custom_llm_provider: "GITHUB_COPILOT", + github_copilot_auth_type: "per_user_oauth", + }, + [], + ); + }); + + it("keeps the shared device login credential free of the per-user auth type", async () => { + const user = userEvent.setup(); + const onSubmit = renderModal({ initialProvider: "GITHUB_COPILOT" }); + + await screen.findByLabelText("API Key"); + fill("Credential Name:", "copilot-shared"); + await user.click(screen.getByRole("button", { name: "Add Credential" })); + + expect(onSubmit).toHaveBeenCalledWith( + { credential_name: "copilot-shared", custom_llm_provider: "GITHUB_COPILOT" }, + [], + ); + }); + + it("opens a stored per-user credential on its auth type and deletes it when switched to shared", async () => { + const user = userEvent.setup(); + const onSubmit = renderModal({ mode: "edit", existingCredential: perUserCopilotCredential }); + + expect(await screen.findByRole("combobox", { name: "Auth Type:" })).toHaveTextContent("Per-user GitHub OAuth"); + await chooseOption(user, /^Auth Type:/, "Shared device login"); + await user.click(screen.getByRole("button", { name: "Update Credential" })); + + const [values, valuesToDelete] = onSubmit.mock.calls[0]; + expect(values).toEqual({ credential_name: "copilot-per-user", custom_llm_provider: "GITHUB_COPILOT" }); + expect(valuesToDelete).toEqual(["github_copilot_auth_type"]); + }); +}); + describe("CredentialModal public JWKS for a LiteLLM-signed Anthropic credential", () => { const signedCredential: CredentialItem = { credential_name: "anthropic-signed", diff --git a/ui/litellm-dashboard/src/components/model_add/GithubCopilotConnectModal.integration.test.tsx b/ui/litellm-dashboard/src/components/model_add/GithubCopilotConnectModal.integration.test.tsx new file mode 100644 index 00000000000..963ebab81d3 --- /dev/null +++ b/ui/litellm-dashboard/src/components/model_add/GithubCopilotConnectModal.integration.test.tsx @@ -0,0 +1,126 @@ +import { act, fireEvent, render, screen } from "@testing-library/react"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; + +import GithubCopilotConnectModal from "./GithubCopilotConnectModal"; + +const networking = vi.hoisted(() => ({ + userConnectionStartCall: vi.fn(), + userConnectionPollCall: vi.fn(), +})); + +vi.mock("@/components/networking", () => networking); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => ({ accessToken: "session-token" }) })); + +const START_RESPONSE = { + user_code: "ABCD-1234", + verification_uri: "https://github.com/login/device", + expires_in: 900, + interval: 5, + flow_handle: "fh-1", +}; + +const renderModal = () => { + const props = { onClose: vi.fn(), onConnected: vi.fn(), onDisconnect: vi.fn() }; + render(); + return props; +}; + +const advance = async (seconds: number) => { + await act(async () => { + await vi.advanceTimersByTimeAsync(seconds * 1000); + }); +}; + +describe("GithubCopilotConnectModal", () => { + beforeEach(() => { + vi.useFakeTimers({ shouldAdvanceTime: false }); + networking.userConnectionStartCall.mockReset().mockResolvedValue(START_RESPONSE); + networking.userConnectionPollCall.mockReset(); + }); + + afterEach(() => { + vi.useRealTimers(); + }); + + it("shows the device code and link, polls at the given interval, honours slow_down and ends connected", async () => { + networking.userConnectionPollCall + .mockResolvedValueOnce({ status: "pending" }) + .mockResolvedValueOnce({ status: "slow_down", interval: 10 }) + .mockResolvedValueOnce({ status: "connected", github_login: "octocat" }); + const props = renderModal(); + await advance(0); + + expect(networking.userConnectionStartCall).toHaveBeenCalledWith("session-token", "copilot-per-user"); + expect(screen.getByTestId("github-device-code")).toHaveTextContent("ABCD-1234"); + expect(screen.getByRole("link", { name: "https://github.com/login/device" })).toHaveAttribute( + "href", + "https://github.com/login/device", + ); + + await advance(4.9); + expect(networking.userConnectionPollCall).not.toHaveBeenCalled(); + await advance(0.1); + expect(networking.userConnectionPollCall).toHaveBeenCalledTimes(1); + expect(networking.userConnectionPollCall).toHaveBeenCalledWith("session-token", "copilot-per-user", "fh-1"); + await advance(5); + expect(networking.userConnectionPollCall).toHaveBeenCalledTimes(2); + + await advance(9.9); + expect(networking.userConnectionPollCall).toHaveBeenCalledTimes(2); + await advance(0.1); + expect(networking.userConnectionPollCall).toHaveBeenCalledTimes(3); + + expect(screen.getByText("Connected as @octocat")).toBeInTheDocument(); + expect(props.onConnected).toHaveBeenCalledTimes(1); + await advance(30); + expect(networking.userConnectionPollCall).toHaveBeenCalledTimes(3); + + fireEvent.click(screen.getByRole("button", { name: "Disconnect" })); + expect(props.onDisconnect).toHaveBeenCalledTimes(1); + }); + + it("stops polling and offers a restart when the device code expires", async () => { + networking.userConnectionStartCall.mockResolvedValue({ ...START_RESPONSE, expires_in: 12 }); + networking.userConnectionPollCall.mockResolvedValue({ status: "pending" }); + renderModal(); + await advance(0); + await advance(5); + await advance(5); + await advance(5); + + expect(networking.userConnectionPollCall).toHaveBeenCalledTimes(2); + expect(screen.getByRole("alert")).toHaveTextContent("The device code expired"); + await advance(30); + expect(networking.userConnectionPollCall).toHaveBeenCalledTimes(2); + + fireEvent.click(screen.getByRole("button", { name: "Start again" })); + await advance(0); + expect(networking.userConnectionStartCall).toHaveBeenCalledTimes(2); + expect(screen.getByTestId("github-device-code")).toBeInTheDocument(); + }); + + it.each([ + ["denied", "denied on GitHub"], + ["no_copilot_seat", "does not have GitHub Copilot access"], + ["expired", "The device code expired"], + ])("shows the %s error and stops polling", async (status, message) => { + networking.userConnectionPollCall.mockResolvedValue({ status }); + const props = renderModal(); + await advance(0); + await advance(5); + + expect(screen.getByRole("alert")).toHaveTextContent(message); + expect(props.onConnected).not.toHaveBeenCalled(); + await advance(30); + expect(networking.userConnectionPollCall).toHaveBeenCalledTimes(1); + }); + + it("shows the proxy error when the start request fails", async () => { + networking.userConnectionStartCall.mockRejectedValue(new Error("Credential not found")); + renderModal(); + await advance(0); + + expect(screen.getByRole("alert")).toHaveTextContent("Credential not found"); + expect(networking.userConnectionPollCall).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/model_add/GithubCopilotConnectModal.tsx b/ui/litellm-dashboard/src/components/model_add/GithubCopilotConnectModal.tsx new file mode 100644 index 00000000000..8c76b47b6db --- /dev/null +++ b/ui/litellm-dashboard/src/components/model_add/GithubCopilotConnectModal.tsx @@ -0,0 +1,136 @@ +"use client"; + +import { useEffect, useReducer } from "react"; + +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { userConnectionPollCall, userConnectionStartCall } from "@/components/networking"; +import CopyButton from "@/components/shared/CopyButton"; +import { Button } from "@/components/ui/button"; +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, +} from "@/components/ui/dialog"; +import { extractProxyErrorMessage } from "@/lib/http/client"; + +import { connectFlowReducer, failureMessage, initialConnectFlowState } from "./github_copilot_connect_flow"; + +interface GithubCopilotConnectModalProps { + readonly credentialName: string; + readonly onClose: () => void; + readonly onConnected: () => void; + readonly onDisconnect: () => void; +} + +export default function GithubCopilotConnectModal({ + credentialName, + onClose, + onConnected, + onDisconnect, +}: GithubCopilotConnectModalProps) { + const { accessToken } = useAuthorized(); + const [state, dispatch] = useReducer(connectFlowReducer, initialConnectFlowState); + + useEffect(() => { + if (state.kind !== "starting" || !accessToken) { + return; + } + let cancelled = false; + userConnectionStartCall(accessToken, credentialName) + .then((response) => !cancelled && dispatch({ type: "started", response, now: Date.now() })) + .catch((error: unknown) => !cancelled && dispatch({ type: "errored", message: extractProxyErrorMessage(error) })); + return () => { + cancelled = true; + }; + }, [state, accessToken, credentialName]); + + useEffect(() => { + if (state.kind !== "awaiting" || !accessToken) { + return; + } + let cancelled = false; + const timer = setTimeout(() => { + if (Date.now() >= state.expiresAt) { + dispatch({ type: "timed_out" }); + return; + } + userConnectionPollCall(accessToken, credentialName, state.flowHandle) + .then((response) => !cancelled && dispatch({ type: "polled", response })) + .catch( + (error: unknown) => !cancelled && dispatch({ type: "errored", message: extractProxyErrorMessage(error) }), + ); + }, state.intervalSeconds * 1000); + return () => { + cancelled = true; + clearTimeout(timer); + }; + }, [state, accessToken, credentialName]); + + useEffect(() => { + if (state.kind === "connected") { + onConnected(); + } + }, [state.kind, onConnected]); + + return ( + !open && onClose()}> + + + Connect GitHub Copilot + Credential: {credentialName} + + + {state.kind === "starting" &&

Requesting a device code from GitHub...

} + + {state.kind === "awaiting" && ( +
+

+ Open{" "} + + {state.verificationUri} + {" "} + and enter this code: +

+
+ + {state.userCode} + + +
+

Waiting for you to approve the request on GitHub...

+
+ )} + + {state.kind === "connected" &&

Connected as @{state.githubLogin}

} + + {state.kind === "failed" && ( +

+ {failureMessage(state)} +

+ )} + + + {state.kind === "failed" && ( + + )} + {state.kind === "connected" && ( + + )} + + +
+
+ ); +} diff --git a/ui/litellm-dashboard/src/components/model_add/UserConnectionsPanel.integration.test.tsx b/ui/litellm-dashboard/src/components/model_add/UserConnectionsPanel.integration.test.tsx new file mode 100644 index 00000000000..7f9dd4c519f --- /dev/null +++ b/ui/litellm-dashboard/src/components/model_add/UserConnectionsPanel.integration.test.tsx @@ -0,0 +1,50 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { describe, expect, it, vi } from "vitest"; + +import { userConnectionsListCall } from "@/components/networking"; + +import UserConnectionsPanel from "./UserConnectionsPanel"; + +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => ({ accessToken: "token", userId: "u", userRole: "internal_user" }), +})); + +vi.mock("@/components/networking", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + userConnectionsListCall: vi.fn(), + userConnectionDeleteCall: vi.fn(), + }; +}); + +function renderPanel() { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return render( + + + , + ); +} + +describe("UserConnectionsPanel", () => { + it("shows an error state with retry instead of an empty list when the query fails", async () => { + vi.mocked(userConnectionsListCall).mockRejectedValue(new Error("boom")); + renderPanel(); + expect(await screen.findByText("Could not load your connections.")).toBeInTheDocument(); + expect(screen.queryByText("No credentials need a personal connection.")).not.toBeInTheDocument(); + }); + + it("refetches when retry is clicked", async () => { + const user = userEvent.setup(); + vi.mocked(userConnectionsListCall) + .mockRejectedValueOnce(new Error("boom")) + .mockResolvedValueOnce({ connections: [] }); + renderPanel(); + await user.click(await screen.findByRole("button", { name: "Retry" })); + expect(await screen.findByText("No credentials need a personal connection.")).toBeInTheDocument(); + expect(vi.mocked(userConnectionsListCall).mock.calls.length).toBeGreaterThanOrEqual(2); + }); +}); diff --git a/ui/litellm-dashboard/src/components/model_add/UserConnectionsPanel.tsx b/ui/litellm-dashboard/src/components/model_add/UserConnectionsPanel.tsx new file mode 100644 index 00000000000..bdf5c3d523d --- /dev/null +++ b/ui/litellm-dashboard/src/components/model_add/UserConnectionsPanel.tsx @@ -0,0 +1,130 @@ +"use client"; + +import { useQueryClient } from "@tanstack/react-query"; +import { useCallback, useState } from "react"; + +import { userConnectionsKeys, useUserConnections } from "@/app/(dashboard)/hooks/credentials/useUserConnections"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { userConnectionDeleteCall, UserProviderConnection } from "@/components/networking"; +import { Providers } from "@/components/provider_info_helpers"; +import { Badge } from "@/components/ui/badge"; +import { Button } from "@/components/ui/button"; +import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; +import { toast } from "@/lib/toast"; + +import GithubCopilotConnectModal from "./GithubCopilotConnectModal"; + +const providerLabel = (provider: string): string => + Providers[provider.toUpperCase() as keyof typeof Providers] ?? + Object.entries(Providers).find(([key]) => key.toLowerCase() === provider.toLowerCase())?.[1] ?? + provider; + +export default function UserConnectionsPanel() { + const { accessToken } = useAuthorized(); + const queryClient = useQueryClient(); + const { data, isLoading, isError, refetch } = useUserConnections(); + const [connectingCredential, setConnectingCredential] = useState(null); + const connections: readonly UserProviderConnection[] = data?.connections ?? []; + const showEmpty = !isLoading && !isError && connections.length === 0; + + const refreshConnections = useCallback( + () => queryClient.invalidateQueries({ queryKey: userConnectionsKeys.all }), + [queryClient], + ); + + const disconnect = async (credentialName: string) => { + if (!accessToken) { + return; + } + try { + await userConnectionDeleteCall(accessToken, credentialName); + toast.success("GitHub Copilot disconnected"); + setConnectingCredential(null); + await refreshConnections(); + } catch (error) { + toast.error("Failed to disconnect GitHub Copilot"); + } + }; + + return ( +
+
+

Your connections

+

+ Credentials that use your own account. Connect once, and your requests to models on that credential use your + own access. +

+
+ + + + Credential + Provider + Status + Actions + + + + {isLoading && ( + + + Loading connections... + + + )} + {isError && ( + + +
+ Could not load your connections. + +
+
+
+ )} + {showEmpty && ( + + + No credentials need a personal connection. + + + )} + {connections.map((connection) => ( + + {connection.credential_name} + {providerLabel(connection.provider)} + + {connection.connected ? ( + Connected as @{connection.github_login} + ) : ( + Not connected + )} + + + {connection.connected ? ( + + ) : ( + + )} + + + ))} +
+
+ {connectingCredential !== null && ( + setConnectingCredential(null)} + onConnected={refreshConnections} + onDisconnect={() => disconnect(connectingCredential)} + /> + )} +
+ ); +} diff --git a/ui/litellm-dashboard/src/components/model_add/github_copilot_connect_flow.test.ts b/ui/litellm-dashboard/src/components/model_add/github_copilot_connect_flow.test.ts new file mode 100644 index 00000000000..72c2e783adf --- /dev/null +++ b/ui/litellm-dashboard/src/components/model_add/github_copilot_connect_flow.test.ts @@ -0,0 +1,86 @@ +import { describe, expect, it } from "vitest"; + +import { connectFlowReducer, type ConnectFlowState, initialConnectFlowState } from "./github_copilot_connect_flow"; + +const started = (verificationUri = "https://github.com/login/device"): ConnectFlowState => + connectFlowReducer(initialConnectFlowState, { + type: "started", + response: { + user_code: "ABCD-1234", + verification_uri: verificationUri, + expires_in: 900, + interval: 5, + flow_handle: "fh-1", + }, + now: 1_000, + }); + +describe("connectFlowReducer", () => { + it("moves from starting to awaiting with the code, expiry and interval from the start response", () => { + const expected: ConnectFlowState = { + kind: "awaiting", + attempt: 0, + pollCount: 0, + userCode: "ABCD-1234", + verificationUri: "https://github.com/login/device", + expiresAt: 901_000, + intervalSeconds: 5, + flowHandle: "fh-1", + }; + expect(started()).toEqual(expected); + }); + + it("never links anywhere but github.com", () => { + const state = started("https://evil.example/login/device"); + expect(state.kind === "awaiting" && state.verificationUri).toBe("https://github.com/login/device"); + }); + + it("keeps waiting on pending and slows down by the server interval or by five seconds", () => { + const pending = connectFlowReducer(started(), { type: "polled", response: { status: "pending" } }); + expect(pending).toMatchObject({ kind: "awaiting", pollCount: 1, intervalSeconds: 5 }); + + const slowed = connectFlowReducer(pending, { type: "polled", response: { status: "slow_down", interval: 12 } }); + expect(slowed).toMatchObject({ kind: "awaiting", pollCount: 2, intervalSeconds: 12 }); + + const slowedAgain = connectFlowReducer(slowed, { type: "polled", response: { status: "slow_down" } }); + expect(slowedAgain).toMatchObject({ kind: "awaiting", intervalSeconds: 17 }); + }); + + it.each(["expired", "denied", "no_copilot_seat"] as const)("fails with %s", (status) => { + expect(connectFlowReducer(started(), { type: "polled", response: { status } })).toEqual({ + kind: "failed", + attempt: 0, + reason: status, + }); + }); + + it("connects with the GitHub login", () => { + expect( + connectFlowReducer(started(), { type: "polled", response: { status: "connected", github_login: "octocat" } }), + ).toEqual({ kind: "connected", attempt: 0, githubLogin: "octocat" }); + }); + + it("expires locally once the device code lifetime has passed", () => { + expect(connectFlowReducer(started(), { type: "timed_out" })).toEqual({ + kind: "failed", + attempt: 0, + reason: "expired", + }); + }); + + it("restarts with a new attempt after a failure", () => { + const failed = connectFlowReducer(started(), { type: "errored", message: "boom" }); + const expected: ConnectFlowState = { kind: "failed", attempt: 0, reason: "error", message: "boom" }; + expect(failed).toEqual(expected); + expect(connectFlowReducer(failed, { type: "restarted" })).toEqual({ kind: "starting", attempt: 1 }); + }); + + it("ignores a late poll result after the flow left the awaiting state", () => { + const connected = connectFlowReducer(started(), { + type: "polled", + response: { status: "connected", github_login: "octocat" }, + }); + expect(connectFlowReducer(connected, { type: "polled", response: { status: "expired" } })).toBe(connected); + expect(connectFlowReducer(connected, { type: "errored", message: "late" })).toBe(connected); + }); +}); diff --git a/ui/litellm-dashboard/src/components/model_add/github_copilot_connect_flow.ts b/ui/litellm-dashboard/src/components/model_add/github_copilot_connect_flow.ts new file mode 100644 index 00000000000..0e38fc91bb1 --- /dev/null +++ b/ui/litellm-dashboard/src/components/model_add/github_copilot_connect_flow.ts @@ -0,0 +1,101 @@ +import type { UserConnectionPollResponse, UserConnectionStartResponse } from "@/components/networking"; + +export const GITHUB_DEVICE_LOGIN_URL = "https://github.com/login/device"; +const SLOW_DOWN_INCREMENT_SECONDS = 5; + +export type ConnectFailureReason = "expired" | "denied" | "no_copilot_seat" | "error"; + +export type ConnectFlowState = + | { readonly kind: "starting"; readonly attempt: number } + | { + readonly kind: "awaiting"; + readonly attempt: number; + readonly pollCount: number; + readonly userCode: string; + readonly verificationUri: string; + readonly expiresAt: number; + readonly intervalSeconds: number; + readonly flowHandle: string; + } + | { readonly kind: "connected"; readonly attempt: number; readonly githubLogin: string } + | { + readonly kind: "failed"; + readonly attempt: number; + readonly reason: ConnectFailureReason; + readonly message?: string; + }; + +export type ConnectFlowEvent = + | { readonly type: "started"; readonly response: UserConnectionStartResponse; readonly now: number } + | { readonly type: "polled"; readonly response: UserConnectionPollResponse } + | { readonly type: "timed_out" } + | { readonly type: "errored"; readonly message: string } + | { readonly type: "restarted" }; + +export const initialConnectFlowState: ConnectFlowState = { kind: "starting", attempt: 0 }; + +const safeVerificationUri = (uri: string): string => + uri.startsWith("https://github.com/") ? uri : GITHUB_DEVICE_LOGIN_URL; + +const afterPoll = ( + state: Extract, + response: UserConnectionPollResponse, +): ConnectFlowState => { + switch (response.status) { + case "pending": + return { ...state, pollCount: state.pollCount + 1 }; + case "slow_down": + return { + ...state, + pollCount: state.pollCount + 1, + intervalSeconds: response.interval ?? state.intervalSeconds + SLOW_DOWN_INCREMENT_SECONDS, + }; + case "connected": + return { kind: "connected", attempt: state.attempt, githubLogin: response.github_login ?? "" }; + case "expired": + case "denied": + case "no_copilot_seat": + return { kind: "failed", attempt: state.attempt, reason: response.status }; + } +}; + +export const connectFlowReducer = (state: ConnectFlowState, event: ConnectFlowEvent): ConnectFlowState => { + switch (event.type) { + case "restarted": + return { kind: "starting", attempt: state.attempt + 1 }; + case "errored": + return state.kind === "connected" + ? state + : { kind: "failed", attempt: state.attempt, reason: "error", message: event.message }; + case "started": + return state.kind !== "starting" + ? state + : { + kind: "awaiting", + attempt: state.attempt, + pollCount: 0, + userCode: event.response.user_code, + verificationUri: safeVerificationUri(event.response.verification_uri), + expiresAt: event.now + event.response.expires_in * 1000, + intervalSeconds: event.response.interval, + flowHandle: event.response.flow_handle, + }; + case "polled": + return state.kind === "awaiting" ? afterPoll(state, event.response) : state; + case "timed_out": + return state.kind === "awaiting" ? { kind: "failed", attempt: state.attempt, reason: "expired" } : state; + } +}; + +export const failureMessage = (state: Extract): string => { + switch (state.reason) { + case "expired": + return "The device code expired before it was approved. Start again to get a new code."; + case "denied": + return "The request was denied on GitHub. Start again if this was a mistake."; + case "no_copilot_seat": + return "This GitHub account does not have GitHub Copilot access, so nothing was saved."; + case "error": + return state.message ?? "Something went wrong while connecting to GitHub."; + } +}; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 9603f85a3a3..21434b66e53 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -2599,6 +2599,59 @@ export const credentialDeleteCall = async (accessToken: string, credentialName: } }; +export interface UserProviderConnection { + credential_name: string; + provider: string; + connected: boolean; + github_login: string | null; + connected_at: string | null; +} + +export interface UserProviderConnectionsResponse { + connections: UserProviderConnection[]; +} + +export interface UserConnectionStartResponse { + user_code: string; + verification_uri: string; + expires_in: number; + interval: number; + flow_handle: string; +} + +export type UserConnectionPollStatus = "pending" | "slow_down" | "expired" | "denied" | "no_copilot_seat" | "connected"; + +export interface UserConnectionPollResponse { + status: UserConnectionPollStatus; + interval?: number | null; + github_login?: string | null; +} + +const userConnectionPath = (credentialName: string): string => + `/credentials/${encodeURIComponent(credentialName)}/user_connection`; + +export const userConnectionsListCall = (accessToken: string): Promise => + apiClient.get("/credentials/user_connections", { accessToken }); + +export const userConnectionStartCall = ( + accessToken: string, + credentialName: string, +): Promise => + apiClient.post(`${userConnectionPath(credentialName)}/start`, { accessToken }); + +export const userConnectionPollCall = ( + accessToken: string, + credentialName: string, + flowHandle: string, +): Promise => + apiClient.post(`${userConnectionPath(credentialName)}/poll`, { + accessToken, + body: { flow_handle: flowHandle }, + }); + +export const userConnectionDeleteCall = (accessToken: string, credentialName: string): Promise => + apiClient.delete(userConnectionPath(credentialName), { accessToken }); + export const credentialUpdateCall = async ( accessToken: string, credentialName: string, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 4c072cfb7ee..ba38e94ff0b 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -3786,6 +3786,26 @@ export interface paths { patch?: never; trace?: never; }; + "/credentials/user_connections": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * List User Connections + * @description List the calling user's per-user provider connections. + */ + get: operations["list_user_connections_credentials_user_connections_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/credentials/{credential_name}": { parameters: { query?: never; @@ -3832,6 +3852,66 @@ export interface paths { patch?: never; trace?: never; }; + "/credentials/{credential_name}/user_connection": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + post?: never; + /** + * Delete User Connection + * @description Disconnect the calling user's stored GitHub token for a per-user credential. Idempotent. + */ + delete: operations["delete_user_connection_credentials__credential_name__user_connection_delete"]; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/credentials/{credential_name}/user_connection/poll": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** + * Poll User Connection + * @description Poll the device flow once and persist the connection on completion. + */ + post: operations["poll_user_connection_credentials__credential_name__user_connection_poll_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/credentials/{credential_name}/user_connection/start": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** + * Start User Connection + * @description Begin a GitHub device flow for the calling user's connection to a per-user credential. + */ + post: operations["start_user_connection_credentials__credential_name__user_connection_start_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/cursor/chat/completions": { parameters: { query?: never; @@ -50380,6 +50460,41 @@ export interface components { */ severity: "info" | "warning" | "error"; }; + /** UserConnectionDeleteResponse */ + UserConnectionDeleteResponse: { + /** Status */ + status: string; + }; + /** UserConnectionPollRequest */ + UserConnectionPollRequest: { + /** Flow Handle */ + flow_handle: string; + }; + /** UserConnectionPollResponse */ + UserConnectionPollResponse: { + /** Github Login */ + github_login?: string | null; + /** Interval */ + interval?: number | null; + /** + * Status + * @enum {string} + */ + status: "pending" | "slow_down" | "expired" | "denied" | "connected" | "no_copilot_seat"; + }; + /** UserConnectionStartResponse */ + UserConnectionStartResponse: { + /** Expires In */ + expires_in: number; + /** Flow Handle */ + flow_handle: string; + /** Interval */ + interval: number; + /** User Code */ + user_code: string; + /** Verification Uri */ + verification_uri: string; + }; /** * UserCreateResult * @description Outcome for one row of `POST /management/v1/users/bulk`. `teams` lists the teams the user was actually @@ -50523,6 +50638,24 @@ export interface components { /** Users */ users: components["schemas"]["LiteLLM_UserTableWithKeyCount"][]; }; + /** UserProviderConnection */ + UserProviderConnection: { + /** Connected */ + connected: boolean; + /** Connected At */ + connected_at?: string | null; + /** Credential Name */ + credential_name: string; + /** Github Login */ + github_login?: string | null; + /** Provider */ + provider: string; + }; + /** UserProviderConnectionsResponse */ + UserProviderConnectionsResponse: { + /** Connections */ + connections: components["schemas"]["UserProviderConnection"][]; + }; /** * UserUpdateResult * @description Result of a single user update operation @@ -57512,6 +57645,26 @@ export interface operations { }; }; }; + list_user_connections_credentials_user_connections_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["UserProviderConnectionsResponse"]; + }; + }; + }; + }; delete_credential_credentials__credential_name__delete: { parameters: { query?: never; @@ -57612,6 +57765,106 @@ export interface operations { }; }; }; + delete_user_connection_credentials__credential_name__user_connection_delete: { + parameters: { + query?: never; + header?: never; + path: { + /** @description The credential name, percent-decoded */ + credential_name: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["UserConnectionDeleteResponse"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + poll_user_connection_credentials__credential_name__user_connection_poll_post: { + parameters: { + query?: never; + header?: never; + path: { + /** @description The credential name, percent-decoded */ + credential_name: string; + }; + cookie?: never; + }; + requestBody: { + content: { + "application/json": components["schemas"]["UserConnectionPollRequest"]; + }; + }; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["UserConnectionPollResponse"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + start_user_connection_credentials__credential_name__user_connection_start_post: { + parameters: { + query?: never; + header?: never; + path: { + /** @description The credential name, percent-decoded */ + credential_name: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["UserConnectionStartResponse"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; cursor_chat_completions_cursor_chat_completions_post: { parameters: { query?: never;