mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): enforce per-caller ownership on proxy-created videos
Record the creating caller's owner scope in LiteLLM_ManagedObjectTable (file_purpose=video, keyed by the provider-native video id) when videos are created, remixed, edited, or extended, and reject status, content, remix, edit, and extension calls on a video the caller does not own with 403. Proxy admins bypass the check, untracked videos are admin-only, and list results are filtered to the caller's videos. Without a database the checks are skipped, matching previous behavior. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
ab4de2f02a
commit
c510bd1aad
3 changed files with 403 additions and 1 deletions
|
|
@ -16,6 +16,11 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
|
|||
get_custom_llm_provider_from_request_query,
|
||||
)
|
||||
from litellm.proxy.image_endpoints.endpoints import batch_to_bytesio
|
||||
from litellm.proxy.video_endpoints.ownership import (
|
||||
assert_user_can_access_video,
|
||||
filter_video_list_for_caller,
|
||||
record_video_owner,
|
||||
)
|
||||
from litellm.proxy.video_endpoints.utils import (
|
||||
encode_character_id_in_response,
|
||||
extract_model_from_target_model_names,
|
||||
|
|
@ -116,6 +121,7 @@ async def video_generation(
|
|||
version=version,
|
||||
)
|
||||
else:
|
||||
await record_video_owner(generated, user_api_key_dict)
|
||||
return generated
|
||||
|
||||
|
||||
|
|
@ -203,7 +209,7 @@ async def video_list(
|
|||
version=version,
|
||||
)
|
||||
else:
|
||||
return listed
|
||||
return await filter_video_list_for_caller(listed, user_api_key_dict)
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -250,6 +256,8 @@ async def video_status(
|
|||
version,
|
||||
)
|
||||
|
||||
await assert_user_can_access_video(video_id, user_api_key_dict)
|
||||
|
||||
# Create data with video_id
|
||||
data: Final[dict[str, object]] = {"video_id": video_id}
|
||||
|
||||
|
|
@ -353,6 +361,8 @@ async def video_content(
|
|||
version,
|
||||
)
|
||||
|
||||
await assert_user_can_access_video(video_id, user_api_key_dict)
|
||||
|
||||
# Create data with video_id
|
||||
data: Final[dict[str, object]] = {"video_id": video_id}
|
||||
|
||||
|
|
@ -462,6 +472,8 @@ async def video_remix(
|
|||
version,
|
||||
)
|
||||
|
||||
await assert_user_can_access_video(video_id, user_api_key_dict)
|
||||
|
||||
data: Final = await _read_request_body(request=request)
|
||||
data["video_id"] = video_id
|
||||
|
||||
|
|
@ -514,6 +526,7 @@ async def video_remix(
|
|||
version=version,
|
||||
)
|
||||
else:
|
||||
await record_video_owner(remixed, user_api_key_dict)
|
||||
return remixed
|
||||
|
||||
|
||||
|
|
@ -780,6 +793,7 @@ async def video_edit(
|
|||
data["video_id"] = ""
|
||||
else:
|
||||
data["video_id"] = video_reference_to_id(uploaded_video)
|
||||
await assert_user_can_access_video(data["video_id"], user_api_key_dict)
|
||||
|
||||
decoded: Final = decode_video_id_with_provider(data["video_id"])
|
||||
provider_from_id: Final = decoded.get("custom_llm_provider")
|
||||
|
|
@ -827,6 +841,7 @@ async def video_edit(
|
|||
version=version,
|
||||
)
|
||||
else:
|
||||
await record_video_owner(edited, user_api_key_dict)
|
||||
return edited
|
||||
|
||||
|
||||
|
|
@ -877,6 +892,7 @@ async def video_extension(
|
|||
|
||||
data: Final = await _read_request_body(request=request)
|
||||
data["video_id"] = video_reference_to_id(data.pop("video", None))
|
||||
await assert_user_can_access_video(data["video_id"], user_api_key_dict)
|
||||
|
||||
decoded: Final = decode_video_id_with_provider(data["video_id"])
|
||||
provider_from_id: Final = decoded.get("custom_llm_provider")
|
||||
|
|
@ -924,4 +940,5 @@ async def video_extension(
|
|||
version=version,
|
||||
)
|
||||
else:
|
||||
await record_video_owner(extended, user_api_key_dict)
|
||||
return extended
|
||||
|
|
|
|||
167
litellm/proxy/video_endpoints/ownership.py
Normal file
167
litellm/proxy/video_endpoints/ownership.py
Normal file
|
|
@ -0,0 +1,167 @@
|
|||
"""Per-caller ownership of videos created through the proxy.
|
||||
|
||||
Ownership rows live in ``LiteLLM_ManagedObjectTable`` with ``file_purpose="video"``
|
||||
(the same table and owner scopes container ownership uses). A row is keyed by the
|
||||
provider-native video id only, so re-wrapping an id with a different provider or
|
||||
model_id cannot dodge the lookup. Proxy admins bypass the check; for everyone else,
|
||||
a video without a row is treated as admin-only.
|
||||
|
||||
Without a connected database there is nowhere to record ownership: recording and
|
||||
checks are skipped and videos remain reachable by any caller allowed on the route.
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from fastapi import HTTPException
|
||||
from typing_extensions import TypeIs # noqa: TID251 # narrows untyped wire payloads without a runtime conversion
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.resource_ownership import (
|
||||
get_primary_resource_owner_scope,
|
||||
get_resource_owner_scopes,
|
||||
is_proxy_admin,
|
||||
user_can_access_resource_owner,
|
||||
)
|
||||
from litellm.repositories.table_repositories import ManagedObjectRepository
|
||||
from litellm.types.videos.utils import extract_original_video_id
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
VIDEO_OBJECT_PURPOSE: Final = "video"
|
||||
|
||||
# Status is polled; a short-lived cache of known owners keeps polling off the DB.
|
||||
# Only positive answers are cached, so a missing row is always re-checked.
|
||||
_VIDEO_OWNER_CACHE: Final = InMemoryCache(max_size_in_memory=10000, default_ttl=60)
|
||||
|
||||
|
||||
def _video_model_object_id(video_id: str) -> str:
|
||||
return f"{VIDEO_OBJECT_PURPOSE}:{extract_original_video_id(video_id)}"
|
||||
|
||||
|
||||
def _video_id_of(item: object) -> str | None:
|
||||
match item:
|
||||
case {"id": str(video_id)} if video_id:
|
||||
return video_id
|
||||
case object(id=str(video_id)) if video_id:
|
||||
return video_id
|
||||
case _:
|
||||
return None
|
||||
|
||||
|
||||
def _is_object_mapping(value: object) -> TypeIs[Mapping[str, object]]: # guard-ok: provider JSON objects have str keys
|
||||
return isinstance(value, Mapping)
|
||||
|
||||
|
||||
def _is_object_sequence(value: object) -> TypeIs[Sequence[object]]: # guard-ok: a list is a Sequence of anything
|
||||
return isinstance(value, list)
|
||||
|
||||
|
||||
def _get_prisma_client() -> "PrismaClient | None":
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
return prisma_client
|
||||
|
||||
|
||||
async def record_video_owner(response: object, user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
"""Stamp the caller as owner of a video the provider just created.
|
||||
|
||||
Failures are logged rather than raised: the provider job already exists and is
|
||||
billed, so the caller still gets its id; the untracked video is admin-only.
|
||||
"""
|
||||
prisma_client: Final = _get_prisma_client()
|
||||
if prisma_client is None:
|
||||
return
|
||||
video_id: Final = _video_id_of(response)
|
||||
owner: Final = get_primary_resource_owner_scope(user_api_key_dict)
|
||||
if video_id is None or owner is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Skipping video ownership tracking: response has no id or caller has no identity scope"
|
||||
)
|
||||
return
|
||||
model_object_id: Final = _video_model_object_id(video_id)
|
||||
row: Final = { # mutable-ok: prisma rows are plain dicts
|
||||
"unified_object_id": video_id,
|
||||
"model_object_id": model_object_id,
|
||||
"file_object": json.dumps({"id": video_id, "object": "video"}), # mutable-ok: serialized immediately
|
||||
"file_purpose": VIDEO_OBJECT_PURPOSE,
|
||||
"created_by": owner,
|
||||
"updated_by": owner,
|
||||
}
|
||||
try:
|
||||
await ManagedObjectRepository(prisma_client).table.upsert(
|
||||
where={"model_object_id": model_object_id}, # mutable-ok: prisma filters are plain dicts
|
||||
data={"create": row, "update": {"updated_by": owner}}, # mutable-ok: prisma payloads are plain dicts
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"Video ownership recording failed; video_id=%s is untracked and admin-only: %s", video_id, e
|
||||
)
|
||||
|
||||
|
||||
async def _get_video_owner(prisma_client: "PrismaClient", model_object_id: str) -> str | None:
|
||||
cached: Final = _VIDEO_OWNER_CACHE.get_cache(model_object_id)
|
||||
if isinstance(cached, str):
|
||||
return cached
|
||||
row: Final = await ManagedObjectRepository(prisma_client).table.find_first(
|
||||
where={ # mutable-ok: prisma filters are plain dicts
|
||||
"model_object_id": model_object_id,
|
||||
"file_purpose": VIDEO_OBJECT_PURPOSE,
|
||||
}
|
||||
)
|
||||
owner: Final = getattr(row, "created_by", None) if row is not None else None
|
||||
if not isinstance(owner, str):
|
||||
return None
|
||||
_VIDEO_OWNER_CACHE.set_cache(model_object_id, owner)
|
||||
return owner
|
||||
|
||||
|
||||
async def assert_user_can_access_video(video_id: str, user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
"""Raise 403 unless the caller owns ``video_id`` (or is a proxy admin)."""
|
||||
if not video_id or is_proxy_admin(user_api_key_dict):
|
||||
return
|
||||
prisma_client: Final = _get_prisma_client()
|
||||
if prisma_client is None:
|
||||
return
|
||||
owner: Final = await _get_video_owner(prisma_client, _video_model_object_id(video_id))
|
||||
if not user_can_access_resource_owner(owner, user_api_key_dict):
|
||||
raise HTTPException(status_code=403, detail="Forbidden")
|
||||
|
||||
|
||||
async def filter_video_list_for_caller(listed: object, user_api_key_dict: UserAPIKeyAuth) -> object:
|
||||
"""Drop videos the caller does not own from a provider list page.
|
||||
|
||||
Pagination cursors are left as the provider returned them so ``after`` keeps
|
||||
walking the provider's list even when a page has no videos owned by the caller.
|
||||
"""
|
||||
if is_proxy_admin(user_api_key_dict) or not _is_object_mapping(listed):
|
||||
return listed
|
||||
prisma_client: Final = _get_prisma_client()
|
||||
items: Final = listed.get("data")
|
||||
if prisma_client is None or not _is_object_sequence(items):
|
||||
return listed
|
||||
ids_by_item: Final = tuple((item, _video_id_of(item)) for item in items)
|
||||
candidate_ids: Final = tuple(
|
||||
_video_model_object_id(video_id) for _, video_id in ids_by_item if video_id is not None
|
||||
)
|
||||
owner_scopes: Final = get_resource_owner_scopes(user_api_key_dict)
|
||||
rows: Final = (
|
||||
await ManagedObjectRepository(prisma_client).table.find_many(
|
||||
where={ # mutable-ok: prisma filters are plain dicts
|
||||
"model_object_id": {"in": list(candidate_ids)}, # mutable-ok: prisma filters are plain dicts
|
||||
"file_purpose": VIDEO_OBJECT_PURPOSE,
|
||||
"created_by": {"in": owner_scopes}, # mutable-ok: prisma filters are plain dicts
|
||||
}
|
||||
)
|
||||
if candidate_ids and owner_scopes
|
||||
else ()
|
||||
)
|
||||
owned: Final = frozenset(row.model_object_id for row in rows)
|
||||
kept: Final = tuple(
|
||||
item for item, video_id in ids_by_item if video_id is not None and _video_model_object_id(video_id) in owned
|
||||
)
|
||||
return {**listed, "data": kept} # mutable-ok: provider list JSON body
|
||||
218
tests/test_litellm/proxy/video_endpoints/test_ownership.py
Normal file
218
tests/test_litellm/proxy/video_endpoints/test_ownership.py
Normal file
|
|
@ -0,0 +1,218 @@
|
|||
"""
|
||||
Ownership tests for the proxy video endpoints.
|
||||
|
||||
Requests go through the real FastAPI routes and the real ownership module. Only two
|
||||
I/O boundaries are replaced: the provider call (``base_process_llm_request``) and the
|
||||
Prisma ``litellm_managedobjecttable`` delegate, which is an in-memory table honoring
|
||||
the subset of the query API the ownership module uses.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth, hash_token
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.video_endpoints import endpoints
|
||||
from litellm.types.videos.utils import encode_video_id_with_provider
|
||||
|
||||
ALICE = UserAPIKeyAuth(user_id="alice", api_key=hash_token("sk-alice"))
|
||||
BOB = UserAPIKeyAuth(user_id="bob", api_key=hash_token("sk-bob"))
|
||||
KEY_ONLY_A = UserAPIKeyAuth(api_key=hash_token("sk-service-a"))
|
||||
KEY_ONLY_B = UserAPIKeyAuth(api_key=hash_token("sk-service-b"))
|
||||
ADMIN = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
|
||||
def _matches(row: SimpleNamespace, where: dict) -> bool:
|
||||
for field, condition in where.items():
|
||||
value = getattr(row, field, None)
|
||||
if isinstance(condition, dict):
|
||||
if value not in condition["in"]:
|
||||
return False
|
||||
elif value != condition:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
class InMemoryManagedObjectTable:
|
||||
def __init__(self, fail_writes: bool = False):
|
||||
self.rows: dict[str, SimpleNamespace] = {}
|
||||
self.fail_writes = fail_writes
|
||||
|
||||
async def upsert(self, where: dict, data: dict) -> SimpleNamespace:
|
||||
if self.fail_writes:
|
||||
raise RuntimeError("database unavailable")
|
||||
key = where["model_object_id"]
|
||||
if key in self.rows:
|
||||
self.rows[key] = SimpleNamespace(**{**vars(self.rows[key]), **data["update"]})
|
||||
else:
|
||||
self.rows[key] = SimpleNamespace(**data["create"])
|
||||
return self.rows[key]
|
||||
|
||||
async def find_first(self, where: dict) -> SimpleNamespace | None:
|
||||
return next((row for row in self.rows.values() if _matches(row, where)), None)
|
||||
|
||||
async def find_many(self, where: dict) -> list[SimpleNamespace]:
|
||||
return [row for row in self.rows.values() if _matches(row, where)]
|
||||
|
||||
|
||||
class Provider:
|
||||
"""The provider behind ``base_process_llm_request``: creates and serves videos."""
|
||||
|
||||
def __init__(self):
|
||||
self.calls: list[str] = []
|
||||
|
||||
async def respond(self, processor, *, route_type: str, **kwargs):
|
||||
self.calls.append(route_type)
|
||||
if route_type in ("avideo_generation", "avideo_remix", "avideo_edit", "avideo_extension"):
|
||||
return {"id": f"video_{uuid.uuid4().hex}", "object": "video", "status": "queued"}
|
||||
if route_type == "avideo_content":
|
||||
return b"mp4-bytes"
|
||||
if route_type == "avideo_list":
|
||||
return {"object": "list", "data": processor.data["listed"], "has_more": False}
|
||||
return {"id": processor.data["video_id"], "object": "video", "status": "completed"}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def provider(monkeypatch) -> Provider:
|
||||
fake = Provider()
|
||||
|
||||
async def base_process_llm_request(processor, **kwargs):
|
||||
return await fake.respond(processor, **kwargs)
|
||||
|
||||
monkeypatch.setattr(ProxyBaseLLMRequestProcessing, "base_process_llm_request", base_process_llm_request)
|
||||
return fake
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def table(monkeypatch) -> InMemoryManagedObjectTable:
|
||||
rows = InMemoryManagedObjectTable()
|
||||
monkeypatch.setattr(
|
||||
proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(litellm_managedobjecttable=rows))
|
||||
)
|
||||
return rows
|
||||
|
||||
|
||||
def _client(auth: UserAPIKeyAuth) -> TestClient:
|
||||
app = FastAPI()
|
||||
app.include_router(endpoints.router)
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: auth
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def _create_video(auth: UserAPIKeyAuth) -> str:
|
||||
response = _client(auth).post("/v1/videos", json={"model": "sora-2", "prompt": "a sunset"})
|
||||
assert response.status_code == 200
|
||||
return response.json()["id"]
|
||||
|
||||
|
||||
def _access(auth: UserAPIKeyAuth, video_id: str) -> dict[str, int]:
|
||||
client = _client(auth)
|
||||
return {
|
||||
"status": client.get(f"/v1/videos/{video_id}").status_code,
|
||||
"content": client.get(f"/v1/videos/{video_id}/content").status_code,
|
||||
"remix": client.post(f"/v1/videos/{video_id}/remix", json={"prompt": "again"}).status_code,
|
||||
"edit": client.post("/v1/videos/edits", json={"prompt": "brighter", "video": {"id": video_id}}).status_code,
|
||||
"extension": client.post(
|
||||
"/v1/videos/extensions", json={"prompt": "continue", "video": {"id": video_id}}
|
||||
).status_code,
|
||||
}
|
||||
|
||||
|
||||
ALL_OK = {"status": 200, "content": 200, "remix": 200, "edit": 200, "extension": 200}
|
||||
ALL_FORBIDDEN = {"status": 403, "content": 403, "remix": 403, "edit": 403, "extension": 403}
|
||||
|
||||
|
||||
def test_owner_can_use_every_video_route(provider, table):
|
||||
video_id = _create_video(ALICE)
|
||||
|
||||
assert table.rows[f"video:{video_id}"].created_by == "alice"
|
||||
assert _access(ALICE, video_id) == ALL_OK
|
||||
|
||||
|
||||
def test_other_user_is_forbidden_and_the_provider_is_never_called(provider, table):
|
||||
video_id = _create_video(ALICE)
|
||||
provider.calls.clear()
|
||||
|
||||
assert _access(BOB, video_id) == ALL_FORBIDDEN
|
||||
assert provider.calls == []
|
||||
|
||||
|
||||
def test_key_scoped_owner_blocks_a_different_key(provider, table):
|
||||
video_id = _create_video(KEY_ONLY_A)
|
||||
|
||||
assert _access(KEY_ONLY_B, video_id) == ALL_FORBIDDEN
|
||||
assert _access(KEY_ONLY_A, video_id) == ALL_OK
|
||||
|
||||
|
||||
def test_proxy_admin_can_access_any_video(provider, table):
|
||||
video_id = _create_video(ALICE)
|
||||
|
||||
assert _access(ADMIN, video_id) == ALL_OK
|
||||
|
||||
|
||||
def test_untracked_video_is_admin_only_when_a_database_is_connected(provider, table):
|
||||
video_id = f"video_{uuid.uuid4().hex}"
|
||||
|
||||
assert _access(ALICE, video_id) == ALL_FORBIDDEN
|
||||
assert _access(ADMIN, video_id) == ALL_OK
|
||||
|
||||
|
||||
def test_rewrapping_the_provider_id_does_not_bypass_the_owner_check(provider, table):
|
||||
video_id = _create_video(ALICE)
|
||||
rewrapped = encode_video_id_with_provider(video_id, "azure", "attacker-deployment")
|
||||
|
||||
assert rewrapped != video_id
|
||||
assert _access(BOB, rewrapped) == ALL_FORBIDDEN
|
||||
|
||||
|
||||
def test_videos_derived_from_an_owned_video_belong_to_the_caller(provider, table):
|
||||
video_id = _create_video(ALICE)
|
||||
remixed = _client(ALICE).post(f"/v1/videos/{video_id}/remix", json={"prompt": "again"}).json()["id"]
|
||||
|
||||
assert _access(ALICE, remixed)["status"] == 200
|
||||
assert _access(BOB, remixed)["status"] == 403
|
||||
|
||||
|
||||
def test_list_only_returns_the_callers_videos(provider, table):
|
||||
alice_video = _create_video(ALICE)
|
||||
bob_video = _create_video(BOB)
|
||||
listed = [{"id": alice_video, "object": "video"}, {"id": bob_video, "object": "video"}]
|
||||
provider_listing = {"listed": listed}
|
||||
|
||||
def list_as(auth: UserAPIKeyAuth) -> list[str]:
|
||||
async def with_listing(processor, **kwargs):
|
||||
processor.data.update(provider_listing)
|
||||
return await provider.respond(processor, **kwargs)
|
||||
|
||||
with pytest.MonkeyPatch.context() as mp:
|
||||
mp.setattr(ProxyBaseLLMRequestProcessing, "base_process_llm_request", with_listing)
|
||||
response = _client(auth).get("/v1/videos")
|
||||
assert response.status_code == 200
|
||||
return [item["id"] for item in response.json()["data"]]
|
||||
|
||||
assert list_as(ALICE) == [alice_video]
|
||||
assert list_as(BOB) == [bob_video]
|
||||
assert list_as(ADMIN) == [alice_video, bob_video]
|
||||
|
||||
|
||||
def test_without_a_database_ownership_is_not_enforced(provider, monkeypatch):
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
video_id = _create_video(ALICE)
|
||||
|
||||
assert _access(BOB, video_id) == ALL_OK
|
||||
|
||||
|
||||
def test_failed_ownership_write_still_returns_the_created_video(provider, table):
|
||||
table.fail_writes = True
|
||||
|
||||
video_id = _create_video(ALICE)
|
||||
|
||||
assert video_id.startswith("video_")
|
||||
assert table.rows == {}
|
||||
assert _access(ALICE, video_id)["status"] == 403
|
||||
Loading…
Add table
Reference in a new issue