feat(proxy): enforce team isolation for provider-format batch ids and output files

Managed batch ownership rows were already written for every create path,
but retrieve/cancel/file-content checks only fired for unified
(litellm_proxy-prefixed) ids, so model-encoded and raw provider batch ids
bypassed isolation entirely.

- enforce can_access_resource on retrieve/cancel for provider-format batch
  ids when an ownership row exists; ids with no row stay accessible so
  pass-through reads keep working
- make batch sync operations update-only so a retrieve/cancel can never
  mint an ownership row attributed to the first caller
- write ownership rows for a synced batch's provider-format output/error
  file ids, inherited from the owning batch row, and enforce them on
  file content/retrieve/delete
This commit is contained in:
Yuneng Jiang 2026-07-16 08:36:46 -07:00
parent 69a491e168
commit f07ea5921b
No known key found for this signature in database
2 changed files with 554 additions and 0 deletions

View file

@ -30,10 +30,12 @@ from litellm.proxy._types import (
)
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
decode_model_from_file_id,
get_batch_id_from_unified_batch_id,
get_content_type_from_file_object,
get_model_id_from_unified_batch_id,
get_models_from_unified_file_id,
get_original_file_id,
normalize_mime_type_for_provider,
)
from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccessIssue]
@ -165,7 +167,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
model_object_id: str,
file_purpose: Literal["batch", "fine-tune", "response"],
user_api_key_dict: UserAPIKeyAuth,
update_only: bool = False,
) -> None:
"""Persist a managed object row.
With ``update_only=True`` an existing row is refreshed but a missing
row is NOT created: sync operations (retrieve/cancel) must never mint
an ownership row attributed to whoever happened to call them first.
"""
verbose_logger.info(
f"Storing LiteLLM Managed {file_purpose} object with id={unified_object_id} in cache"
)
@ -175,6 +184,27 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
file_purpose=file_purpose,
file_object=file_object,
)
if update_only:
updated_count = (
await self.prisma_client.db.litellm_managedobjecttable.update_many(
where={"unified_object_id": unified_object_id},
data={
"file_object": file_object.model_dump_json(),
"status": file_object.status,
"updated_by": user_api_key_dict.user_id,
},
)
)
if updated_count == 0:
return
await self.internal_usage_cache.async_set_cache(
key=unified_object_id,
value=litellm_managed_object.model_dump(),
litellm_parent_otel_span=litellm_parent_otel_span,
)
return
await self.internal_usage_cache.async_set_cache(
key=unified_object_id,
value=litellm_managed_object.model_dump(),
@ -284,6 +314,103 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
detail=f"Object not found: {unified_object_id}",
)
async def enforce_batch_object_access(
self, object_id: str, user_api_key_dict: UserAPIKeyAuth
) -> None:
"""Deny access to a provider-format batch id owned by another caller.
Ids with no ownership row (batches created before ownership tracking,
or directly on the provider account) stay accessible so pass-through
reads keep working.
"""
if self.prisma_client is None:
return
managed_object = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={"unified_object_id": object_id}
)
)
if managed_object is None:
return
if not can_access_resource(
user_api_key_dict=user_api_key_dict,
created_by=managed_object.created_by,
resource_team_id=managed_object.team_id,
):
raise HTTPException(
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to the object {object_id}",
)
async def enforce_provider_file_access(
self, file_id: str, user_api_key_dict: UserAPIKeyAuth
) -> None:
"""Deny access to a provider-format file id owned by another caller.
Ownership rows for provider-format ids are written when a managed
batch's output/error files are first synced; ids with no row stay
accessible so pass-through reads keep working.
"""
if self.prisma_client is None:
return
managed_file = (
await self.prisma_client.db.litellm_managedfiletable.find_first(
where={"unified_file_id": file_id}
)
)
if managed_file is None:
return
if not can_access_resource(
user_api_key_dict=user_api_key_dict,
created_by=managed_file.created_by,
resource_team_id=managed_file.team_id,
):
raise HTTPException(
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to the file {file_id}",
)
async def store_batch_output_file_ownership(
self, response: LiteLLMBatch, litellm_parent_otel_span: Optional[Span]
) -> None:
"""Record ownership rows for a batch's provider-format output/error
file ids, inherited from the owning batch row (never the caller), so
file reads can be isolation-checked."""
provider_file_ids = tuple(
file_id
for file_id in (
getattr(response, "output_file_id", None),
getattr(response, "error_file_id", None),
)
if file_id and not _is_base64_encoded_unified_file_id(file_id)
)
if not provider_file_ids:
return
if self.prisma_client is None:
return
batch_row = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={"unified_object_id": response.id}
)
)
if batch_row is None or (
batch_row.created_by is None and batch_row.team_id is None
):
return
owner_identity = UserAPIKeyAuth(
user_id=batch_row.created_by, team_id=batch_row.team_id
)
for file_id in provider_file_ids:
model_name = decode_model_from_file_id(file_id)
raw_file_id = get_original_file_id(file_id)
await self.store_unified_file_id(
file_id=file_id,
file_object=None,
litellm_parent_otel_span=litellm_parent_otel_span,
model_mappings={model_name: raw_file_id} if model_name else {},
user_api_key_dict=owner_identity,
)
async def list_user_batches(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -397,6 +524,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to the file {retrieve_file_id}",
)
if retrieve_file_id:
await self.enforce_provider_file_access(
retrieve_file_id, user_api_key_dict
)
return False
async def check_file_ids_access(
@ -580,6 +711,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
data[accessor_key] = get_batch_id_from_unified_batch_id(
potential_llm_object_id
)
elif retrieve_object_id and accessor_key == "batch_id":
await self.enforce_batch_object_access(
retrieve_object_id, user_api_key_dict
)
elif call_type == CallTypes.acreate_fine_tuning_job.value:
input_file_id = cast(Optional[str], data.get("training_file"))
if input_file_id:
@ -1183,6 +1318,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
model_mappings={model_id: provider_file_id},
user_api_key_dict=user_api_key_dict,
)
is_batch_create = "completion_window" in data or "input_file_id" in data
await self.store_unified_object_id(
unified_object_id=response.id,
file_object=response,
@ -1190,7 +1326,13 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
model_object_id=original_response_id,
file_purpose="batch",
user_api_key_dict=user_api_key_dict,
update_only=not is_batch_create,
)
if not is_batch_create:
await self.store_batch_output_file_ownership(
response=response,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
)
# Only record batch creation metric on actual create (not retrieve/cancel).
# unified_file_id in _hidden_params is only set by the create_batch endpoint.

View file

@ -11,6 +11,7 @@ from litellm.caching import DualCache
from litellm.proxy._types import CallTypes
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
encode_file_id_with_model,
)
@ -2367,3 +2368,414 @@ async def test_same_user_different_keys_can_access_batch():
assert "batch_id" in result2
# Both keys should get the same result
assert result1["batch_id"] == result2["batch_id"]
MODEL_ENCODED_BATCH_ID = encode_file_id_with_model(
"batch_provider123", "gpt-4o-team-alias", id_type="batch"
)
MODEL_ENCODED_OUTPUT_FILE_ID = encode_file_id_with_model(
"file-output456", "gpt-4o-team-alias", id_type="file"
)
RAW_PROVIDER_BATCH_ID = "batch_provider123"
RAW_PROVIDER_FILE_ID = "file-output456"
def _owned_record(created_by, team_id):
record = MagicMock()
record.created_by = created_by
record.team_id = team_id
return record
def _batch_response(batch_id, output_file_id=None):
from litellm.types.utils import LiteLLMBatch
return LiteLLMBatch(
id=batch_id,
completion_window="24h",
created_at=1700000000,
endpoint="/v1/chat/completions",
input_file_id="file-input789",
object="batch",
status="completed",
output_file_id=output_file_id,
)
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", ["aretrieve_batch", "acancel_batch"])
@pytest.mark.parametrize(
"batch_id", [MODEL_ENCODED_BATCH_ID, RAW_PROVIDER_BATCH_ID]
)
async def test_team_b_cannot_access_team_a_provider_format_batch(
call_type, batch_id
):
"""
Cross-team retrieve/cancel of a model-encoded or raw provider batch id
must 403 when an ownership row exists for another team.
Regression test: before this check only unified (litellm_proxy-prefixed)
batch ids were enforced, so any key could read any model-encoded or raw
provider batch.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
with pytest.raises(HTTPException) as exc_info:
await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"batch_id": batch_id},
call_type=call_type,
)
assert exc_info.value.status_code == 403
prisma_client.db.litellm_managedobjecttable.find_first.assert_awaited_once_with(
where={"unified_object_id": batch_id}
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"caller_kwargs",
[
{"user_id": "user_a", "team_id": "team_a"},
{"user_id": "teammate_of_a", "team_id": "team_a"},
{"user_id": "admin_user", "user_role": "proxy_admin"},
],
)
async def test_authorized_callers_can_access_provider_format_batch(caller_kwargs):
"""
The creator, a same-team member, and a proxy admin can all retrieve a
model-encoded batch owned by team_a. Data must pass through unmodified so
the endpoint's own routing still applies.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
result = await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
parent_otel_span=MagicMock(), **caller_kwargs
),
cache=MagicMock(),
data={"batch_id": MODEL_ENCODED_BATCH_ID},
call_type="aretrieve_batch",
)
assert result["batch_id"] == MODEL_ENCODED_BATCH_ID
assert "model" not in result
@pytest.mark.asyncio
async def test_provider_format_batch_without_ownership_row_stays_accessible():
"""
A provider-format batch id with no ownership row (created before ownership
tracking, or directly on the provider account) must stay retrievable.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first.return_value = None
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
result = await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"batch_id": RAW_PROVIDER_BATCH_ID},
call_type="aretrieve_batch",
)
assert result["batch_id"] == RAW_PROVIDER_BATCH_ID
@pytest.mark.asyncio
async def test_fine_tuning_provider_format_id_not_enforced():
"""
Provider-format fine-tuning job ids are deliberately out of scope for
ownership enforcement; only unified fine-tuning ids are checked.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
result = await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"fine_tuning_job_id": "ftjob-abc123"},
call_type="aretrieve_fine_tuning_job",
)
assert result["fine_tuning_job_id"] == "ftjob-abc123"
prisma_client.db.litellm_managedobjecttable.find_first.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"call_type", ["afile_content", "afile_retrieve", "afile_delete"]
)
@pytest.mark.parametrize(
"file_id", [MODEL_ENCODED_OUTPUT_FILE_ID, RAW_PROVIDER_FILE_ID]
)
async def test_team_b_cannot_access_team_a_provider_format_file(
call_type, file_id
):
"""
Cross-team content/retrieve/delete of a model-encoded or raw provider
file id must 403 when an ownership row exists for another team.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedfiletable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(), prisma_client=prisma_client
)
with pytest.raises(HTTPException) as exc_info:
await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"file_id": file_id},
call_type=call_type,
)
assert exc_info.value.status_code == 403
prisma_client.db.litellm_managedfiletable.find_first.assert_awaited_once_with(
where={"unified_file_id": file_id}
)
@pytest.mark.asyncio
async def test_same_team_can_access_provider_format_file():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedfiletable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(), prisma_client=prisma_client
)
result = await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="teammate_of_a", team_id="team_a", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"file_id": MODEL_ENCODED_OUTPUT_FILE_ID},
call_type="afile_content",
)
assert result["file_id"] == MODEL_ENCODED_OUTPUT_FILE_ID
@pytest.mark.asyncio
async def test_provider_format_file_without_ownership_row_stays_accessible():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedfiletable.find_first.return_value = None
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(), prisma_client=prisma_client
)
result = await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"file_id": RAW_PROVIDER_FILE_ID},
call_type="afile_content",
)
assert result["file_id"] == RAW_PROVIDER_FILE_ID
@pytest.mark.asyncio
async def test_post_call_batch_create_stores_ownership_row():
"""
Batch creation (request data carries completion_window/input_file_id)
must write an ownership row attributed to the creating key.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
)
await proxy_managed_files.async_post_call_success_hook(
data={
"input_file_id": "file-input789",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
},
user_api_key_dict=UserAPIKeyAuth(
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
),
response=_batch_response(MODEL_ENCODED_BATCH_ID),
)
upsert_call = prisma_client.db.litellm_managedobjecttable.upsert.await_args
assert upsert_call.kwargs["where"] == {
"unified_object_id": MODEL_ENCODED_BATCH_ID
}
create_data = upsert_call.kwargs["data"]["create"]
assert create_data["created_by"] == "user_a"
assert create_data["team_id"] == "team_a"
prisma_client.db.litellm_managedobjecttable.update_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_post_call_batch_sync_does_not_claim_ownership():
"""
Retrieve/cancel of a batch with no ownership row must NOT create one:
otherwise the first foreign key to touch a legacy batch would become its
owner and lock out the real creator once enforcement is on.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 0
internal_usage_cache = MagicMock(async_set_cache=AsyncMock())
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
internal_usage_cache, prisma_client=prisma_client
)
await proxy_managed_files.async_post_call_success_hook(
data={"batch_id": MODEL_ENCODED_BATCH_ID},
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
response=_batch_response(MODEL_ENCODED_BATCH_ID),
)
prisma_client.db.litellm_managedobjecttable.upsert.assert_not_awaited()
internal_usage_cache.async_set_cache.assert_not_awaited()
@pytest.mark.asyncio
async def test_post_call_batch_sync_updates_existing_row():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 1
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
)
await proxy_managed_files.async_post_call_success_hook(
data={"batch_id": MODEL_ENCODED_BATCH_ID},
user_api_key_dict=UserAPIKeyAuth(
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
),
response=_batch_response(MODEL_ENCODED_BATCH_ID),
)
update_call = prisma_client.db.litellm_managedobjecttable.update_many.await_args
assert update_call.kwargs["where"] == {
"unified_object_id": MODEL_ENCODED_BATCH_ID
}
assert update_call.kwargs["data"]["status"] == "completed"
prisma_client.db.litellm_managedobjecttable.upsert.assert_not_awaited()
@pytest.mark.asyncio
async def test_post_call_batch_sync_stores_output_file_ownership_from_batch_row():
"""
When a synced batch reports a provider-format output file id, an
ownership row for that file must be written with the BATCH row's
created_by/team_id, not the caller's identity.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 1
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
)
await proxy_managed_files.async_post_call_success_hook(
data={"batch_id": MODEL_ENCODED_BATCH_ID},
user_api_key_dict=UserAPIKeyAuth(
user_id="admin_user",
user_role="proxy_admin",
parent_otel_span=MagicMock(),
),
response=_batch_response(
MODEL_ENCODED_BATCH_ID, output_file_id=MODEL_ENCODED_OUTPUT_FILE_ID
),
)
file_upsert = prisma_client.db.litellm_managedfiletable.upsert.await_args
assert file_upsert.kwargs["where"] == {
"unified_file_id": MODEL_ENCODED_OUTPUT_FILE_ID
}
create_data = file_upsert.kwargs["data"]["create"]
assert create_data["created_by"] == "user_a"
assert create_data["team_id"] == "team_a"
assert create_data["flat_model_file_ids"] == ["file-output456"]
@pytest.mark.asyncio
async def test_post_call_batch_create_does_not_store_output_file_ownership():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
)
await proxy_managed_files.async_post_call_success_hook(
data={
"input_file_id": "file-input789",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
},
user_api_key_dict=UserAPIKeyAuth(
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
),
response=_batch_response(
MODEL_ENCODED_BATCH_ID, output_file_id=MODEL_ENCODED_OUTPUT_FILE_ID
),
)
prisma_client.db.litellm_managedfiletable.upsert.assert_not_awaited()