mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge litellm_control-plane-backend into frontend branch
Brings in: control_plane_url gate on /v3/login, _validate_return_to with RFC 3986 case-insensitive hostname comparison, and SSO return_to validation.
This commit is contained in:
commit
c1327dd386
4 changed files with 124 additions and 13 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue