mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
test(proxy): audit project team ownership across endpoints, bulk paths and chaos
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
8d0f36bc59
commit
bbbd0f62b3
1 changed files with 861 additions and 4 deletions
|
|
@ -1,19 +1,38 @@
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
import signal
|
||||
import threading
|
||||
from collections.abc import Iterator
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
import psutil
|
||||
import psycopg
|
||||
import pytest
|
||||
from integration._support.client import JSON_OBJECT, Gateway, gateway_from_environment, object_value, string_value
|
||||
from integration._support.client import (
|
||||
JSON_OBJECT,
|
||||
Gateway,
|
||||
delete_key_if_present,
|
||||
eventually,
|
||||
gateway_from_environment,
|
||||
object_value,
|
||||
string_value,
|
||||
)
|
||||
from integration._support.database import read_rows, write_rows
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.process import group_members, owned_proxy, owned_proxy_process
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
from litellm.models.user import LiteLLM_UserTable
|
||||
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
|
||||
|
||||
_OWNED_PROXY_SALT_KEY: Final = "sk-integration-salt"
|
||||
_OWNERSHIP_MARKER: Final = re.compile(rb"ownership-[0-9a-f]{32}")
|
||||
|
||||
|
||||
def _project_rows(project_id: str) -> list[dict[str, JsonValue]]:
|
||||
|
|
@ -40,6 +59,22 @@ def _key_permission_and_budget_ids(key: str) -> list[dict[str, JsonValue]]:
|
|||
)
|
||||
|
||||
|
||||
def _key_state_rows(key: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT to_jsonb(k) AS key_row, to_jsonb(b) AS budget_row FROM "LiteLLM_VerificationToken" AS k '
|
||||
'LEFT JOIN "LiteLLM_BudgetTable" AS b ON b.budget_id = k.budget_id WHERE k.token = %s',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
)
|
||||
|
||||
|
||||
def _project_state_rows(project_id: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT to_jsonb(p) AS project_row, to_jsonb(b) AS budget_row FROM "LiteLLM_ProjectTable" AS p '
|
||||
'LEFT JOIN "LiteLLM_BudgetTable" AS b ON b.budget_id = p.budget_id WHERE p.project_id = %s',
|
||||
(project_id,),
|
||||
)
|
||||
|
||||
|
||||
def _object_permission_rows(permission_id: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT * FROM "LiteLLM_ObjectPermissionTable" WHERE object_permission_id = %s',
|
||||
|
|
@ -47,6 +82,13 @@ def _object_permission_rows(permission_id: str) -> list[dict[str, JsonValue]]:
|
|||
)
|
||||
|
||||
|
||||
def _object_permission_table_rows() -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT to_jsonb(p) AS row FROM "LiteLLM_ObjectPermissionTable" AS p ORDER BY object_permission_id',
|
||||
(),
|
||||
)
|
||||
|
||||
|
||||
def _budget_rows(budget_id: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b WHERE budget_id = %s',
|
||||
|
|
@ -67,14 +109,14 @@ def _clear_key_object_permission(key: str, permission_id: str) -> None:
|
|||
|
||||
def _deleted_key_rows(key: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT token FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s',
|
||||
'SELECT to_jsonb(d) AS row FROM "LiteLLM_DeletedVerificationToken" AS d WHERE d.token = %s',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
)
|
||||
|
||||
|
||||
def _deprecated_key_rows(key: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT token FROM "LiteLLM_DeprecatedVerificationToken" WHERE token = %s',
|
||||
'SELECT to_jsonb(d) AS row FROM "LiteLLM_DeprecatedVerificationToken" AS d WHERE d.token = %s',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
)
|
||||
|
||||
|
|
@ -109,6 +151,821 @@ def _discard_unexpected_key(candidate: Gateway, response: httpx.Response) -> Non
|
|||
candidate.post("/key/delete", {"keys": [string_value(body["key"])]})
|
||||
|
||||
|
||||
def _ownership_chat_reply(identity: str, stream: bool) -> Reply:
|
||||
if not stream:
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": identity,
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-4.1-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "ownership ok"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=(
|
||||
f"data: {json.dumps({'id': identity, 'object': 'chat.completion.chunk', 'created': 1, 'model': 'gpt-4.1-mini', 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': 'ownership'}}]})}\n\n".encode(),
|
||||
f"data: {json.dumps({'id': identity, 'object': 'chat.completion.chunk', 'created': 1, 'model': 'gpt-4.1-mini', 'choices': [{'index': 0, 'delta': {'content': ' ok'}, 'finish_reason': 'stop'}], 'usage': {'prompt_tokens': 7, 'completion_tokens': 2, 'total_tokens': 9}})}\n\n".encode(),
|
||||
b"data: [DONE]\n\n",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _ownership_responses_reply(identity: str, stream: bool) -> Reply:
|
||||
response: Final = {
|
||||
"id": identity,
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-4.1-mini",
|
||||
"output": [
|
||||
{
|
||||
"id": f"msg_{identity}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": "ownership ok", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9},
|
||||
}
|
||||
if not stream:
|
||||
return Reply(body=json.dumps(response).encode())
|
||||
events: Final = (
|
||||
{
|
||||
"type": "response.created",
|
||||
"sequence_number": 0,
|
||||
"response": {**response, "status": "in_progress", "output": []},
|
||||
},
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"sequence_number": 1,
|
||||
"item_id": f"msg_{identity}",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": "ownership ok",
|
||||
},
|
||||
{"type": "response.completed", "sequence_number": 2, "response": response},
|
||||
)
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
|
||||
)
|
||||
|
||||
|
||||
def _ownership_upstream(request: Request) -> Reply:
|
||||
found: Final = _OWNERSHIP_MARKER.search(request.body)
|
||||
if found is None:
|
||||
return Reply(status=400, body=b'{"error":"missing ownership marker"}')
|
||||
marker: Final = found.group(0).decode()
|
||||
body: Final = JSON_OBJECT.validate_json(request.body)
|
||||
stream: Final = body.get("stream") is True
|
||||
if request.target.endswith("/responses"):
|
||||
return _ownership_responses_reply(f"resp_{marker}", stream)
|
||||
return _ownership_chat_reply(f"chatcmpl-{marker}", stream)
|
||||
|
||||
|
||||
def _ownership_sse_events(text: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
return tuple(
|
||||
JSON_OBJECT.validate_json(line[6:])
|
||||
for line in text.splitlines()
|
||||
if line.startswith("data: ") and line != "data: [DONE]"
|
||||
)
|
||||
|
||||
|
||||
def _ownership_request_marker(request: Request) -> str:
|
||||
found: Final = _OWNERSHIP_MARKER.search(request.body)
|
||||
assert found is not None, request.body
|
||||
return found.group(0).decode()
|
||||
|
||||
|
||||
def _ownership_request_payload(
|
||||
path: str, model: str, marker: str, stream: bool
|
||||
) -> tuple[dict[str, JsonValue], dict[str, str]]:
|
||||
if path.endswith("/messages"):
|
||||
return (
|
||||
{
|
||||
"model": model,
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": marker}],
|
||||
"stream": stream,
|
||||
},
|
||||
{"anthropic-version": "2023-06-01"},
|
||||
)
|
||||
if path.endswith("/responses"):
|
||||
return {"model": model, "input": marker, "stream": stream}, {}
|
||||
return {"model": model, "messages": [{"role": "user", "content": marker}], "stream": stream}, {}
|
||||
|
||||
|
||||
def _ownership_response_id(response: httpx.Response, path: str) -> str:
|
||||
if not response.headers.get("content-type", "").startswith("text/event-stream"):
|
||||
return string_value(JSON_OBJECT.validate_json(response.content)["id"])
|
||||
events: Final = _ownership_sse_events(response.text)
|
||||
identities: Final = (
|
||||
tuple(
|
||||
string_value(object_value(event["message"])["id"])
|
||||
for event in events
|
||||
if event.get("type") == "message_start"
|
||||
)
|
||||
if path.endswith("/messages")
|
||||
else tuple(
|
||||
string_value(object_value(event["response"])["id"])
|
||||
for event in events
|
||||
if event.get("type") == "response.completed"
|
||||
)
|
||||
if path.endswith("/responses")
|
||||
else tuple(string_value(event["id"]) for event in events if "id" in event)
|
||||
)
|
||||
unique_identities: Final = frozenset(identities)
|
||||
assert len(unique_identities) == 1, response.text
|
||||
return next(iter(unique_identities))
|
||||
|
||||
|
||||
def _ownership_serving_call(
|
||||
candidate: Gateway,
|
||||
key: str,
|
||||
model: str,
|
||||
index: int,
|
||||
marker: str,
|
||||
) -> tuple[int, str, str]:
|
||||
stream: Final = index % 2 == 0
|
||||
route: Final = index % 3
|
||||
paths: Final = ("/v1/chat/completions", "/v1/responses", "/v1/messages")
|
||||
path: Final = paths[route]
|
||||
body, headers = _ownership_request_payload(path, model, marker, stream)
|
||||
response: Final = candidate.request("POST", path, body, key=key, headers=headers)
|
||||
response.read()
|
||||
if response.status_code != 200:
|
||||
return response.status_code, "", response.text
|
||||
return response.status_code, _ownership_response_id(response, path), response.text
|
||||
|
||||
|
||||
def _create_mcp_server(candidate: Gateway, server_id: str, server_name: str, alias: str) -> None:
|
||||
response: Final = candidate.request(
|
||||
"POST",
|
||||
"/v1/mcp/server",
|
||||
{
|
||||
"server_id": server_id,
|
||||
"server_name": server_name,
|
||||
"alias": alias,
|
||||
"transport": "sse",
|
||||
"url": "http://127.0.0.1:9/mcp",
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201, response.text
|
||||
|
||||
|
||||
def _delete_mcp_server(candidate: Gateway, server_id: str) -> None:
|
||||
response: Final = candidate.request("DELETE", f"/v1/mcp/server/{server_id}")
|
||||
assert response.status_code == 202, response.text
|
||||
|
||||
|
||||
def test_service_account_generate_rejects_foreign_team_project_without_writing_key(ownership_gateway: Gateway) -> None:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
team_a: Final = scenario.team(models=[model])
|
||||
team_b: Final = scenario.team(models=[model])
|
||||
project_b: Final = scenario.project(team_b, models=[model])
|
||||
alias: Final = f"service-account-{uuid4().hex}"
|
||||
response: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/service-account/generate",
|
||||
{"team_id": team_a, "project_id": project_b, "key_alias": alias, "models": [model]},
|
||||
)
|
||||
_discard_unexpected_key(ownership_gateway, response)
|
||||
assert response.status_code == 400, response.text
|
||||
assert (
|
||||
read_rows(
|
||||
'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND key_alias = %s',
|
||||
(project_b, alias),
|
||||
)
|
||||
== []
|
||||
)
|
||||
|
||||
|
||||
def test_key_generate_unowned_project_accepts_any_team(ownership_gateway: Gateway) -> None:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
team_a: Final = scenario.team(models=[model])
|
||||
team_b: Final = scenario.team(models=[model])
|
||||
project: Final = scenario.project(team_a, models=[model])
|
||||
write_rows('UPDATE "LiteLLM_ProjectTable" SET team_id = NULL WHERE project_id = %s', (project,))
|
||||
team_key: Final = scenario.key(team_id=team_b, project_id=project, models=[model])
|
||||
teamless_key: Final = scenario.key(project_id=project, models=[model])
|
||||
chats: Final = tuple(
|
||||
ownership_gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": "legacy unowned"}]},
|
||||
key=key,
|
||||
)
|
||||
for key in (team_key, teamless_key)
|
||||
)
|
||||
assert all(chat.status_code == 200 for chat in chats), tuple(chat.text for chat in chats)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("path", "stream"),
|
||||
(
|
||||
("/v1/chat/completions", False),
|
||||
("/v1/chat/completions", True),
|
||||
("/v1/messages", False),
|
||||
("/v1/messages", True),
|
||||
("/v1/responses", False),
|
||||
("/v1/responses", True),
|
||||
),
|
||||
)
|
||||
def test_legacy_mismatched_key_keeps_serving_and_stays_editable(
|
||||
ownership_gateway: Gateway, path: str, stream: bool
|
||||
) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
return _ownership_upstream(request)
|
||||
|
||||
with wire_server(respond) as upstream:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(api_base=f"{upstream.url}/v1")
|
||||
team_a: Final = scenario.team(models=[model])
|
||||
team_b: Final = scenario.team(models=[model])
|
||||
project: Final = scenario.project(team_a, models=[model])
|
||||
key: Final = scenario.key(team_id=team_a, project_id=project, models=[model])
|
||||
write_rows(
|
||||
'UPDATE "LiteLLM_VerificationToken" SET team_id = %s WHERE token = %s',
|
||||
(team_b, sha256(key.encode()).hexdigest()),
|
||||
)
|
||||
marker: Final = f"ownership-{uuid4().hex}"
|
||||
body, headers = _ownership_request_payload(path, model, marker, stream)
|
||||
response: Final = ownership_gateway.request("POST", path, body, key=key, headers=headers)
|
||||
response.read()
|
||||
assert response.status_code == 200, response.text
|
||||
requests: Final = upstream.drain()
|
||||
assert len(requests) == 1
|
||||
expected_target: Final = "/chat/completions" if path.endswith("/chat/completions") else "/responses"
|
||||
assert requests[0].target.endswith(expected_target), requests[0].target
|
||||
assert marker.encode() in requests[0].body
|
||||
assert _ownership_response_id(response, path) != ""
|
||||
|
||||
|
||||
def test_stored_team_mismatch_allows_key_edits_and_detach(ownership_gateway: Gateway) -> None:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
team_a: Final = scenario.team(models=[model])
|
||||
team_b: Final = scenario.team(models=[model])
|
||||
project: Final = scenario.project(team_a, models=[model])
|
||||
key: Final = scenario.key(team_id=team_a, project_id=project, models=[model])
|
||||
write_rows(
|
||||
'UPDATE "LiteLLM_VerificationToken" SET team_id = %s WHERE token = %s',
|
||||
(team_b, sha256(key.encode()).hexdigest()),
|
||||
)
|
||||
alias: Final = f"stored-mismatch-{uuid4().hex}"
|
||||
alias_update: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/update",
|
||||
{"key": key, "key_alias": alias},
|
||||
)
|
||||
assert alias_update.status_code == 200, alias_update.text
|
||||
aliased_key: Final = _key_rows(key)
|
||||
assert aliased_key[0]["key_alias"] == alias
|
||||
unchanged: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/update",
|
||||
{"key": key, "team_id": team_b},
|
||||
)
|
||||
assert unchanged.status_code == 200, unchanged.text
|
||||
assert _key_rows(key) == aliased_key
|
||||
detached: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/update",
|
||||
{"key": key, "project_id": None},
|
||||
)
|
||||
assert detached.status_code == 200, detached.text
|
||||
detached_key: Final = _key_rows(key)
|
||||
assert detached_key[0]["team_id"] == team_b
|
||||
assert detached_key[0]["project_id"] is None
|
||||
|
||||
|
||||
def test_key_regenerate_cross_team_with_object_permission_writes_nothing(ownership_gateway: Gateway) -> None:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
team_a: Final = scenario.team(models=[model])
|
||||
team_b: Final = scenario.team(models=[model])
|
||||
project_a: Final = scenario.project(team_a, models=[model])
|
||||
project_b: Final = scenario.project(team_b, models=[model])
|
||||
key: Final = scenario.key(
|
||||
team_id=team_a,
|
||||
project_id=project_a,
|
||||
models=[model],
|
||||
object_permission={"vector_stores": ["before"]},
|
||||
)
|
||||
ids: Final = _key_permission_and_budget_ids(key)
|
||||
assert len(ids) == 1
|
||||
permission_id: Final = string_value(ids[0]["object_permission_id"])
|
||||
scenario.cleanups.callback(_clear_key_object_permission, key, permission_id)
|
||||
before: Final = (
|
||||
_object_permission_rows(permission_id),
|
||||
_object_permission_table_rows(),
|
||||
_key_rows(key),
|
||||
_key_state_rows(key),
|
||||
_key_permission_and_budget_ids(key),
|
||||
_deleted_key_rows(key),
|
||||
_deprecated_key_rows(key),
|
||||
)
|
||||
response: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
f"/key/{key}/regenerate",
|
||||
{"project_id": project_b, "object_permission": {"vector_stores": ["after"]}},
|
||||
)
|
||||
_discard_unexpected_key(ownership_gateway, response)
|
||||
assert response.status_code == 400, response.text
|
||||
assert (
|
||||
_object_permission_rows(permission_id),
|
||||
_object_permission_table_rows(),
|
||||
_key_rows(key),
|
||||
_key_state_rows(key),
|
||||
_key_permission_and_budget_ids(key),
|
||||
_deleted_key_rows(key),
|
||||
_deprecated_key_rows(key),
|
||||
) == before
|
||||
|
||||
|
||||
def test_key_bulk_update_cross_team_with_object_permission_preserves_permission(ownership_gateway: Gateway) -> None:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
team_a: Final = scenario.team(models=[model])
|
||||
team_b: Final = scenario.team(models=[model])
|
||||
project: Final = scenario.project(team_a, models=[model])
|
||||
key: Final = scenario.key(
|
||||
team_id=team_a,
|
||||
project_id=project,
|
||||
models=[model],
|
||||
object_permission={"vector_stores": ["before"]},
|
||||
)
|
||||
ids: Final = _key_permission_and_budget_ids(key)
|
||||
assert len(ids) == 1
|
||||
permission_id: Final = string_value(ids[0]["object_permission_id"])
|
||||
scenario.cleanups.callback(_clear_key_object_permission, key, permission_id)
|
||||
before: Final = (
|
||||
_object_permission_rows(permission_id),
|
||||
_object_permission_table_rows(),
|
||||
_key_rows(key),
|
||||
_key_state_rows(key),
|
||||
_key_permission_and_budget_ids(key),
|
||||
_deleted_key_rows(key),
|
||||
_deprecated_key_rows(key),
|
||||
)
|
||||
response: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/bulk_update",
|
||||
{
|
||||
"keys": [
|
||||
{
|
||||
"key": key,
|
||||
"team_id": team_b,
|
||||
"object_permission": {"vector_stores": ["after"]},
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
failed_updates: Final = JSON_OBJECT.validate_json(response.content)["failed_updates"]
|
||||
assert isinstance(failed_updates, list)
|
||||
assert len(failed_updates) == 1
|
||||
failed_update: Final = object_value(failed_updates[0])
|
||||
assert f"Project {project} belongs to team {team_a}" in string_value(failed_update["failed_reason"])
|
||||
assert (
|
||||
_object_permission_rows(permission_id),
|
||||
_object_permission_table_rows(),
|
||||
_key_rows(key),
|
||||
_key_state_rows(key),
|
||||
_key_permission_and_budget_ids(key),
|
||||
_deleted_key_rows(key),
|
||||
_deprecated_key_rows(key),
|
||||
) == before
|
||||
|
||||
|
||||
def test_key_update_ambiguous_mcp_permission_error_precedes_project_ownership(ownership_gateway: Gateway) -> None:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
identifier: Final = f"ambiguous{uuid4().hex}"
|
||||
first_id: Final = f"mcp{uuid4().hex}"
|
||||
second_id: Final = f"mcp{uuid4().hex}"
|
||||
_create_mcp_server(ownership_gateway, first_id, identifier, f"alias{uuid4().hex}")
|
||||
scenario.cleanups.callback(_delete_mcp_server, ownership_gateway, first_id)
|
||||
_create_mcp_server(ownership_gateway, second_id, f"name{uuid4().hex}", f"alias{uuid4().hex}")
|
||||
scenario.cleanups.callback(_delete_mcp_server, ownership_gateway, second_id)
|
||||
write_rows(
|
||||
'UPDATE "LiteLLM_MCPServerTable" SET alias = %s WHERE server_id = %s',
|
||||
(identifier, second_id),
|
||||
)
|
||||
team_permissions: Final = {"mcp_servers": [first_id, second_id]}
|
||||
team_a: Final = scenario.team(models=[model], object_permission=team_permissions)
|
||||
team_b: Final = scenario.team(models=[model], object_permission=team_permissions)
|
||||
project: Final = scenario.project(team_a, models=[model])
|
||||
key: Final = scenario.key(
|
||||
team_id=team_a,
|
||||
project_id=project,
|
||||
models=[model],
|
||||
object_permission={"vector_stores": ["before"]},
|
||||
)
|
||||
ids: Final = _key_permission_and_budget_ids(key)
|
||||
assert len(ids) == 1
|
||||
permission_id: Final = string_value(ids[0]["object_permission_id"])
|
||||
scenario.cleanups.callback(_clear_key_object_permission, key, permission_id)
|
||||
before: Final = (
|
||||
_object_permission_rows(permission_id),
|
||||
_object_permission_table_rows(),
|
||||
_key_rows(key),
|
||||
_key_state_rows(key),
|
||||
_key_permission_and_budget_ids(key),
|
||||
_deleted_key_rows(key),
|
||||
_deprecated_key_rows(key),
|
||||
)
|
||||
response: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/update",
|
||||
{
|
||||
"key": key,
|
||||
"team_id": team_b,
|
||||
"object_permission": {"mcp_tool_permissions": {identifier: ["tool"]}},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 400, response.text
|
||||
assert "ambiguous" in response.text.lower(), response.text
|
||||
assert "project" not in response.text.lower(), response.text
|
||||
assert (
|
||||
_object_permission_rows(permission_id),
|
||||
_object_permission_table_rows(),
|
||||
_key_rows(key),
|
||||
_key_state_rows(key),
|
||||
_key_permission_and_budget_ids(key),
|
||||
_deleted_key_rows(key),
|
||||
_deprecated_key_rows(key),
|
||||
) == before
|
||||
|
||||
|
||||
def test_team_key_bulk_update_rejects_foreign_team_project(ownership_gateway: Gateway) -> None:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
team_a: Final = scenario.team(models=[model])
|
||||
team_b: Final = scenario.team(models=[model])
|
||||
project_b: Final = scenario.project(team_b, models=[model])
|
||||
key: Final = scenario.key(team_id=team_a, models=[model])
|
||||
before: Final = _key_rows(key)
|
||||
response: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/team/key/bulk_update",
|
||||
{
|
||||
"team_id": team_a,
|
||||
"key_ids": [sha256(key.encode()).hexdigest()],
|
||||
"update_fields": {"project_id": project_b, "team_id": team_b},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 422, response.text
|
||||
assert "project_id" in response.text, response.text
|
||||
assert "team_id" in response.text, response.text
|
||||
assert _key_rows(key) == before
|
||||
|
||||
|
||||
def test_bulk_update_and_regenerate_new_object_permission_is_served(ownership_gateway: Gateway) -> None:
|
||||
with wire_server(_ownership_upstream) as upstream:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(api_base=f"{upstream.url}/v1")
|
||||
team: Final = scenario.team(models=[model])
|
||||
project: Final = scenario.project(team, models=[model])
|
||||
bulk_key: Final = scenario.key(team_id=team, project_id=project, models=[model])
|
||||
bulk_response: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/bulk_update",
|
||||
{
|
||||
"keys": [
|
||||
{
|
||||
"key": bulk_key,
|
||||
"object_permission": {"vector_stores": ["bulk"]},
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
assert bulk_response.status_code == 200, bulk_response.text
|
||||
bulk_info: Final = ownership_gateway.request("GET", "/key/info", params={"key": bulk_key})
|
||||
assert bulk_info.status_code == 200, bulk_info.text
|
||||
bulk_row: Final = object_value(JSON_OBJECT.validate_json(bulk_info.content)["info"])
|
||||
bulk_permission_id: Final = string_value(bulk_row["object_permission_id"])
|
||||
assert bulk_permission_id != ""
|
||||
assert len(_object_permission_rows(bulk_permission_id)) == 1
|
||||
scenario.cleanups.callback(_clear_key_object_permission, bulk_key, bulk_permission_id)
|
||||
bulk_chat: Final = ownership_gateway.chat(model, key=bulk_key, text=f"ownership-{uuid4().hex}")
|
||||
assert string_value(bulk_chat["id"]) != ""
|
||||
regenerated: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/generate",
|
||||
{"team_id": team, "project_id": project, "models": [model]},
|
||||
)
|
||||
assert regenerated.status_code == 200, regenerated.text
|
||||
regenerated_key: Final = string_value(JSON_OBJECT.validate_json(regenerated.content)["key"])
|
||||
scenario.cleanups.callback(delete_key_if_present, ownership_gateway, regenerated_key)
|
||||
regeneration: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
f"/key/{regenerated_key}/regenerate",
|
||||
{"object_permission": {"vector_stores": ["regenerated"]}},
|
||||
)
|
||||
assert regeneration.status_code == 200, regeneration.text
|
||||
new_key: Final = string_value(JSON_OBJECT.validate_json(regeneration.content)["key"])
|
||||
scenario.cleanups.callback(delete_key_if_present, ownership_gateway, new_key)
|
||||
info: Final = ownership_gateway.request("GET", "/key/info", params={"key": new_key})
|
||||
assert info.status_code == 200, info.text
|
||||
info_row: Final = object_value(JSON_OBJECT.validate_json(info.content)["info"])
|
||||
regenerated_permission_id: Final = string_value(info_row["object_permission_id"])
|
||||
assert regenerated_permission_id != ""
|
||||
assert regenerated_permission_id != bulk_permission_id
|
||||
assert len(_object_permission_rows(regenerated_permission_id)) == 1
|
||||
scenario.cleanups.callback(_clear_key_object_permission, new_key, regenerated_permission_id)
|
||||
chat: Final = ownership_gateway.chat(model, key=new_key, text=f"ownership-{uuid4().hex}")
|
||||
assert string_value(chat["id"]) != ""
|
||||
|
||||
|
||||
def test_ownership_rejections_during_concurrent_traffic_burst(ownership_gateway: Gateway) -> None:
|
||||
with wire_server(_ownership_upstream) as upstream:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(api_base=f"{upstream.url}/v1")
|
||||
team_a: Final = scenario.team(models=[model])
|
||||
team_b: Final = scenario.team(models=[model])
|
||||
project_a: Final = scenario.project(team_a, models=[model])
|
||||
serving_project: Final = scenario.project(team_a, models=[model])
|
||||
key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model])
|
||||
serving_key: Final = scenario.key(team_id=team_a, project_id=serving_project, models=[model])
|
||||
markers: Final = tuple(f"ownership-{uuid4().hex}" for _ in range(30))
|
||||
key_before: Final = (
|
||||
_object_permission_table_rows(),
|
||||
_key_rows(key),
|
||||
_key_state_rows(key),
|
||||
_key_permission_and_budget_ids(key),
|
||||
_deleted_key_rows(key),
|
||||
_deprecated_key_rows(key),
|
||||
)
|
||||
serving_key_before: Final = (
|
||||
_key_rows(serving_key),
|
||||
_key_permission_and_budget_ids(serving_key),
|
||||
_deleted_key_rows(serving_key),
|
||||
_deprecated_key_rows(serving_key),
|
||||
)
|
||||
project_before: Final = _project_state_rows(project_a)
|
||||
|
||||
def serve(index: int) -> tuple[int, str, str]:
|
||||
return _ownership_serving_call(ownership_gateway, serving_key, model, index, markers[index])
|
||||
|
||||
def update_key() -> httpx.Response:
|
||||
return ownership_gateway.request("POST", "/key/update", {"key": key, "team_id": team_b})
|
||||
|
||||
def bulk_update() -> httpx.Response:
|
||||
return ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/bulk_update",
|
||||
{"keys": [{"key": key, "team_id": team_b}]},
|
||||
)
|
||||
|
||||
def move_project() -> httpx.Response:
|
||||
return ownership_gateway.request(
|
||||
"POST", "/project/update", {"project_id": project_a, "team_id": team_b}
|
||||
)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=33) as pool:
|
||||
serving_futures: Final = tuple(pool.submit(serve, index) for index in range(30))
|
||||
ownership_futures: Final = (
|
||||
pool.submit(update_key),
|
||||
pool.submit(bulk_update),
|
||||
pool.submit(move_project),
|
||||
)
|
||||
serving_results: Final = tuple(future.result(timeout=90) for future in serving_futures)
|
||||
ownership_results: Final = tuple(future.result(timeout=90) for future in ownership_futures)
|
||||
assert all(status == 200 for status, _, _ in serving_results), serving_results
|
||||
assert all(response_id for _, response_id, _ in serving_results), serving_results
|
||||
assert len({response_id for _, response_id, _ in serving_results}) == 30
|
||||
key_update_response: Final = ownership_results[0]
|
||||
bulk_update_response: Final = ownership_results[1]
|
||||
project_update_response: Final = ownership_results[2]
|
||||
assert key_update_response.status_code == 400, key_update_response.text
|
||||
assert bulk_update_response.status_code == 200, bulk_update_response.text
|
||||
bulk_result: Final = JSON_OBJECT.validate_json(bulk_update_response.content)
|
||||
successful_updates: Final = bulk_result["successful_updates"]
|
||||
failed_updates: Final = bulk_result["failed_updates"]
|
||||
assert isinstance(successful_updates, list), bulk_update_response.text
|
||||
assert successful_updates == [], bulk_update_response.text
|
||||
assert isinstance(failed_updates, list), bulk_update_response.text
|
||||
assert len(failed_updates) == 1, bulk_update_response.text
|
||||
failed_update: Final = object_value(failed_updates[0])
|
||||
assert f"Project {project_a} belongs to team {team_a}" in string_value(failed_update["failed_reason"]), (
|
||||
bulk_update_response.text
|
||||
)
|
||||
assert project_update_response.status_code == 400, project_update_response.text
|
||||
upstream_requests: Final = upstream.drain()
|
||||
assert len(upstream_requests) == 30
|
||||
received_markers: Final = tuple(_ownership_request_marker(request) for request in upstream_requests)
|
||||
assert tuple(sorted(received_markers)) == tuple(sorted(markers))
|
||||
assert (
|
||||
_object_permission_table_rows(),
|
||||
_key_rows(key),
|
||||
_key_state_rows(key),
|
||||
_key_permission_and_budget_ids(key),
|
||||
_deleted_key_rows(key),
|
||||
_deprecated_key_rows(key),
|
||||
) == key_before
|
||||
assert (
|
||||
_key_rows(serving_key),
|
||||
_key_permission_and_budget_ids(serving_key),
|
||||
_deleted_key_rows(serving_key),
|
||||
_deprecated_key_rows(serving_key),
|
||||
) == serving_key_before
|
||||
assert _project_state_rows(project_a) == project_before
|
||||
|
||||
|
||||
def test_ownership_checks_wait_for_row_locks_without_failing_readiness(ownership_gateway: Gateway) -> None:
|
||||
with ownership_gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
team: Final = scenario.team(models=[model])
|
||||
project: Final = scenario.project(team, models=[model], max_budget=7)
|
||||
key: Final = scenario.key(team_id=team, project_id=project, models=[model])
|
||||
alias: Final = f"locked-{uuid4().hex}"
|
||||
locked: Final = threading.Event()
|
||||
release: Final = threading.Event()
|
||||
|
||||
def hold_locks() -> None:
|
||||
with psycopg.connect(os.environ["DATABASE_URL"]) as connection:
|
||||
connection.execute("BEGIN")
|
||||
connection.execute(
|
||||
'SELECT project_id FROM "LiteLLM_ProjectTable" WHERE project_id = %s FOR UPDATE',
|
||||
(project,),
|
||||
)
|
||||
connection.execute(
|
||||
'SELECT token FROM "LiteLLM_VerificationToken" WHERE token = %s FOR UPDATE',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
)
|
||||
locked.set()
|
||||
assert release.wait(timeout=30)
|
||||
connection.commit()
|
||||
|
||||
with ThreadPoolExecutor(max_workers=3) as pool:
|
||||
lock_future: Final = pool.submit(hold_locks)
|
||||
try:
|
||||
assert locked.wait(timeout=10)
|
||||
key_future: Final = pool.submit(
|
||||
ownership_gateway.request,
|
||||
"POST",
|
||||
"/key/update",
|
||||
{"key": key, "key_alias": alias},
|
||||
)
|
||||
project_future: Final = pool.submit(
|
||||
ownership_gateway.request,
|
||||
"POST",
|
||||
"/project/update",
|
||||
{"project_id": project, "team_id": team, "max_budget": 19},
|
||||
)
|
||||
lock_waiters: Final = eventually(
|
||||
lambda: read_rows(
|
||||
"SELECT pid FROM pg_stat_activity WHERE datname = current_database() "
|
||||
"AND wait_event_type = 'Lock' AND cardinality(pg_blocking_pids(pid)) > 0",
|
||||
(),
|
||||
),
|
||||
lambda rows: len(rows) >= 2,
|
||||
seconds=10,
|
||||
)
|
||||
assert len(lock_waiters) >= 2
|
||||
readiness: Final = eventually(
|
||||
lambda: ownership_gateway.request("GET", "/health/readiness").status_code,
|
||||
lambda status: status == 200,
|
||||
seconds=10,
|
||||
)
|
||||
assert readiness == 200
|
||||
finally:
|
||||
release.set()
|
||||
key_response: Final = key_future.result(timeout=90)
|
||||
project_response: Final = project_future.result(timeout=90)
|
||||
lock_future.result(timeout=30)
|
||||
assert key_response.status_code == 200, key_response.text
|
||||
assert project_response.status_code == 200, project_response.text
|
||||
assert _key_rows(key)[0]["key_alias"] == alias
|
||||
assert _project_rows(project)[0]["max_budget"] == 19.0
|
||||
|
||||
|
||||
def test_ownership_enforced_after_worker_kill(ownership_gateway: Gateway, tmp_path: Path) -> None:
|
||||
started: Final = threading.Event()
|
||||
release_stream: Final = threading.Event()
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
started.set()
|
||||
reply: Final = _ownership_upstream(request)
|
||||
return Reply(
|
||||
status=reply.status,
|
||||
body=reply.body,
|
||||
content_type=reply.content_type,
|
||||
chunks=reply.chunks,
|
||||
gate_after_first=release_stream,
|
||||
)
|
||||
|
||||
with wire_server(respond) as upstream:
|
||||
with owned_proxy_process(
|
||||
ownership_gateway,
|
||||
tmp_path / "worker-kill",
|
||||
{"LITELLM_SALT_KEY": _OWNED_PROXY_SALT_KEY},
|
||||
workers=2,
|
||||
) as owned:
|
||||
with owned.gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(api_base=f"{upstream.url}/v1")
|
||||
team_a: Final = scenario.team(models=[model])
|
||||
team_b: Final = scenario.team(models=[model])
|
||||
project_a: Final = scenario.project(team_a, models=[model])
|
||||
project_b: Final = scenario.project(team_b, models=[model])
|
||||
key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model])
|
||||
markers: Final = tuple(f"ownership-{uuid4().hex}" for _ in range(24))
|
||||
|
||||
def serve(index: int) -> tuple[int, int, str]:
|
||||
try:
|
||||
response: Final = owned.gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": markers[index]}],
|
||||
"stream": True,
|
||||
},
|
||||
key=key,
|
||||
)
|
||||
except httpx.TransportError as error:
|
||||
return index, 0, str(error)
|
||||
return index, response.status_code, response.text
|
||||
|
||||
candidate_port: Final = owned.gateway.client.base_url.port
|
||||
assert candidate_port is not None
|
||||
workers: Final = tuple(
|
||||
process
|
||||
for process in group_members(owned.process.pid)
|
||||
if process.pid != owned.process.pid
|
||||
and any(
|
||||
connection.laddr.port == candidate_port and connection.status == psutil.CONN_LISTEN
|
||||
for connection in process.net_connections(kind="inet")
|
||||
)
|
||||
)
|
||||
assert len(workers) == 2, tuple(process.pid for process in workers)
|
||||
victim: Final = workers[0]
|
||||
survivor: Final = workers[1]
|
||||
with ThreadPoolExecutor(max_workers=24) as pool:
|
||||
futures: Final = tuple(pool.submit(serve, index) for index in range(24))
|
||||
try:
|
||||
assert started.wait(timeout=10)
|
||||
os.kill(victim.pid, signal.SIGKILL)
|
||||
release_stream.set()
|
||||
psutil.wait_procs((victim,), timeout=10)
|
||||
assert not psutil.pid_exists(victim.pid), victim.pid
|
||||
finally:
|
||||
release_stream.set()
|
||||
burst: Final = tuple(future.result(timeout=90) for future in futures)
|
||||
assert psutil.pid_exists(survivor.pid), survivor.pid
|
||||
assert len(burst) == 24
|
||||
burst_requests: Final = upstream.drain()
|
||||
empty_body_count: Final = sum(not request.body for request in burst_requests)
|
||||
burst_markers: Final = tuple(
|
||||
_ownership_request_marker(request) for request in burst_requests if request.body
|
||||
)
|
||||
successful_markers: Final = tuple(markers[index] for index, status, _ in burst if status == 200)
|
||||
assert burst_markers, f"empty_body_captures={empty_body_count}; burst={burst}"
|
||||
assert all(burst_markers.count(marker) == 1 for marker in successful_markers), (
|
||||
f"empty_body_captures={empty_body_count}; "
|
||||
f"successful_markers={successful_markers}; upstream_markers={burst_markers}; burst={burst}"
|
||||
)
|
||||
cross_team: Final = owned.gateway.request(
|
||||
"POST",
|
||||
"/key/generate",
|
||||
{"team_id": team_a, "project_id": project_b, "models": [model]},
|
||||
)
|
||||
_discard_unexpected_key(owned.gateway, cross_team)
|
||||
assert cross_team.status_code == 400, cross_team.text
|
||||
marker: Final = f"ownership-{uuid4().hex}"
|
||||
chat: Final = owned.gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": marker}]},
|
||||
key=key,
|
||||
)
|
||||
assert chat.status_code == 200, chat.text
|
||||
post_kill_requests: Final = upstream.drain()
|
||||
post_kill_empty_body_count: Final = sum(not request.body for request in post_kill_requests)
|
||||
post_kill_markers: Final = tuple(
|
||||
_ownership_request_marker(request) for request in post_kill_requests if request.body
|
||||
)
|
||||
assert post_kill_markers.count(marker) == 1, (
|
||||
f"empty_body_captures={post_kill_empty_body_count}; "
|
||||
f"post_kill_markers={post_kill_markers}; response={chat.text}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def ownership_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
with gateway_from_environment() as gateway:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue