diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index c71127a4c3a..765e90bc896 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -2152,6 +2152,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
master_key: Optional[str] = Field(
None, description="require a key for all calls to proxy"
)
+ allow_cli_sso_verification_uri_complete: bool | None = Field(
+ None,
+ description="opt-in to RFC 8628 verification_uri_complete for the CLI SSO device flow, pre-filling the user_code in the browser. Off by default; intended for same-host clients where the device that starts the flow and the browser run on the same machine",
+ )
database_url: Optional[str] = Field(
None,
description="connect to a postgres db - needed for generating temporary keys + tracking spend / key",
diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py
index 2bf12880a75..91a5c109acf 100644
--- a/litellm/proxy/management_endpoints/ui_sso.py
+++ b/litellm/proxy/management_endpoints/ui_sso.py
@@ -145,6 +145,9 @@ _CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS = 60
_CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS = 30
_CLI_SSO_USER_CODE_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
_CLI_SSO_LOGIN_ID_RE = re.compile(r"^cli-[A-Za-z0-9_-]{12,124}$")
+_CLI_SSO_USER_CODE_RE = re.compile(
+ rf"^[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}-[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}$"
+)
_CLI_SSO_SCALAR_TYPES = (str, int, float, bool)
_CLI_SSO_DEST_KEY_RE = re.compile(r"^[A-Za-z0-9_.-]+$")
_CLI_SSO_SECRET_KEY_FRAGMENTS = frozenset(
@@ -182,6 +185,45 @@ def _is_valid_cli_sso_login_id(login_id: Optional[str]) -> bool:
return isinstance(login_id, str) and bool(_CLI_SSO_LOGIN_ID_RE.fullmatch(login_id))
+def _is_valid_cli_sso_user_code(user_code: str | None) -> bool:
+ return isinstance(user_code, str) and bool(
+ _CLI_SSO_USER_CODE_RE.fullmatch(user_code)
+ )
+
+
+def _cli_sso_verification_uri_complete_enabled() -> bool:
+ from litellm.proxy.proxy_server import general_settings
+
+ return bool(
+ general_settings.get( # any-ok: operator opt-in read from the untyped general_settings dict
+ "allow_cli_sso_verification_uri_complete", False
+ )
+ )
+
+
+def _cli_sso_start_response_body(
+ *,
+ login_id: str,
+ poll_secret: str,
+ user_code: str,
+ verification_uri_complete: str | None,
+) -> dict[str, str | int]:
+ if verification_uri_complete is None:
+ return {
+ "login_id": login_id,
+ "poll_secret": poll_secret,
+ "user_code": user_code,
+ "expires_in": CLI_SSO_SESSION_TTL_SECONDS,
+ }
+ return {
+ "login_id": login_id,
+ "poll_secret": poll_secret,
+ "user_code": user_code,
+ "verification_uri_complete": verification_uri_complete,
+ "expires_in": CLI_SSO_SESSION_TTL_SECONDS,
+ }
+
+
def _get_cli_sso_start_rate_limit_cache_key(
request: Request, use_x_forwarded_for: Optional[bool] = False
) -> str:
@@ -478,10 +520,20 @@ def _cli_poll_attribution_metadata_from_session(
def _render_cli_sso_verification_page(
- verify_url: str, browser_complete_token: str
+ verify_url: str,
+ browser_complete_token: str,
+ prefill_user_code: str | None = None,
) -> str:
escaped_verify_url = escape(verify_url, quote=True)
escaped_browser_complete_token = escape(browser_complete_token, quote=True)
+ user_code_value_attr = (
+ f' value="{escape(prefill_user_code, quote=True)}"' if prefill_user_code else ""
+ )
+ instructions = (
+ "Confirm the verification code below to finish this login."
+ if prefill_user_code
+ else "Enter the verification code shown in your terminal to finish this login."
+ )
return f"""
@@ -535,11 +587,11 @@ def _render_cli_sso_verification_page(
Complete CLI Login
- Enter the verification code shown in your terminal to finish this login.
+ {instructions}
@@ -573,12 +625,29 @@ async def cli_sso_start(request: Request):
}
_set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow)
- return {
- "login_id": login_id,
- "poll_secret": poll_secret,
- "user_code": user_code,
- "expires_in": CLI_SSO_SESSION_TTL_SECONDS,
- }
+ verification_uri_complete: str | None = (
+ (
+ get_custom_url(
+ request_base_url=str(request.base_url), route="sso/key/generate"
+ )
+ + "?"
+ + urlencode(
+ {
+ "source": LITELLM_CLI_SOURCE_IDENTIFIER,
+ "key": login_id,
+ "user_code": user_code,
+ }
+ )
+ )
+ if _cli_sso_verification_uri_complete_enabled()
+ else None
+ )
+ return _cli_sso_start_response_body(
+ login_id=login_id,
+ poll_secret=poll_secret,
+ user_code=user_code,
+ verification_uri_complete=verification_uri_complete,
+ )
@router.post(
@@ -829,6 +898,7 @@ async def google_login(
key: Optional[str] = None,
existing_key: Optional[str] = None,
return_to: Optional[str] = None,
+ user_code: str | None = None,
):
"""
Create Proxy API Keys using Google Workspace SSO. Requires setting PROXY_BASE_URL in .env
@@ -897,6 +967,7 @@ async def google_login(
cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state(
source=source,
key=key,
+ user_code=(user_code if _cli_sso_verification_uri_complete_enabled() else None),
)
# check if user defined a custom auth sso sign in handler, if yes, use it
@@ -1921,14 +1992,16 @@ async def auth_callback(request: Request, state: Optional[str] = None):
)
if state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"):
- # State format: {PREFIX}:{login_id}
- state_parts = state.split(":", 1)
+ # State format: {PREFIX}:{login_id}[:{user_code}]
+ state_parts = state.split(":", 2)
key_id = state_parts[1] if len(state_parts) > 1 else None
+ prefill_user_code = state_parts[2] if len(state_parts) > 2 else None
verbose_proxy_logger.info("CLI SSO callback detected")
return await cli_sso_callback(
request=request,
key=key_id,
+ prefill_user_code=prefill_user_code,
result=result,
received_response=received_response,
)
@@ -2008,6 +2081,7 @@ async def _complete_cli_sso_callback_session(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
+ prefill_user_code: str | None = None,
):
from fastapi.responses import HTMLResponse
@@ -2071,6 +2145,7 @@ async def _complete_cli_sso_callback_session(
content=_render_cli_sso_verification_page(
verify_url=verify_url,
browser_complete_token=browser_complete_token,
+ prefill_user_code=prefill_user_code,
),
status_code=200,
)
@@ -2081,6 +2156,7 @@ async def cli_sso_callback(
key: Optional[str] = None,
result: Optional[Union[OpenID, dict]] = None,
received_response: Optional[dict] = None,
+ prefill_user_code: str | None = None,
):
"""CLI SSO callback - stores session info for JWT generation on polling"""
verbose_proxy_logger.info("CLI SSO callback")
@@ -2137,6 +2213,7 @@ async def cli_sso_callback(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
+ prefill_user_code=prefill_user_code,
)
except ProxyException:
raise
@@ -3053,21 +3130,27 @@ class SSOAuthenticationHandler:
@staticmethod
def _get_cli_state(
- source: Optional[str], key: Optional[str], existing_key: Optional[str] = None
+ source: str | None,
+ key: str | None,
+ existing_key: str | None = None,
+ user_code: str | None = None,
) -> Optional[str]:
"""
Checks the request 'source' if a cli state token was passed in
This is used to authenticate through the CLI login flow.
- The state parameter format is: {PREFIX}:{login_id}
+ The state parameter format is: {PREFIX}:{login_id}[:{user_code}]
- The state parameter is used to pass data through the OAuth flow without changing the callback URL
+ - user_code is appended only for the opt-in verification_uri_complete flow so the verify page can pre-fill it
"""
from litellm.constants import (
LITELLM_CLI_SESSION_TOKEN_PREFIX,
)
if source == LITELLM_CLI_SOURCE_IDENTIFIER and key:
+ if _is_valid_cli_sso_user_code(user_code):
+ return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}:{user_code}"
return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}"
else:
return None
diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
index 2efec3e0b34..acca357e641 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
@@ -2120,6 +2120,292 @@ class TestCLIKeyRegenerationFlow:
assert exc_info.value.status_code == 429
mock_cache.set_cache.assert_not_called()
+ @pytest.mark.asyncio
+ async def test_cli_sso_start_returns_verification_uri_complete_when_enabled(self):
+ """Test CLI SSO start returns a verification_uri_complete that round-trips the user_code only when the operator opts in"""
+ from urllib.parse import parse_qs, urlparse
+
+ from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
+ from litellm.proxy.management_endpoints.ui_sso import cli_sso_start
+
+ mock_request = MagicMock(spec=Request)
+ mock_request.client = SimpleNamespace(host="127.0.0.1")
+ mock_request.headers = {}
+ mock_request.base_url = "https://proxy.example.com/"
+ mock_cache = MagicMock()
+ mock_cache.increment_cache.return_value = 1
+
+ with (
+ patch.dict(
+ os.environ,
+ {"PROXY_BASE_URL": "https://proxy.example.com", "SERVER_ROOT_PATH": ""},
+ ),
+ patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
+ patch(
+ "litellm.proxy.proxy_server.general_settings",
+ {"allow_cli_sso_verification_uri_complete": True},
+ ),
+ ):
+ result = await cli_sso_start(request=mock_request)
+
+ verification_uri_complete = result["verification_uri_complete"]
+ parsed = urlparse(verification_uri_complete)
+ query = parse_qs(parsed.query)
+
+ assert parsed.path.endswith("/sso/key/generate")
+ assert query["source"] == [LITELLM_CLI_SOURCE_IDENTIFIER]
+ assert query["key"] == [result["login_id"]]
+ assert query["user_code"] == [result["user_code"]]
+
+ @pytest.mark.asyncio
+ async def test_cli_sso_start_omits_verification_uri_complete_by_default(self):
+ """Test CLI SSO start does NOT advertise verification_uri_complete unless the operator enables it (default off)"""
+ from litellm.proxy.management_endpoints.ui_sso import cli_sso_start
+
+ mock_request = MagicMock(spec=Request)
+ mock_request.client = SimpleNamespace(host="127.0.0.1")
+ mock_request.headers = {}
+ mock_request.base_url = "https://proxy.example.com/"
+ mock_cache = MagicMock()
+ mock_cache.increment_cache.return_value = 1
+
+ with (
+ patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
+ patch("litellm.proxy.proxy_server.general_settings", {}),
+ ):
+ result = await cli_sso_start(request=mock_request)
+
+ assert "verification_uri_complete" not in result
+ assert result["user_code"]
+ assert result["login_id"].startswith("cli-")
+
+ def test_cli_sso_verification_uri_complete_enabled_reads_general_settings(self):
+ """Test the operator opt-in flag is read from general_settings and defaults off"""
+ from litellm.proxy.management_endpoints.ui_sso import (
+ _cli_sso_verification_uri_complete_enabled,
+ )
+
+ with patch("litellm.proxy.proxy_server.general_settings", {}):
+ assert _cli_sso_verification_uri_complete_enabled() is False
+ with patch(
+ "litellm.proxy.proxy_server.general_settings",
+ {"allow_cli_sso_verification_uri_complete": True},
+ ):
+ assert _cli_sso_verification_uri_complete_enabled() is True
+
+ @pytest.mark.asyncio
+ async def test_google_login_only_threads_user_code_when_enabled(self):
+ """Test google_login forwards user_code into the OAuth state only when the operator opt-in is on, dropping it otherwise"""
+ from litellm.proxy.management_endpoints.ui_sso import google_login
+
+ mock_request = MagicMock(spec=Request)
+ mock_request.base_url = "https://proxy.example.com/"
+ mock_cache = MagicMock()
+ mock_cache.get_cache.return_value = {"poll_secret_hash": "h"}
+
+ async def drive(enabled: bool):
+ with (
+ patch.dict(os.environ, {}, clear=True),
+ patch("litellm.proxy.proxy_server.premium_user", True),
+ patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
+ patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
+ patch(
+ "litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler",
+ None,
+ ),
+ patch(
+ "litellm.proxy.proxy_server.general_settings",
+ {"allow_cli_sso_verification_uri_complete": enabled},
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.ui_sso.show_missing_vars_in_env",
+ return_value=None,
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.get_redirect_url_for_sso",
+ return_value="https://proxy.example.com/sso/callback",
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler._get_cli_state",
+ return_value=None,
+ ) as mock_get_cli_state,
+ ):
+ try:
+ await google_login(
+ request=mock_request,
+ source="litellm-cli",
+ key="cli-validsessionkey123456",
+ user_code="WXYZ-2345",
+ )
+ except Exception:
+ pass
+ return mock_get_cli_state.call_args.kwargs["user_code"]
+
+ assert await drive(enabled=True) == "WXYZ-2345"
+ assert await drive(enabled=False) is None
+
+ def test_get_cli_state_appends_user_code_for_prefill(self):
+ """Test the OAuth state carries the user_code only for the opt-in prefill flow"""
+ from litellm.constants import (
+ LITELLM_CLI_SESSION_TOKEN_PREFIX,
+ LITELLM_CLI_SOURCE_IDENTIFIER,
+ )
+ from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
+
+ manual_state = SSOAuthenticationHandler._get_cli_state(
+ source=LITELLM_CLI_SOURCE_IDENTIFIER, key="cli-abc123"
+ )
+ prefill_state = SSOAuthenticationHandler._get_cli_state(
+ source=LITELLM_CLI_SOURCE_IDENTIFIER,
+ key="cli-abc123",
+ user_code="WXYZ-2345",
+ )
+
+ assert manual_state == f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123"
+ assert (
+ prefill_state == f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123:WXYZ-2345"
+ )
+ assert (
+ SSOAuthenticationHandler._get_cli_state(
+ source="not-cli", key="cli-abc123", user_code="WXYZ-2345"
+ )
+ is None
+ )
+
+ def test_get_cli_state_drops_malformed_user_code(self):
+ """Test a user_code that is not a server-issued code is dropped before reaching the size-limited OAuth state"""
+ from litellm.constants import (
+ LITELLM_CLI_SESSION_TOKEN_PREFIX,
+ LITELLM_CLI_SOURCE_IDENTIFIER,
+ )
+ from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
+
+ manual_only = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123"
+ for bad_user_code in ("A" * 4096, "not-a-code", "WXYZ2345", "WXYZ-234", ""):
+ assert (
+ SSOAuthenticationHandler._get_cli_state(
+ source=LITELLM_CLI_SOURCE_IDENTIFIER,
+ key="cli-abc123",
+ user_code=bad_user_code,
+ )
+ == manual_only
+ )
+
+ def test_is_valid_cli_sso_user_code_matches_generated_format(self):
+ """Test the user_code validator accepts a freshly generated code and rejects malformed input"""
+ from litellm.proxy.management_endpoints.ui_sso import (
+ _generate_cli_sso_user_code,
+ _is_valid_cli_sso_user_code,
+ )
+
+ assert _is_valid_cli_sso_user_code(_generate_cli_sso_user_code())
+ assert _is_valid_cli_sso_user_code("WXYZ-2345")
+ assert not _is_valid_cli_sso_user_code("WXYZ-2340") # 0 is not in the alphabet
+ assert not _is_valid_cli_sso_user_code("wxyz-2345")
+ assert not _is_valid_cli_sso_user_code("WXYZ2345")
+ assert not _is_valid_cli_sso_user_code("A" * 64)
+ assert not _is_valid_cli_sso_user_code(None)
+
+ def test_cli_state_round_trips_user_code_to_callback_parser(self):
+ """Test the callback's state parser recovers login_id and user_code from the state _get_cli_state builds"""
+ from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
+ from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
+
+ state = SSOAuthenticationHandler._get_cli_state(
+ source=LITELLM_CLI_SOURCE_IDENTIFIER,
+ key="cli-abc123",
+ user_code="WXYZ-2345",
+ )
+
+ state_parts = state.split(":", 2)
+ key_id = state_parts[1] if len(state_parts) > 1 else None
+ prefill_user_code = state_parts[2] if len(state_parts) > 2 else None
+
+ assert key_id == "cli-abc123"
+ assert prefill_user_code == "WXYZ-2345"
+
+ def test_render_cli_sso_verification_page_prefills_user_code(self):
+ """Test the verify page pre-fills the user_code input (HTML-escaped) when provided"""
+ from litellm.proxy.management_endpoints.ui_sso import (
+ _render_cli_sso_verification_page,
+ )
+
+ html = _render_cli_sso_verification_page(
+ verify_url="https://proxy.example.com/sso/cli/complete/cli-abc123",
+ browser_complete_token="browser-token",
+ prefill_user_code='WXYZ-2345">