mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
197 lines
7.8 KiB
Python
197 lines
7.8 KiB
Python
import os
|
|
import shutil
|
|
import subprocess
|
|
from typing import Final
|
|
|
|
import pytest
|
|
import test_offline_image_migration
|
|
|
|
offline_postgres: Final = test_offline_image_migration.offline_postgres
|
|
|
|
IMAGE: Final = os.getenv("LITELLM_IMAGE")
|
|
SCHEMA_IMAGE: Final = os.getenv("LITELLM_ADMIN_MCP_SCHEMA_IMAGE")
|
|
COMPONENT: Final = os.getenv("LITELLM_IMAGE_COMPONENT", "unified")
|
|
PROBE: Final = """
|
|
import asyncio
|
|
import base64
|
|
import importlib
|
|
import json
|
|
import sys
|
|
from typing import Final
|
|
|
|
import httpx2
|
|
from cryptography.hazmat.primitives import hashes
|
|
from cryptography.hazmat.primitives.asymmetric import padding, rsa
|
|
from fastapi import HTTPException
|
|
from litellm_admin_mcp.server import create_http_app
|
|
from litellm.proxy import proxy_server
|
|
|
|
module_name: Final = {
|
|
"unified": "litellm.proxy.proxy_server",
|
|
"backend": "backend.main",
|
|
"gateway": "gateway.main",
|
|
}[sys.argv[2]]
|
|
app: Final = importlib.import_module(module_name).app
|
|
|
|
if sys.argv[1] == "base":
|
|
private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
|
message: Final = json.dumps({"expiration_date": "2999-01-01", "user_id": "image-test"}).encode()
|
|
signature: Final = private_key.sign(
|
|
message,
|
|
padding.PSS(mgf=padding.MGF1(hashes.SHA256()), salt_length=padding.PSS.MAX_LENGTH),
|
|
hashes.SHA256(),
|
|
)
|
|
proxy_server._license_check.public_key = private_key.public_key()
|
|
proxy_server._license_check.license_str = base64.b64encode(message + b"." + signature).decode()
|
|
|
|
async def verify_admin_tools(client: httpx2.AsyncClient) -> None:
|
|
user: Final = await client.post(
|
|
"/user/new",
|
|
headers={"Authorization": "Bearer sk-0123456789abcdef0123456789abcdef"},
|
|
json={"user_id": "image-admin", "user_role": "proxy_admin", "auto_create_key": True},
|
|
)
|
|
assert user.status_code == 200, "Admin provisioning failed: " + str(user.status_code)
|
|
headers: Final = {
|
|
"Authorization": "Bearer " + user.json()["key"],
|
|
"Accept": "application/json, text/event-stream",
|
|
}
|
|
discovery: Final = await client.post(
|
|
"/admin/mcp", headers=headers, json={"jsonrpc": "2.0", "id": 1, "method": "tools/list"}
|
|
)
|
|
assert discovery.status_code == 200, discovery.text
|
|
assert {"create_team", "get_team", "delete_teams"} <= {
|
|
tool["name"] for tool in discovery.json()["result"]["tools"]
|
|
}
|
|
|
|
async def call_tool(name: str, arguments: dict[str, object]) -> str:
|
|
response: Final = await client.post(
|
|
"/admin/mcp", headers=headers,
|
|
json={"jsonrpc": "2.0", "id": 2, "method": "tools/call",
|
|
"params": {"name": name, "arguments": arguments}},
|
|
)
|
|
assert response.status_code == 200, response.text
|
|
result: Final = response.json()["result"]
|
|
assert not result.get("isError"), response.text
|
|
return result["content"][0]["text"]
|
|
|
|
team_id: Final = "image-admin-mcp-team"
|
|
created: Final = json.loads(await call_tool(
|
|
"create_team", {"body": {"team_id": team_id, "team_alias": team_id, "max_budget": 25}}
|
|
))
|
|
assert created["team_id"] == team_id and created["max_budget"] == 25, created
|
|
read: Final = json.loads(await call_tool("get_team", {"query": {"team_id": team_id}}))
|
|
assert read["team_info"]["team_id"] == team_id and read["team_info"]["max_budget"] == 25, read
|
|
await call_tool("delete_teams", {"body": {"team_ids": [team_id]}})
|
|
deleted: Final = await client.get("/team/info", headers=headers, params={"team_id": team_id})
|
|
assert deleted.status_code == 404, deleted.text
|
|
print("admin-mcp-tools-ok")
|
|
|
|
async def probe() -> None:
|
|
try:
|
|
async with app.router.lifespan_context(app):
|
|
async with httpx2.AsyncClient(
|
|
transport=httpx2.ASGITransport(app=app), base_url="http://localhost:4000"
|
|
) as client:
|
|
response: Final = await client.post(
|
|
"/admin/mcp",
|
|
json={"jsonrpc": "2.0", "id": 1, "method": "tools/list"},
|
|
headers={"Accept": "application/json, text/event-stream"},
|
|
)
|
|
print("admin-mcp-status=" + str(response.status_code))
|
|
if sys.argv[3] == "tools":
|
|
assert response.status_code == 401, response.text
|
|
await verify_admin_tools(client)
|
|
except HTTPException as exc:
|
|
assert exc.status_code == 403 and "LITELLM_LICENSE" in str(exc.detail)
|
|
assert sys.argv[1] == "none"
|
|
print("admin-mcp-license-required")
|
|
|
|
asyncio.run(probe())
|
|
"""
|
|
|
|
pytestmark = [
|
|
pytest.mark.skipif(IMAGE is None, reason="requires a built image (set LITELLM_IMAGE)"),
|
|
pytest.mark.skipif(shutil.which("docker") is None, reason="requires the docker CLI"),
|
|
]
|
|
|
|
|
|
def _run_probe(
|
|
enabled: str | None,
|
|
license_mode: str,
|
|
network: str = "none",
|
|
database_url: str | None = None,
|
|
) -> subprocess.CompletedProcess[str]:
|
|
assert IMAGE is not None
|
|
enabled_args: Final = () if enabled is None else ("--env", "LITELLM_ENABLE_ADMIN_MCP=" + enabled)
|
|
database_args: Final = () if database_url is None else ("--env", "DATABASE_URL=" + database_url)
|
|
return subprocess.run(
|
|
[
|
|
"docker",
|
|
"run",
|
|
"--rm",
|
|
"--network",
|
|
network,
|
|
"--user",
|
|
"12345:0",
|
|
*enabled_args,
|
|
*database_args,
|
|
"--env",
|
|
"LITELLM_LOCAL_MODEL_COST_MAP=true",
|
|
"--env",
|
|
"LITELLM_MASTER_KEY=sk-0123456789abcdef0123456789abcdef",
|
|
"--entrypoint",
|
|
"python",
|
|
IMAGE,
|
|
"-c",
|
|
PROBE,
|
|
license_mode,
|
|
COMPONENT,
|
|
"tools" if database_url is not None else "visibility",
|
|
],
|
|
capture_output=True,
|
|
text=True,
|
|
timeout=120,
|
|
check=False,
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"enabled,license_mode,expected",
|
|
[
|
|
(None, "none", "status=404"),
|
|
("false", "base", "status=404"),
|
|
("true", "none", "license-required"),
|
|
("true", "base", "status=404" if COMPONENT == "gateway" else "status=401"),
|
|
],
|
|
)
|
|
def test_image_admin_mcp_requires_opt_in_license_and_management_component(
|
|
enabled: str | None, license_mode: str, expected: str
|
|
) -> None:
|
|
result: Final = _run_probe(enabled, license_mode)
|
|
assert result.returncode == 0 and f"admin-mcp-{expected}" in result.stdout, (
|
|
f"Admin MCP image probe failed with component={COMPONENT}, enabled={enabled}, license={license_mode}\n"
|
|
f"{result.stdout}\n{result.stderr}"
|
|
)
|
|
|
|
|
|
@pytest.mark.skipif(COMPONENT == "gateway", reason="the gateway excludes management endpoints")
|
|
def test_image_admin_mcp_personal_admin_manages_team(offline_postgres: tuple[str, str]) -> None:
|
|
assert SCHEMA_IMAGE is not None, "set LITELLM_ADMIN_MCP_SCHEMA_IMAGE to the matching builder image"
|
|
network, postgres = offline_postgres
|
|
database_url: Final = f"postgresql://postgres:pw@{postgres}:5432/litellm"
|
|
schema: Final = subprocess.run(
|
|
[
|
|
"docker", "run", "--rm", "--network", network,
|
|
"--env", "DATABASE_URL=" + database_url,
|
|
"--env", "HOME=/opt/prisma", "--env", "XDG_CACHE_HOME=/opt/prisma/.cache",
|
|
"--env", "PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries",
|
|
"--entrypoint", "prisma", SCHEMA_IMAGE,
|
|
"db", "push", "--schema", "/app/schema.prisma", "--skip-generate", "--accept-data-loss",
|
|
],
|
|
capture_output=True, text=True, timeout=180, check=False,
|
|
)
|
|
assert schema.returncode == 0, f"Schema provisioning failed\n{schema.stdout}\n{schema.stderr}"
|
|
result: Final = _run_probe("true", "base", network, database_url)
|
|
assert result.returncode == 0 and "admin-mcp-tools-ok" in result.stdout, (
|
|
f"Admin MCP management failed in {COMPONENT}\n{result.stdout}\n{result.stderr}"
|
|
)
|