mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge branch 'litellm_control-plane-backend' into litellm_control-plane-frontend
This commit is contained in:
commit
0e3705d63a
4 changed files with 286 additions and 14 deletions
|
|
@ -1791,7 +1791,15 @@ class SSOAuthenticationHandler:
|
|||
detail="return_to is not allowed: control_plane_url is not configured",
|
||||
)
|
||||
|
||||
if urlparse(return_to).hostname != urlparse(control_plane_url).hostname:
|
||||
def _origin(url: str) -> tuple:
|
||||
parsed = urlparse(url)
|
||||
scheme = (parsed.scheme or "").lower()
|
||||
hostname = (parsed.hostname or "").lower()
|
||||
default_port = 443 if scheme == "https" else 80
|
||||
port = parsed.port if parsed.port is not None else default_port
|
||||
return (scheme, hostname, port)
|
||||
|
||||
if _origin(return_to) != _origin(control_plane_url):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="return_to does not match the configured control_plane_url",
|
||||
|
|
@ -2405,6 +2413,7 @@ class SSOAuthenticationHandler:
|
|||
master_key,
|
||||
premium_user,
|
||||
proxy_logging_obj,
|
||||
redis_usage_cache,
|
||||
user_api_key_cache,
|
||||
user_custom_sso,
|
||||
)
|
||||
|
|
@ -2573,18 +2582,32 @@ class SSOAuthenticationHandler:
|
|||
algorithm="HS256",
|
||||
)
|
||||
|
||||
# Control-plane cross-origin: redirect back to the control plane UI
|
||||
# with the token in the URL (cookie won't work cross-origin)
|
||||
# Control-plane cross-origin: store JWT behind a single-use opaque
|
||||
# code (60s TTL) so the token never appears in browser history / logs.
|
||||
# The control plane redeems it via POST /v3/login/exchange.
|
||||
if return_to is not None:
|
||||
SSOAuthenticationHandler._validate_return_to(return_to)
|
||||
|
||||
code = secrets.token_urlsafe(32)
|
||||
cache_key = f"login_code:{code}"
|
||||
cache_value = {"token": jwt_token, "redirect_url": return_to}
|
||||
if redis_usage_cache is not None:
|
||||
await redis_usage_cache.async_set_cache(
|
||||
key=cache_key, value=cache_value, ttl=60
|
||||
)
|
||||
else:
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=cache_key, value=cache_value, ttl=60
|
||||
)
|
||||
|
||||
separator = "&" if "?" in return_to else "?"
|
||||
redirect_url = (
|
||||
return_to
|
||||
+ separator
|
||||
+ urlencode({"login": "success", "token": jwt_token})
|
||||
+ urlencode({"login": "success", "code": code})
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"Cross-origin SSO: redirecting to control plane"
|
||||
"Cross-origin SSO: redirecting to control plane with login code"
|
||||
)
|
||||
redirect_response = RedirectResponse(url=redirect_url, status_code=303)
|
||||
redirect_response.delete_cookie("litellm_cp_return_to")
|
||||
|
|
|
|||
|
|
@ -1546,6 +1546,7 @@ user_custom_key_generate = None
|
|||
# Sentinel: prevents PKCE-no-Redis advisory from re-logging on config hot-reload.
|
||||
# Tests that need to reset it can patch 'litellm.proxy.proxy_server._pkce_no_redis_warning_emitted'.
|
||||
_pkce_no_redis_warning_emitted: bool = False
|
||||
_cp_no_redis_warning_emitted: bool = False
|
||||
user_custom_sso = None
|
||||
user_custom_ui_sso_sign_in_handler = None
|
||||
use_background_health_checks = None
|
||||
|
|
@ -3096,6 +3097,21 @@ class ProxyConfig:
|
|||
"Set PKCE_STRICT_CACHE_MISS=true to fail fast with a 401 on cache misses "
|
||||
"instead of continuing without a code_verifier."
|
||||
)
|
||||
|
||||
### CONTROL PLANE CODE-EXCHANGE PREREQUISITE CHECK ###
|
||||
cp_url = general_settings.get("control_plane_url")
|
||||
if cp_url and redis_usage_cache is None:
|
||||
global _cp_no_redis_warning_emitted
|
||||
if not _cp_no_redis_warning_emitted:
|
||||
_cp_no_redis_warning_emitted = True
|
||||
verbose_proxy_logger.warning(
|
||||
"control_plane_url is configured but Redis is not configured for LiteLLM caching. "
|
||||
"Login codes (SSO and /v3/login) will not be shared across instances — "
|
||||
"the /v3/login/exchange call may land on a different pod and fail with 401. "
|
||||
"Configure Redis via the 'cache' section in your proxy config, "
|
||||
"or ensure sticky sessions for single-instance deployments."
|
||||
)
|
||||
|
||||
### STORE MODEL IN DB ### feature flag for `/model/new`
|
||||
store_model_in_db = general_settings.get("store_model_in_db", False)
|
||||
if store_model_in_db is None:
|
||||
|
|
@ -11120,12 +11136,23 @@ async def login_v3(request: Request): # noqa: PLR0915
|
|||
litellm_dashboard_ui += "/ui/"
|
||||
litellm_dashboard_ui += "?login=success"
|
||||
|
||||
json_response = JSONResponse(
|
||||
content={"redirect_url": litellm_dashboard_ui, "token": jwt_token},
|
||||
# Store JWT behind a single-use opaque code (60s TTL)
|
||||
code = secrets.token_urlsafe(32)
|
||||
cache_key = f"login_code:{code}"
|
||||
cache_value = {"token": jwt_token, "redirect_url": litellm_dashboard_ui}
|
||||
if redis_usage_cache is not None:
|
||||
await redis_usage_cache.async_set_cache(
|
||||
key=cache_key, value=cache_value, ttl=60
|
||||
)
|
||||
else:
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=cache_key, value=cache_value, ttl=60
|
||||
)
|
||||
|
||||
return JSONResponse(
|
||||
content={"code": code, "expires_in": 60},
|
||||
status_code=status.HTTP_200_OK,
|
||||
)
|
||||
json_response.set_cookie(key="token", value=jwt_token)
|
||||
return json_response
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.proxy_server.login_v3(): Exception occurred - {}".format(
|
||||
|
|
@ -11151,6 +11178,74 @@ async def login_v3(request: Request): # noqa: PLR0915
|
|||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v3/login/exchange", include_in_schema=False
|
||||
) # exchange single-use opaque code for JWT
|
||||
async def login_v3_exchange(request: Request):
|
||||
try:
|
||||
if not general_settings.get("control_plane_url"):
|
||||
raise ProxyException(
|
||||
message="/v3/login/exchange is only available on workers with control_plane_url configured",
|
||||
type=ProxyErrorTypes.not_found_error,
|
||||
param="control_plane_url",
|
||||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
|
||||
body = await request.json()
|
||||
code = body.get("code")
|
||||
if not code:
|
||||
raise ProxyException(
|
||||
message="Missing 'code' parameter",
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param="code",
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
cache_key = f"login_code:{code}"
|
||||
if redis_usage_cache is not None:
|
||||
cached_data = await redis_usage_cache.async_get_cache(key=cache_key)
|
||||
else:
|
||||
cached_data = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
|
||||
if not cached_data or not isinstance(cached_data, dict):
|
||||
raise ProxyException(
|
||||
message="Invalid or expired login code",
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param="code",
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
)
|
||||
|
||||
# Single-use: delete immediately
|
||||
if redis_usage_cache is not None:
|
||||
await redis_usage_cache.async_delete_cache(key=cache_key)
|
||||
else:
|
||||
await user_api_key_cache.async_delete_cache(key=cache_key)
|
||||
|
||||
json_response = JSONResponse(
|
||||
content={
|
||||
"token": cached_data["token"],
|
||||
"redirect_url": cached_data["redirect_url"],
|
||||
},
|
||||
status_code=status.HTTP_200_OK,
|
||||
)
|
||||
json_response.set_cookie(key="token", value=cached_data["token"])
|
||||
return json_response
|
||||
except ProxyException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.proxy_server.login_v3_exchange(): Exception occurred - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
raise ProxyException(
|
||||
message=str(e),
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param="None",
|
||||
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
)
|
||||
|
||||
|
||||
@app.get("/onboarding/get_token", include_in_schema=False)
|
||||
async def onboarding(invite_link: str, request: Request):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -5220,3 +5220,39 @@ class TestValidateReturnTo:
|
|||
# Should not raise
|
||||
SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui")
|
||||
|
||||
def test_rejects_scheme_mismatch(self, monkeypatch):
|
||||
"""http:// must be rejected when control_plane_url uses https://."""
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"control_plane_url": "https://cp.example.com"},
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
SSOAuthenticationHandler._validate_return_to("http://cp.example.com/ui")
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
def test_rejects_port_mismatch(self, monkeypatch):
|
||||
"""Non-default port must be rejected."""
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"control_plane_url": "https://cp.example.com"},
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
SSOAuthenticationHandler._validate_return_to("https://cp.example.com:8443/ui")
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
def test_allows_explicit_default_port(self, monkeypatch):
|
||||
"""https://host:443 should match https://host (default port normalisation)."""
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"control_plane_url": "https://cp.example.com"},
|
||||
)
|
||||
SSOAuthenticationHandler._validate_return_to("https://cp.example.com:443/ui")
|
||||
|
||||
def test_allows_matching_custom_port(self, monkeypatch):
|
||||
"""Both sides on the same custom port should match."""
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"control_plane_url": "https://cp.example.com:3000"},
|
||||
)
|
||||
SSOAuthenticationHandler._validate_return_to("https://cp.example.com:3000/ui")
|
||||
|
||||
|
|
|
|||
|
|
@ -253,8 +253,8 @@ def test_login_v3_rejected_without_control_plane_url(monkeypatch):
|
|||
assert "control_plane_url" in response.json()["error"]["message"]
|
||||
|
||||
|
||||
def test_login_v3_includes_token_in_body(monkeypatch):
|
||||
"""v3/login returns token in body when control_plane_url is configured."""
|
||||
def test_login_v3_returns_code(monkeypatch):
|
||||
"""v3/login returns an opaque code, not the JWT directly."""
|
||||
mock_prisma_client = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||||
|
|
@ -285,9 +285,127 @@ def test_login_v3_includes_token_in_body(monkeypatch):
|
|||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["token"] == "signed-token"
|
||||
assert response.json()["redirect_url"] == "http://testserver/ui/?login=success"
|
||||
assert response.cookies.get("token") == "signed-token"
|
||||
data = response.json()
|
||||
assert "code" in data
|
||||
assert data["expires_in"] == 60
|
||||
assert "token" not in data
|
||||
|
||||
|
||||
def test_login_v3_exchange_happy_path(monkeypatch):
|
||||
"""Full flow: v3/login returns code, v3/login/exchange redeems it for JWT."""
|
||||
mock_prisma_client = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||||
AsyncMock(return_value={"user_id": "test-user"}),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.login_utils.create_ui_token_object",
|
||||
MagicMock(return_value={"user_id": "test-user"}),
|
||||
)
|
||||
monkeypatch.setattr("jwt.encode", MagicMock(return_value="signed-token"))
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"control_plane_url": "https://cp.example.com"},
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
mock_config = MagicMock()
|
||||
mock_config.worker_registry = []
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", mock_config)
|
||||
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
|
||||
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
|
||||
|
||||
client = TestClient(app)
|
||||
|
||||
# Step 1: login — get code
|
||||
login_response = client.post(
|
||||
"/v3/login",
|
||||
json={"username": "alice", "password": "secret"},
|
||||
)
|
||||
assert login_response.status_code == 200
|
||||
code = login_response.json()["code"]
|
||||
|
||||
# Step 2: exchange — get JWT
|
||||
exchange_response = client.post(
|
||||
"/v3/login/exchange",
|
||||
json={"code": code},
|
||||
)
|
||||
assert exchange_response.status_code == 200
|
||||
exchange_data = exchange_response.json()
|
||||
assert exchange_data["token"] == "signed-token"
|
||||
assert "redirect_url" in exchange_data
|
||||
assert exchange_response.cookies.get("token") == "signed-token"
|
||||
|
||||
|
||||
def test_login_v3_exchange_single_use(monkeypatch):
|
||||
"""Code can only be redeemed once."""
|
||||
mock_prisma_client = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||||
AsyncMock(return_value={"user_id": "test-user"}),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.login_utils.create_ui_token_object",
|
||||
MagicMock(return_value={"user_id": "test-user"}),
|
||||
)
|
||||
monkeypatch.setattr("jwt.encode", MagicMock(return_value="signed-token"))
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"control_plane_url": "https://cp.example.com"},
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
mock_config = MagicMock()
|
||||
mock_config.worker_registry = []
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", mock_config)
|
||||
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
|
||||
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
|
||||
|
||||
client = TestClient(app)
|
||||
|
||||
login_response = client.post(
|
||||
"/v3/login",
|
||||
json={"username": "alice", "password": "secret"},
|
||||
)
|
||||
code = login_response.json()["code"]
|
||||
|
||||
# First exchange succeeds
|
||||
first = client.post("/v3/login/exchange", json={"code": code})
|
||||
assert first.status_code == 200
|
||||
|
||||
# Second exchange fails
|
||||
second = client.post("/v3/login/exchange", json={"code": code})
|
||||
assert second.status_code == 401
|
||||
|
||||
|
||||
def test_login_v3_exchange_invalid_code(monkeypatch):
|
||||
"""Random code returns 401."""
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"control_plane_url": "https://cp.example.com"},
|
||||
)
|
||||
client = TestClient(app)
|
||||
response = client.post(
|
||||
"/v3/login/exchange",
|
||||
json={"code": "nonexistent-code"},
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_login_v3_exchange_rejected_without_control_plane_url(monkeypatch):
|
||||
"""v3/login/exchange returns 404 when control_plane_url is not configured."""
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||||
|
||||
client = TestClient(app)
|
||||
response = client.post(
|
||||
"/v3/login/exchange",
|
||||
json={"code": "some-code"},
|
||||
)
|
||||
|
||||
assert response.status_code == 404
|
||||
assert "control_plane_url" in response.json()["error"]["message"]
|
||||
|
||||
|
||||
def test_login_v3_returns_json_on_proxy_exception(monkeypatch):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue