mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Add /v3/login endpoint for cross-origin control-plane login
Workers don't have worker_registry, so /v2/login omits the token from the response body. The control plane UI needs the token in the body to set the cookie cross-origin (document.cookie) when authenticating against a worker. /v3/login always includes it.
This commit is contained in:
parent
6686a1213f
commit
96cd74440c
2 changed files with 141 additions and 0 deletions
|
|
@ -11075,6 +11075,78 @@ async def login_v2(request: Request): # noqa: PLR0915
|
|||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v3/login", include_in_schema=False
|
||||
) # control-plane login — always returns token in body for cross-origin use
|
||||
async def login_v3(request: Request): # noqa: PLR0915
|
||||
global premium_user, general_settings, master_key
|
||||
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object
|
||||
from litellm.proxy.utils import get_custom_url
|
||||
|
||||
try:
|
||||
body = await request.json()
|
||||
username = str(body.get("username"))
|
||||
password = str(body.get("password"))
|
||||
|
||||
login_result = await authenticate_user(
|
||||
username=username,
|
||||
password=password,
|
||||
master_key=master_key,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
returned_ui_token_object = create_ui_token_object(
|
||||
login_result=login_result,
|
||||
general_settings=general_settings,
|
||||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
import jwt
|
||||
|
||||
jwt_token = jwt.encode(
|
||||
cast(dict, returned_ui_token_object),
|
||||
cast(str, master_key),
|
||||
algorithm="HS256",
|
||||
)
|
||||
|
||||
litellm_dashboard_ui = get_custom_url(str(request.base_url))
|
||||
if litellm_dashboard_ui.endswith("/"):
|
||||
litellm_dashboard_ui += "ui/"
|
||||
else:
|
||||
litellm_dashboard_ui += "/ui/"
|
||||
litellm_dashboard_ui += "?login=success"
|
||||
|
||||
json_response = JSONResponse(
|
||||
content={"redirect_url": litellm_dashboard_ui, "token": jwt_token},
|
||||
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(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
if isinstance(e, ProxyException):
|
||||
raise e
|
||||
elif isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "detail", str(e)),
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR),
|
||||
)
|
||||
else:
|
||||
error_msg = f"{str(e)}"
|
||||
raise ProxyException(
|
||||
message=error_msg,
|
||||
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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -277,6 +277,75 @@ 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."""
|
||||
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", {})
|
||||
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)
|
||||
response = client.post(
|
||||
"/v3/login",
|
||||
json={"username": "alice", "password": "secret"},
|
||||
)
|
||||
|
||||
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"
|
||||
|
||||
|
||||
def test_login_v3_returns_json_on_proxy_exception(monkeypatch):
|
||||
"""Test that /v3/login returns JSON error when ProxyException is raised"""
|
||||
from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_authenticate_user = AsyncMock(
|
||||
side_effect=ProxyException(
|
||||
message="Invalid credentials",
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param="password",
|
||||
code=401,
|
||||
)
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||||
mock_authenticate_user,
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
client = TestClient(app)
|
||||
response = client.post(
|
||||
"/v3/login",
|
||||
json={"username": "alice", "password": "wrong"},
|
||||
)
|
||||
|
||||
assert response.status_code == 401
|
||||
assert response.headers["content-type"] == "application/json"
|
||||
data = response.json()
|
||||
assert "error" in data
|
||||
assert data["error"]["message"] == "Invalid credentials"
|
||||
assert data["error"]["type"] == "auth_error"
|
||||
|
||||
|
||||
def test_fallback_login_has_no_deprecation_banner(client_no_auth):
|
||||
response = client_no_auth.get("/fallback/login")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue