mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
fbf4af301c
commit
d36b66e19e
72 changed files with 8154 additions and 210 deletions
|
|
@ -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");
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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__(
|
||||
|
|
|
|||
|
|
@ -293,6 +293,7 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset(
|
|||
"LiteLLM_BackgroundInteractionSettlement",
|
||||
"LiteLLM_ErrorLogs",
|
||||
"LiteLLM_UserNotifications",
|
||||
"LiteLLM_UserProviderCredentials",
|
||||
"LiteLLM_TeamMembership",
|
||||
"LiteLLM_OrganizationMembership",
|
||||
"LiteLLM_InvitationLink",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
741
litellm/llms/github_copilot/per_user_auth.py
Normal file
741
litellm/llms/github_copilot/per_user_auth.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
351
litellm/proxy/credential_endpoints/user_provider_credentials.py
Normal file
351
litellm/proxy/credential_endpoints/user_provider_credentials.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
1084
tests/unit/llms/github_copilot/test_per_user_auth.py
Normal file
1084
tests/unit/llms/github_copilot/test_per_user_auth.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
});
|
||||
};
|
||||
|
|
@ -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([
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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"]);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
@ -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.";
|
||||
}
|
||||
};
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
253
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
253
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue