Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/competent-bassi-2731b1

This commit is contained in:
Yuneng Jiang 2026-06-08 17:08:41 -07:00
commit e7113dd8b3
No known key found for this signature in database
14 changed files with 319 additions and 16 deletions

View file

@ -80,11 +80,15 @@ class FocusLiteLLMDatabase:
vt.team_id,
vt.key_alias as api_key_alias,
tt.team_alias,
ut.user_email as user_email
ut.user_email as user_email,
COALESCE(vt.organization_id, tt.organization_id) as organization_id,
ot.organization_alias as organization_alias
FROM "LiteLLM_DailyUserSpend" dus
LEFT JOIN "LiteLLM_VerificationToken" vt ON dus.api_key = vt.token
LEFT JOIN "LiteLLM_TeamTable" tt ON vt.team_id = tt.team_id
LEFT JOIN "LiteLLM_UserTable" ut ON dus.user_id = ut.user_id
LEFT JOIN "LiteLLM_OrganizationTable" ot
ON ot.organization_id = COALESCE(vt.organization_id, tt.organization_id)
{where_clause}
ORDER BY dus.date DESC, dus.created_at DESC
{limit_clause}

View file

@ -12,6 +12,8 @@ from .schema import FOCUS_NORMALIZED_SCHEMA
_TAG_KEYS = (
"team_id",
"team_alias",
"organization_id",
"organization_alias",
"user_id",
"user_email",
"api_key_alias",

View file

@ -561,7 +561,7 @@ async def _setup_new_team_model_assignment(
async def _get_team_deployments(
team_id: str, prisma_client: PrismaClient
team_id: str, prisma_client: PrismaClient, table: Optional[Any] = None
) -> List[LiteLLM_ProxyModelTable]:
"""
Fetch all deployments for a given team_id from the database.
@ -572,9 +572,13 @@ async def _get_team_deployments(
Note: prisma-client-py 0.11.0 does not support JSON path filtering, so we filter
by the model_name prefix (team models use "model_name_{team_id}_*") and confirm
team_id in model_info with Python-side filtering.
Pass ``table`` (a transaction's proxy-model table) to run the read inside an
existing transaction.
"""
prefix = f"model_name_{team_id}_"
response = await ModelRepository(prisma_client).table.find_many(
table = table or ModelRepository(prisma_client).table
response = await table.find_many(
where={
"model_name": {"startswith": prefix},
}
@ -596,6 +600,42 @@ async def _get_team_deployments(
return result
async def delete_team_models(
team_ids: List[str],
prisma_client: PrismaClient,
llm_router: Optional[Any],
) -> List[str]:
"""
Delete every BYOK model owned by the given teams, from the DB and the router.
The DB rows are removed inside a single transaction, so deletion is atomic
across all team_ids. Each team's rows are deleted by the exact model_ids read
in the same transaction, which keeps the deleted set identical to the set
handed to the router. The router is synced only after the transaction commits,
so a rollback can never leave a deployment live in the router without its row.
Returns the model_ids that were deleted.
"""
deleted_model_ids: List[str] = []
async with prisma_client.db.tx() as tx:
for team_id in team_ids:
rows = await _get_team_deployments(
team_id, prisma_client, table=tx.litellm_proxymodeltable
)
model_ids = [row.model_id for row in rows]
if model_ids:
await tx.litellm_proxymodeltable.delete_many(
where={"model_id": {"in": model_ids}}
)
deleted_model_ids.extend(model_ids)
if llm_router is not None:
for model_id in deleted_model_ids:
llm_router.delete_deployment(id=model_id)
return deleted_model_ids
async def _get_team_public_model_names(
team_id: str,
prisma_client: PrismaClient,

View file

@ -3289,6 +3289,20 @@ async def delete_team(
await prisma_client.delete_data(team_id_list=data.team_ids, table_name="key")
## DELETE ASSOCIATED BYOK MODELS
# Runs before the team rows are deleted so a mid-flight failure never leaves
# the team gone with its models orphaned.
from litellm.proxy.management_endpoints.model_management_endpoints import (
delete_team_models,
)
from litellm.proxy.proxy_server import llm_router
await delete_team_models(
team_ids=data.team_ids,
prisma_client=prisma_client,
llm_router=llm_router,
)
# ## DELETE TEAM MEMBERSHIPS
for team_row in team_rows:
### get all team members

View file

@ -53,7 +53,7 @@ proxy = [
"orjson>=3.11.6,<4.0",
"apscheduler>=3.11.2,<4.0",
"fastapi-sso>=0.19.0,<1.0",
"PyJWT>=2.12.0,<3.0",
"PyJWT>=2.13.0,<3.0",
"python-multipart>=0.0.27,<1.0",
"cryptography>=46.0.7,<47.0",
"pynacl>=1.6.2,<2.0",

View file

@ -72,3 +72,18 @@ async def test_should_reject_invalid_limit(monkeypatch: pytest.MonkeyPatch):
await db.get_usage_data(limit="invalid")
assert query_mock.await_count == 0
@pytest.mark.asyncio
async def test_should_join_organization_table(monkeypatch: pytest.MonkeyPatch):
db, query_mock = _setup_db(monkeypatch, [])
await db.get_usage_data()
query_text, *_ = query_mock.await_args.args
assert (
"COALESCE(vt.organization_id, tt.organization_id) as organization_id"
in query_text
)
assert "ot.organization_alias as organization_alias" in query_text
assert 'LEFT JOIN "LiteLLM_OrganizationTable" ot' in query_text

View file

@ -0,0 +1,62 @@
"""Tests for FocusTransformer organization metadata in Tags."""
from __future__ import annotations
import json
from datetime import date
import polars as pl
from litellm.integrations.focus.transformer import FocusTransformer
def test_should_include_organization_fields_in_tags():
frame = pl.DataFrame(
{
"date": [date(2024, 1, 2)],
"spend": [1.25],
"api_requests": [1],
"api_key": ["hashed-key"],
"api_key_alias": ["prod-key"],
"model": ["gpt-4o"],
"model_group": ["gpt-4o"],
"custom_llm_provider": ["openai"],
"team_id": ["team-1"],
"team_alias": ["Platform"],
"organization_id": ["org-123"],
"organization_alias": ["Acme Corp"],
"user_id": ["user-1"],
"user_email": ["user@example.com"],
}
)
normalized = FocusTransformer().transform(frame)
tags = json.loads(normalized["Tags"][0])
assert tags["organization_id"] == "org-123"
assert tags["organization_alias"] == "Acme Corp"
assert tags["team_id"] == "team-1"
def test_should_omit_missing_organization_fields_from_tags():
frame = pl.DataFrame(
{
"date": [date(2024, 1, 2)],
"spend": [0.5],
"api_requests": [1],
"api_key": ["hashed-key"],
"api_key_alias": ["prod-key"],
"model": ["gpt-4o-mini"],
"model_group": ["gpt-4o-mini"],
"custom_llm_provider": ["openai"],
"team_id": ["team-1"],
"team_alias": ["Platform"],
}
)
normalized = FocusTransformer().transform(frame)
tags = json.loads(normalized["Tags"][0])
assert "organization_id" not in tags
assert "organization_alias" not in tags
assert tags["team_id"] == "team-1"

View file

@ -24,6 +24,7 @@ from litellm.proxy.management_endpoints.model_management_endpoints import (
ModelManagementAuthChecks,
_get_team_deployments,
clear_cache,
delete_team_models,
)
from litellm.proxy.utils import PrismaClient
from litellm.types.router import Deployment, LiteLLM_Params, updateDeployment
@ -2184,6 +2185,148 @@ class TestGetTeamDeployments:
assert result[0] is dep1
def _model_row(model_id: str, team_id: str):
row = MagicMock()
row.model_id = model_id
row.model_name = f"model_name_{team_id}_{model_id}"
row.model_info = {"team_id": team_id}
return row
class _TxProxyModelTable:
"""Transactional proxy-model table that records the order of DB writes."""
def __init__(self, rows, events):
self._rows = list(rows)
self.events = events
async def find_many(self, where):
prefix = where["model_name"]["startswith"]
return [r for r in self._rows if r.model_name.startswith(prefix)]
async def delete_many(self, where):
ids = list(where["model_id"]["in"])
self.events.append(("delete_many", tuple(ids)))
self._rows = [r for r in self._rows if r.model_id not in ids]
return len(ids)
class _TxPrismaClient:
"""Minimal prisma stub whose ``db.tx()`` yields a transaction and records commit."""
def __init__(self, rows):
self.events: list = []
self._table = _TxProxyModelTable(rows, self.events)
tx = MagicMock()
tx.litellm_proxymodeltable = self._table
outer = self
class _TxCM:
async def __aenter__(self):
return tx
async def __aexit__(self, *exc):
outer.events.append(("commit",))
return False
self.db = MagicMock()
self.db.tx = MagicMock(return_value=_TxCM())
class _RecordingRouter:
def __init__(self, events):
self.events = events
self.deleted: list = []
def delete_deployment(self, id): # noqa: A002 - matches router signature
self.events.append(("router", id))
self.deleted.append(id)
class TestDeleteTeamModels:
"""delete_team_models must remove every team's BYOK models in one transaction
and sync the in-memory router only after that transaction commits."""
@pytest.mark.asyncio
async def test_deletes_all_teams_models_and_syncs_router(self):
rows = [_model_row("a1", "team_a"), _model_row("b1", "team_b")]
prisma = _TxPrismaClient(rows)
router = _RecordingRouter(prisma.events)
deleted = await delete_team_models(
team_ids=["team_a", "team_b"],
prisma_client=prisma,
llm_router=router,
)
assert sorted(deleted) == ["a1", "b1"]
assert sorted(router.deleted) == ["a1", "b1"]
@pytest.mark.asyncio
async def test_router_sync_happens_after_commit(self):
"""Race-safety: the router is touched only once the DB transaction has
committed, so a rollback can never leave a deployment without its row."""
rows = [_model_row("a1", "team_a"), _model_row("b1", "team_b")]
prisma = _TxPrismaClient(rows)
router = _RecordingRouter(prisma.events)
await delete_team_models(
team_ids=["team_a", "team_b"], prisma_client=prisma, llm_router=router
)
commit_idx = prisma.events.index(("commit",))
router_indices = [i for i, e in enumerate(prisma.events) if e[0] == "router"]
delete_indices = [
i for i, e in enumerate(prisma.events) if e[0] == "delete_many"
]
assert router_indices, "router was never synced"
assert all(i > commit_idx for i in router_indices)
assert all(i < commit_idx for i in delete_indices)
@pytest.mark.asyncio
async def test_only_owning_team_models_deleted(self):
"""A row sharing the prefix but a different model_info.team_id is left alone."""
mine = _model_row("a1", "team_a")
intruder = MagicMock()
intruder.model_id = "x9"
intruder.model_name = "model_name_team_a_x9"
intruder.model_info = {"team_id": "someone_else"}
prisma = _TxPrismaClient([mine, intruder])
router = _RecordingRouter(prisma.events)
deleted = await delete_team_models(
team_ids=["team_a"], prisma_client=prisma, llm_router=router
)
assert deleted == ["a1"]
assert router.deleted == ["a1"]
@pytest.mark.asyncio
async def test_no_models_no_writes(self):
prisma = _TxPrismaClient([])
router = _RecordingRouter(prisma.events)
deleted = await delete_team_models(
team_ids=["team_a"], prisma_client=prisma, llm_router=router
)
assert deleted == []
assert router.deleted == []
assert not any(e[0] == "delete_many" for e in prisma.events)
@pytest.mark.asyncio
async def test_missing_router_is_safe(self):
rows = [_model_row("a1", "team_a")]
prisma = _TxPrismaClient(rows)
deleted = await delete_team_models(
team_ids=["team_a"], prisma_client=prisma, llm_router=None
)
assert deleted == ["a1"]
assert any(e[0] == "delete_many" for e in prisma.events)
def _build_db_model_for_blocked_test():
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo

View file

@ -6350,6 +6350,14 @@ async def test_delete_team_persists_deleted_teams(monkeypatch):
mock_find_many_keys = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many_keys
# delete_team now deletes team BYOK models inside a transaction; this team has none.
mock_tx = AsyncMock()
mock_tx.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
mock_tx_cm = MagicMock()
mock_tx_cm.__aenter__ = AsyncMock(return_value=mock_tx)
mock_tx_cm.__aexit__ = AsyncMock(return_value=False)
mock_prisma_client.db.tx = MagicMock(return_value=mock_tx_cm)
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma_client,
@ -8499,6 +8507,8 @@ async def test_new_team_encrypts_callback_vars(
assert cv["langfuse_secret_key"] != "sk-real"
recovered = decrypt_callback_vars(metadata)["logging"][0]["callback_vars"]
assert recovered["langfuse_secret_key"] == "sk-real"
def _non_admin_auth():
return UserAPIKeyAuth(
user_id="u-team-admin", user_role=LitellmUserRoles.INTERNAL_USER

View file

@ -13770,9 +13770,9 @@
}
},
"node_modules/ws": {
"version": "8.19.0",
"resolved": "https://registry.npmjs.org/ws/-/ws-8.19.0.tgz",
"integrity": "sha512-blAT2mjOEIi0ZzruJfIhb3nps74PRWTCz1IjglWEEpQl5XS/UNama6u2/rjFkDDouqr4L67ry+1aGIALViWjDg==",
"version": "8.20.1",
"resolved": "https://registry.npmjs.org/ws/-/ws-8.20.1.tgz",
"integrity": "sha512-It4dO0K5v//JtTXuPkfEOaI3uUN87iYPnqo/ZzqCoG3g8uhA66QUMs/SrM0YK7/NAu+r4LMh/9dq2A7k+rHs+w==",
"devOptional": true,
"license": "MIT",
"engines": {

View file

@ -90,7 +90,7 @@
"glob": "13.0.0",
"minimatch": "10.2.4",
"lodash": "4.18.1",
"ws": "8.19.0",
"ws": "8.20.1",
"braces": "3.0.3",
"axios": "1.13.6",
"postcss": "8.5.13"

View file

@ -115,8 +115,16 @@ describe("useProjects", () => {
expect(global.fetch).not.toHaveBeenCalled();
});
it("should not fetch when userRole is not an admin role", () => {
it("should fetch when userRole is an internal user role", async () => {
mockUseAuthorized.mockReturnValue({ accessToken: "test-token", userRole: "Internal User" });
(global.fetch as any).mockResolvedValue({ ok: true, json: async () => mockProjects });
const { result } = renderHook(() => useProjects(), { wrapper: makeWrapper(queryClient) });
await waitFor(() => expect(result.current.isSuccess).toBe(true));
expect(global.fetch).toHaveBeenCalled();
});
it("should not fetch when userRole cannot read projects", () => {
mockUseAuthorized.mockReturnValue({ accessToken: "test-token", userRole: "regular_user" });
const { result } = renderHook(() => useProjects(), { wrapper: makeWrapper(queryClient) });
expect(result.current.isFetched).toBe(false);
expect(global.fetch).not.toHaveBeenCalled();

View file

@ -2,7 +2,7 @@ import { useQuery } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
import { getProxyBaseUrl, getGlobalLitellmHeaderName, deriveErrorMessage, handleError } from "@/components/networking";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { all_admin_roles } from "@/utils/roles";
import { all_admin_roles, internalUserRoles } from "@/utils/roles";
// ── Types ────────────────────────────────────────────────────────────────────
@ -42,6 +42,8 @@ export interface ProjectResponse {
export const projectKeys = createQueryKeys("projects");
const projectReaderRoles = [...all_admin_roles, ...internalUserRoles];
// ── Fetch function ───────────────────────────────────────────────────────────
const fetchProjects = async (accessToken: string): Promise<ProjectResponse[]> => {
@ -74,6 +76,6 @@ export const useProjects = () => {
return useQuery<ProjectResponse[]>({
queryKey: projectKeys.list({}),
queryFn: async () => fetchProjects(accessToken!),
enabled: Boolean(accessToken) && all_admin_roles.includes(userRole!),
enabled: Boolean(accessToken) && projectReaderRoles.includes(userRole!),
});
};

13
uv.lock generated
View file

@ -9,7 +9,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-06-03T21:40:52.018333Z"
exclude-newer = "2026-06-05T23:18:37.734017Z"
exclude-newer-span = "P3D"
[manifest]
@ -3522,7 +3522,7 @@ requires-dist = [
{ name = "prometheus-client", marker = "extra == 'proxy-runtime'", specifier = ">=0.20.0,<1.0" },
{ name = "pydantic", specifier = ">=2.10.0,<3.0.0" },
{ name = "pydantic-settings", marker = "extra == 'proxy'", specifier = ">=2.14.1,<3.0" },
{ name = "pyjwt", marker = "extra == 'proxy'", specifier = ">=2.12.0,<3.0" },
{ name = "pyjwt", marker = "extra == 'proxy'", specifier = ">=2.13.0,<3.0" },
{ name = "pynacl", marker = "extra == 'proxy'", specifier = ">=1.6.2,<2.0" },
{ name = "pypdf", marker = "python_full_version < '3.14' and extra == 'proxy-runtime'", specifier = ">=6.10.2,<7.0" },
{ name = "pyroscope-io", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = ">=0.8.16,<1.0" },
@ -5983,11 +5983,14 @@ wheels = [
[[package]]
name = "pyjwt"
version = "2.12.0"
version = "2.13.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/a8/10/e8192be5f38f3e8e7e046716de4cae33d56fd5ae08927a823bb916be36c1/pyjwt-2.12.0.tar.gz", hash = "sha256:2f62390b667cd8257de560b850bb5a883102a388829274147f1d724453f8fb02", size = 102511, upload-time = "2026-03-12T17:15:30.831Z" }
dependencies = [
{ name = "typing-extensions", marker = "python_full_version < '3.11'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/3b/81/58d0ac84e1ef3a3843791d6954d94c0b33d526c75eeb1efbce9d0a4c4077/pyjwt-2.13.0.tar.gz", hash = "sha256:41571c89ca91598c79e8ef18a2d07367d4810fbbd6f637794879baf1b7703423", size = 107515, upload-time = "2026-05-21T19:54:36.618Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/15/70/70f895f404d363d291dcf62c12c85fdd47619ad9674ac0f53364d035925a/pyjwt-2.12.0-py3-none-any.whl", hash = "sha256:9bb459d1bdd0387967d287f5656bf7ec2b9a26645d1961628cda1764e087fd6e", size = 29700, upload-time = "2026-03-12T17:15:29.257Z" },
{ url = "https://files.pythonhosted.org/packages/a3/5e/ecf12fdb62546d64385c158514e9b2b671f7832108ef2ecd2020ce0af2d1/pyjwt-2.13.0-py3-none-any.whl", hash = "sha256:66adcc2aff09b3f1bbd95fc1e1577df8ac8723c978552fd43304c8a290ac5728", size = 31274, upload-time = "2026-05-21T19:54:35.362Z" },
]
[package.optional-dependencies]