feat(github_copilot): per-user GitHub OAuth connections for Copilot credentials (#45241)

* feat(github_copilot): per-user GitHub OAuth connections for Copilot credentials

Related to LIT-9306

* fix(ui): make per-user GitHub Copilot OAuth credential-only in Add Model

* fix(proxy): render LiteLLM_UserProviderCredentials prisma spans

* style(ui): prettier-format add_model tests

* fix(lint): annotate connection endpoints and narrow per-user credential excepts

* fix(lint): satisfy type-discipline gate for per-user copilot code

* fix(types): clear basedpyright and strict-lint gate regressions

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(types): keep runtime guards on untyped responses input

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* style: ruff format per-user Copilot files

* fix(lint): keep type-discipline suppressions on flagged lines

* fix(github_copilot): resolve team and deployment-id per-user routes, share token cache invalidation

* test(github_copilot): scope response-cache mutation via monkeypatch, assert route allowlist

* fix(github_copilot): fall back to the database on Redis errors and skip worker-local token caching

* fix(lint): suppress mutable return annotation after base merge

* style: shorten rebind-ok reason so the dispatch line stays formatted

* fix(github_copilot): harden per-user token cache, overrides, fallbacks, headers and device flow

* fix(lint): suppress MutableMapping param on session auth pinning

* test(github_copilot): narrow pytest.raises to HTTPException for service-key override

* fix(github_copilot): fail closed on revoke errors, strip caller auth headers, widen fallback discovery

* fix(github_copilot): verify revoke tombstone, bound fallback discovery, cover aliased fallbacks

* fix(github_copilot): make fallback discovery a router superset and guard credential override at deployment

* style(router): drop explanatory comment on per-user credential guard

* fix(github_copilot): gate fallback discovery on per-user credentials and bound aggregate targets

* fix(github_copilot): count every fallback mapping toward the discovery limit

* fix(github_copilot): suppress banned typing cast imports with per-line reasons

* test(github_copilot): cover per-user 401 cooldown skip through the router failure callback

* test(github_copilot): cover per-user 401 cooldown skip through router and fallback paths

* fix(github_copilot): cover auto-router discovery, per-user fallback cooldown and connect cache race

* fix(github_copilot): cover every strategy router kind and verify the connect cache write

* fix merge conflict marker remnant in credential endpoint tests

* fix(github_copilot): do not purge user connections on label-only credential updates

* fix(github_copilot): read user provider credentials from the writer engine

* test(github_copilot): move per-user auth type form test to the integration tier

* test(github_copilot): annotate new credential regression test locals as Final

* fix(github_copilot): type per-user OAuth JSON and DB boundaries

* fix(github_copilot): validate per-user OAuth boundaries and restore truncated comments

* style(github_copilot): complete cut-off suppression reasons on new casts

* refactor(github_copilot): annotate new boundary-typing locals as Final

---------

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

View file

@ -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");

View file

@ -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 {

View file

@ -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,

View file

@ -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)

View file

@ -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"

View file

@ -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__(

View file

@ -293,6 +293,7 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset(
"LiteLLM_BackgroundInteractionSettlement",
"LiteLLM_ErrorLogs",
"LiteLLM_UserNotifications",
"LiteLLM_UserProviderCredentials",
"LiteLLM_TeamMembership",
"LiteLLM_OrganizationMembership",
"LiteLLM_InvitationLink",

View file

@ -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)

View file

@ -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,
}

View file

@ -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,

View file

@ -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,

View file

@ -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:
"""

View file

@ -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":

View file

@ -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}"

View file

@ -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

View file

@ -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"

View file

@ -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)

View file

@ -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

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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:
"""

View file

@ -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,

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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 {

View file

@ -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)

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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]]

View file

@ -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))

View file

@ -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)

View file

@ -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 {

View file

@ -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

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -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

File diff suppressed because it is too large Load diff

View file

@ -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 = {}

View file

@ -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

View file

@ -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,

View file

@ -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()

View file

@ -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

File diff suppressed because it is too large Load diff

View file

@ -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

View file

@ -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",

View file

@ -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",

View file

@ -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",

View file

@ -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<UserProviderConnectionsResponse>({
queryKey: userConnectionsKeys.list({}),
queryFn: async () => await userConnectionsListCall(accessToken!),
enabled: Boolean(accessToken),
});
};

View file

@ -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([

View file

@ -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";

View file

@ -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<object>("@/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(
<QueryClientProvider client={new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } })}>
<LlmCredentialsPanel />
</QueryClientProvider>,
);
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));
});
});

View file

@ -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 <CredentialsPanel />;
const { userRole } = useAuthorized();
return (
<div className="flex flex-col gap-6">
{all_admin_roles.includes(userRole ?? "") && <CredentialsPanel />}
<UserConnectionsPanel />
</div>
);
}

View file

@ -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(<AddModelForm {...props} />);
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", () => {

View file

@ -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"]);
});
});

View file

@ -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<Record<keyof typeof Providers, readonly ProviderAuthType[]>> = {
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",

View file

@ -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<MountedFormValues>();
return <output data-testid="github-copilot-auth-type">{String(watch(GITHUB_COPILOT_AUTH_TYPE_KEY) ?? "")}</output>;
};
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(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ProviderSpecificFields
selectedProvider="GITHUB_COPILOT"
context="model"
initialAuthTypeId="per_user_oauth"
onCreateCredential={onCreateCredential}
/>
<GitHubCopilotAuthTypeProbe />
</MountedFormHost>
</QueryClientProvider>,
);
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");
});
});

View file

@ -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",

View file

@ -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(<GithubCopilotConnectModal credentialName="copilot-per-user" {...props} />);
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();
});
});

View file

@ -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 (
<Dialog open onOpenChange={(open) => !open && onClose()}>
<DialogContent>
<DialogHeader>
<DialogTitle>Connect GitHub Copilot</DialogTitle>
<DialogDescription>Credential: {credentialName}</DialogDescription>
</DialogHeader>
{state.kind === "starting" && <p className="text-sm">Requesting a device code from GitHub...</p>}
{state.kind === "awaiting" && (
<div className="flex flex-col gap-3">
<p className="text-sm">
Open{" "}
<a
href={state.verificationUri}
target="_blank"
rel="noopener noreferrer"
className="text-primary underline-offset-4 hover:underline"
>
{state.verificationUri}
</a>{" "}
and enter this code:
</p>
<div className="flex items-center gap-2">
<code data-testid="github-device-code" className="rounded-md bg-muted px-3 py-2 font-mono text-lg">
{state.userCode}
</code>
<CopyButton value={state.userCode} label="Copy code" copiedLabel="Code copied" />
</div>
<p className="text-sm text-muted-foreground">Waiting for you to approve the request on GitHub...</p>
</div>
)}
{state.kind === "connected" && <p className="text-sm">Connected as @{state.githubLogin}</p>}
{state.kind === "failed" && (
<p role="alert" className="text-sm text-destructive">
{failureMessage(state)}
</p>
)}
<DialogFooter>
{state.kind === "failed" && (
<Button variant="outline" onClick={() => dispatch({ type: "restarted" })}>
Start again
</Button>
)}
{state.kind === "connected" && (
<Button variant="destructive" onClick={onDisconnect}>
Disconnect
</Button>
)}
<Button onClick={onClose}>{state.kind === "connected" ? "Done" : "Close"}</Button>
</DialogFooter>
</DialogContent>
</Dialog>
);
}

View file

@ -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<typeof import("@/components/networking")>();
return {
...actual,
userConnectionsListCall: vi.fn(),
userConnectionDeleteCall: vi.fn(),
};
});
function renderPanel() {
const client = new QueryClient({ defaultOptions: { queries: { retry: false } } });
return render(
<QueryClientProvider client={client}>
<UserConnectionsPanel />
</QueryClientProvider>,
);
}
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);
});
});

View file

@ -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<string | null>(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 (
<section className="mx-auto flex w-full flex-col gap-3 p-2">
<div>
<h3 className="text-base font-semibold">Your connections</h3>
<p className="text-sm text-muted-foreground">
Credentials that use your own account. Connect once, and your requests to models on that credential use your
own access.
</p>
</div>
<Table>
<TableHeader>
<TableRow>
<TableHead>Credential</TableHead>
<TableHead>Provider</TableHead>
<TableHead>Status</TableHead>
<TableHead className="text-right">Actions</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{isLoading && (
<TableRow>
<TableCell colSpan={4} className="text-sm text-muted-foreground">
Loading connections...
</TableCell>
</TableRow>
)}
{isError && (
<TableRow>
<TableCell colSpan={4} className="text-sm">
<div className="flex items-center justify-between gap-2">
<span className="text-muted-foreground">Could not load your connections.</span>
<Button variant="outline" size="sm" onClick={() => refetch()}>
Retry
</Button>
</div>
</TableCell>
</TableRow>
)}
{showEmpty && (
<TableRow>
<TableCell colSpan={4} className="text-sm text-muted-foreground">
No credentials need a personal connection.
</TableCell>
</TableRow>
)}
{connections.map((connection) => (
<TableRow key={connection.credential_name}>
<TableCell>{connection.credential_name}</TableCell>
<TableCell>{providerLabel(connection.provider)}</TableCell>
<TableCell>
{connection.connected ? (
<Badge>Connected as @{connection.github_login}</Badge>
) : (
<Badge variant="outline">Not connected</Badge>
)}
</TableCell>
<TableCell className="text-right">
{connection.connected ? (
<Button variant="outline" size="sm" onClick={() => disconnect(connection.credential_name)}>
Disconnect
</Button>
) : (
<Button size="sm" onClick={() => setConnectingCredential(connection.credential_name)}>
Connect
</Button>
)}
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
{connectingCredential !== null && (
<GithubCopilotConnectModal
credentialName={connectingCredential}
onClose={() => setConnectingCredential(null)}
onConnected={refreshConnections}
onDisconnect={() => disconnect(connectingCredential)}
/>
)}
</section>
);
}

View file

@ -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);
});
});

View file

@ -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<ConnectFlowState, { kind: "awaiting" }>,
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<ConnectFlowState, { kind: "failed" }>): 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.";
}
};

View file

@ -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<UserProviderConnectionsResponse> =>
apiClient.get<UserProviderConnectionsResponse>("/credentials/user_connections", { accessToken });
export const userConnectionStartCall = (
accessToken: string,
credentialName: string,
): Promise<UserConnectionStartResponse> =>
apiClient.post<UserConnectionStartResponse>(`${userConnectionPath(credentialName)}/start`, { accessToken });
export const userConnectionPollCall = (
accessToken: string,
credentialName: string,
flowHandle: string,
): Promise<UserConnectionPollResponse> =>
apiClient.post<UserConnectionPollResponse>(`${userConnectionPath(credentialName)}/poll`, {
accessToken,
body: { flow_handle: flowHandle },
});
export const userConnectionDeleteCall = (accessToken: string, credentialName: string): Promise<void> =>
apiClient.delete<void>(userConnectionPath(credentialName), { accessToken });
export const credentialUpdateCall = async (
accessToken: string,
credentialName: string,

View file

@ -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;