diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 611a0e89bf4..39e015f3928 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -16,6 +16,7 @@ import os import secrets from copy import deepcopy from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast +from urllib.parse import urlencode, urlparse if TYPE_CHECKING: import httpx @@ -403,6 +404,7 @@ async def google_login( state=cli_state, ) if return_to is not None and sso_redirect is not None: + SSOAuthenticationHandler._validate_return_to(return_to) sso_redirect.set_cookie( key="litellm_cp_return_to", value=return_to, @@ -1771,6 +1773,30 @@ class SSOAuthenticationHandler: Handler for SSO Authentication across all SSO providers """ + @staticmethod + def _validate_return_to(return_to: str) -> None: + """ + Validate that return_to matches the configured control_plane_url origin. + + Raises HTTPException(400) if: + - control_plane_url is not configured in general_settings + - return_to origin does not match control_plane_url origin + """ + from litellm.proxy.proxy_server import general_settings + + control_plane_url = general_settings.get("control_plane_url") + if control_plane_url is None: + raise HTTPException( + status_code=400, + detail="return_to is not allowed: control_plane_url is not configured", + ) + + if urlparse(return_to).hostname != urlparse(control_plane_url).hostname: + raise HTTPException( + status_code=400, + detail="return_to does not match the configured control_plane_url", + ) + @staticmethod async def get_sso_login_redirect( redirect_url: str, @@ -2550,14 +2576,7 @@ class SSOAuthenticationHandler: # Control-plane cross-origin: redirect back to the control plane UI # with the token in the URL (cookie won't work cross-origin) if return_to is not None: - from urllib.parse import urlencode, urlparse - - parsed = urlparse(return_to) - if parsed.scheme not in ("http", "https"): - raise HTTPException( - status_code=400, - detail="return_to must be an HTTP(S) URL", - ) + SSOAuthenticationHandler._validate_return_to(return_to) separator = "&" if "?" in return_to else "?" redirect_url = ( return_to @@ -2565,7 +2584,7 @@ class SSOAuthenticationHandler: + urlencode({"login": "success", "token": jwt_token}) ) verbose_proxy_logger.info( - f"Cross-origin SSO: redirecting to control plane at {parsed.netloc}" + "Cross-origin SSO: redirecting to control plane" ) redirect_response = RedirectResponse(url=redirect_url, status_code=303) redirect_response.delete_cookie("litellm_cp_return_to") diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 48d9d18fe46..73b224d4bba 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11080,6 +11080,14 @@ async def login_v3(request: Request): # noqa: PLR0915 from litellm.proxy.utils import get_custom_url try: + if not general_settings.get("control_plane_url"): + raise ProxyException( + message="/v3/login 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() username = str(body.get("username")) password = str(body.get("password")) 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 d43b2c4ba05..4d531da67ca 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest -from fastapi import Request +from fastapi import HTTPException, Request from litellm._uuid import uuid @@ -5160,3 +5160,63 @@ def test_generic_response_convertor_extra_attributes_missing_field(monkeypatch): assert result.extra_fields["missing_field"] is None assert result.extra_fields["another_missing"] is None + +class TestValidateReturnTo: + """Tests for SSOAuthenticationHandler._validate_return_to""" + + def test_rejects_when_no_control_plane_url_configured(self, monkeypatch): + """return_to should be rejected if control_plane_url is not in general_settings.""" + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", {} + ) + with pytest.raises(HTTPException) as exc_info: + SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui") + assert exc_info.value.status_code == 400 + assert "not configured" in exc_info.value.detail + + def test_allows_matching_origin(self, monkeypatch): + """return_to matching the configured control_plane_url origin should pass.""" + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", + {"control_plane_url": "https://cp.example.com"}, + ) + # Should not raise + SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui?page=models") + + def test_allows_matching_origin_with_trailing_slash(self, monkeypatch): + """Trailing slash on control_plane_url should not affect origin comparison.""" + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", + {"control_plane_url": "https://cp.example.com/"}, + ) + SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui") + + def test_rejects_prefix_attack(self, monkeypatch): + """return_to like cp.example.com.evil.com must be rejected (not just prefix match).""" + 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.evil.com/steal") + assert exc_info.value.status_code == 400 + + def test_rejects_different_origin(self, monkeypatch): + """return_to pointing to a completely different domain should 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://evil.com/phish") + assert exc_info.value.status_code == 400 + + def test_case_insensitive_hostname(self, monkeypatch): + """Hostname comparison should be case-insensitive per RFC 3986.""" + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", + {"control_plane_url": "https://CP.Example.COM"}, + ) + # Should not raise + SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui") + diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index e6269936b5c..f23f5a39353 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -236,8 +236,25 @@ def test_login_v2_returns_json_on_invalid_json_body(monkeypatch): assert isinstance(data["error"], dict) -def test_login_v3_always_includes_token_in_body(monkeypatch): - """v3/login always returns token in body, even without worker_registry.""" +def test_login_v3_rejected_without_control_plane_url(monkeypatch): + """v3/login returns 404 when control_plane_url is not configured.""" + mock_prisma_client = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key") + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + client = TestClient(app) + response = client.post( + "/v3/login", + json={"username": "alice", "password": "secret"}, + ) + + assert response.status_code == 404 + 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.""" mock_prisma_client = MagicMock() monkeypatch.setattr( "litellm.proxy.auth.login_utils.authenticate_user", @@ -249,7 +266,10 @@ def test_login_v3_always_includes_token_in_body(monkeypatch): ) 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", {}) + 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() @@ -289,6 +309,10 @@ def test_login_v3_returns_json_on_proxy_exception(monkeypatch): mock_authenticate_user, ) 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.prisma_client", mock_prisma_client) client = TestClient(app)