mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
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:
parent
69a491e168
commit
f07ea5921b
2 changed files with 554 additions and 0 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue