mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Converge on staging's delete plumbing (_S3DeleteContext read from the logging call's additional_args, _sign_s3_request_without_body, the credential-stripping delete_data in the managed-files hook) and keep this PR's listing support, the 400 mapping for out-of-bucket file ids, the proxy-admin-only raw cloud id rule, and the OpenAI FileDeleted delete response. Two staging tests move to this PR's contract: an out-of-bucket delete raises BedrockError 400 instead of ValueError, and deleting a stored provider output returns FileDeleted rather than the stored file object.
1979 lines
85 KiB
Python
1979 lines
85 KiB
Python
# What is this?
|
|
## This hook is used to check for LiteLLM managed files in the request body, and replace them with model-specific file id
|
|
|
|
import base64
|
|
import json
|
|
from collections.abc import Mapping, Sequence
|
|
from types import MappingProxyType
|
|
from typing import (
|
|
TYPE_CHECKING,
|
|
Any,
|
|
Dict,
|
|
Final,
|
|
List,
|
|
Literal,
|
|
Optional,
|
|
Protocol,
|
|
TypedDict,
|
|
Union,
|
|
cast,
|
|
)
|
|
from uuid import NAMESPACE_URL, uuid5
|
|
|
|
from fastapi import HTTPException
|
|
from pydantic import ValidationError
|
|
|
|
import litellm
|
|
from litellm import Router, verbose_logger
|
|
from litellm._uuid import uuid
|
|
from litellm.caching.caching import DualCache
|
|
from litellm.constants import MAX_FILE_LIST_LIMIT
|
|
from litellm.integrations.custom_logger import CustomLogger
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|
extract_file_metadata,
|
|
)
|
|
from openai.types.file_deleted import FileDeleted
|
|
|
|
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
|
|
from litellm.llms.base_llm.managed_resources.isolation import (
|
|
build_list_page,
|
|
build_owner_filter,
|
|
can_access_resource,
|
|
resolve_resource_owner_id,
|
|
)
|
|
from litellm.proxy._types import (
|
|
CallTypes,
|
|
LiteLLM_ManagedFileTable,
|
|
LiteLLM_ManagedObjectTable,
|
|
ProxyException,
|
|
UserAPIKeyAuth,
|
|
)
|
|
from litellm.proxy.openai_files_endpoints.common_utils import (
|
|
BATCH_CREATE_HIDDEN_PARAM,
|
|
FILE_LIST_CONTINUATION_CHUNK_SIZE,
|
|
_is_base64_encoded_unified_file_id,
|
|
apply_unified_file_ids,
|
|
decode_model_from_file_id,
|
|
ensure_batch_response_managed_file_ids,
|
|
get_batch_id_from_unified_batch_id,
|
|
get_content_type_from_file_object,
|
|
get_model_id_from_unified_batch_id,
|
|
get_original_file_id,
|
|
map_raw_file_ids_to_unified,
|
|
normalize_mime_type_for_provider,
|
|
resolve_managed_output_file_model_name,
|
|
validate_file_list_limit,
|
|
validate_file_list_purpose,
|
|
)
|
|
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import (
|
|
request_tags_from_metadata,
|
|
)
|
|
from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccessIssue]
|
|
AllMessageValues,
|
|
AsyncCursorPage,
|
|
ChatCompletionFileObject,
|
|
CreateFileRequest,
|
|
FileListPage,
|
|
FileObject,
|
|
OpenAIFileObject,
|
|
ResponsesAPIResponse,
|
|
)
|
|
from litellm.types.utils import (
|
|
CallTypesLiteral,
|
|
LiteLLMBatch,
|
|
LiteLLMFineTuningJob,
|
|
LLMResponseTypes,
|
|
SpecialEnums,
|
|
)
|
|
|
|
if TYPE_CHECKING:
|
|
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
|
|
|
|
|
if TYPE_CHECKING:
|
|
from opentelemetry.trace import Span as _Span
|
|
from prisma.models import (
|
|
LiteLLM_ManagedObjectTable as PrismaManagedObjectRow,
|
|
)
|
|
|
|
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
|
|
from litellm.proxy.utils import PrismaClient as _PrismaClient
|
|
|
|
Span = Union[_Span, Any]
|
|
InternalUsageCache = _InternalUsageCache
|
|
PrismaClient = _PrismaClient
|
|
else:
|
|
Span = Any
|
|
InternalUsageCache = Any
|
|
PrismaClient = Any
|
|
|
|
|
|
def _sanitized_parse_error(e: Exception) -> str:
|
|
return (
|
|
str(e.errors(include_input=False, include_url=False, include_context=False))
|
|
if isinstance(e, ValidationError)
|
|
else type(e).__name__
|
|
)
|
|
|
|
|
|
def _decode_json_blob(blob: object) -> object:
|
|
return json.loads(blob) if isinstance(blob, str) else blob
|
|
|
|
|
|
def _parse_managed_batch_row(row: "PrismaManagedObjectRow") -> Optional[LiteLLMBatch]:
|
|
try:
|
|
batch_obj: Final = LiteLLMBatch.model_validate(_decode_json_blob(row.file_object))
|
|
except Exception as e:
|
|
verbose_logger.warning(f"Failed to parse batch object {row.unified_object_id}: {_sanitized_parse_error(e)}")
|
|
return None
|
|
batch_obj.id = row.unified_object_id
|
|
return batch_obj
|
|
|
|
|
|
def _parse_managed_file_object(raw_file_object: object, unified_file_id: str) -> Optional[OpenAIFileObject]:
|
|
if raw_file_object is None:
|
|
return None
|
|
try:
|
|
return OpenAIFileObject.model_validate(raw_file_object)
|
|
except Exception as e:
|
|
verbose_logger.warning(f"Failed to parse managed file object {unified_file_id}: {_sanitized_parse_error(e)}")
|
|
return None
|
|
|
|
|
|
class _ManagedFileRow(Protocol):
|
|
unified_file_id: str
|
|
file_object: OpenAIFileObject
|
|
storage_backend: Optional[str]
|
|
storage_url: Optional[str]
|
|
created_by: Optional[str]
|
|
team_id: Optional[str]
|
|
|
|
def model_dump(self) -> Mapping[str, object]: ...
|
|
|
|
|
|
class _ManagedFileTableActions(Protocol):
|
|
async def find_first(self, where: Mapping[str, object]) -> Optional[_ManagedFileRow]: ...
|
|
|
|
async def find_many(
|
|
self,
|
|
where: Mapping[str, object],
|
|
take: int = ...,
|
|
order: Union[Mapping[str, str], Sequence[Mapping[str, str]]] = ...,
|
|
cursor: Mapping[str, str] = ...,
|
|
skip: int = ...,
|
|
) -> Sequence[_ManagedFileRow]: ...
|
|
|
|
async def upsert(self, where: Mapping[str, str], data: Mapping[str, Mapping[str, object]]) -> _ManagedFileRow: ...
|
|
|
|
async def delete(self, where: Mapping[str, str]) -> Optional[_ManagedFileRow]: ...
|
|
|
|
|
|
class _ManagedObjectTableActions(Protocol):
|
|
async def find_first(self, where: Mapping[str, object]) -> "Optional[PrismaManagedObjectRow]": ...
|
|
|
|
async def find_many(
|
|
self,
|
|
where: Mapping[str, object],
|
|
take: int,
|
|
order: Union[Mapping[str, str], Sequence[Mapping[str, str]]],
|
|
cursor: Mapping[str, str] = ...,
|
|
skip: int = ...,
|
|
) -> "Sequence[PrismaManagedObjectRow]": ...
|
|
|
|
async def upsert(
|
|
self, where: Mapping[str, str], data: Mapping[str, Mapping[str, object]]
|
|
) -> "PrismaManagedObjectRow": ...
|
|
|
|
async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
|
|
|
|
|
|
class _SchedulerWithJobLookup(Protocol):
|
|
def get_job(self, job_id: str) -> object: ...
|
|
|
|
|
|
class _CursorPageArgs(TypedDict, total=False):
|
|
cursor: Mapping[str, str]
|
|
skip: int
|
|
|
|
|
|
def _managed_file_table(prisma_client: PrismaClient) -> _ManagedFileTableActions:
|
|
return prisma_client.db.litellm_managedfiletable
|
|
|
|
|
|
def _managed_object_table(prisma_client: PrismaClient) -> _ManagedObjectTableActions:
|
|
return prisma_client.db.litellm_managedobjecttable
|
|
|
|
|
|
class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|
# Class variables or attributes
|
|
def __init__(self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient):
|
|
self.internal_usage_cache = internal_usage_cache
|
|
self.prisma_client = prisma_client
|
|
|
|
@staticmethod
|
|
def _get_prometheus_logger():
|
|
"""Find PrometheusLogger from litellm.callbacks, if registered."""
|
|
from litellm.integrations.prometheus import PrometheusLogger
|
|
|
|
return PrometheusLogger.get_instance()
|
|
|
|
async def store_unified_file_id(
|
|
self,
|
|
file_id: str,
|
|
file_object: Optional[OpenAIFileObject],
|
|
litellm_parent_otel_span: Optional[Span],
|
|
model_mappings: Dict[str, str],
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> None:
|
|
verbose_logger.info(f"Storing LiteLLM Managed File object with id={file_id} in cache")
|
|
if file_object is not None:
|
|
litellm_managed_file_object = LiteLLM_ManagedFileTable(
|
|
unified_file_id=file_id,
|
|
file_object=file_object,
|
|
model_mappings=model_mappings,
|
|
flat_model_file_ids=list(model_mappings.values()),
|
|
created_by=resolve_resource_owner_id(user_api_key_dict),
|
|
team_id=user_api_key_dict.team_id,
|
|
updated_by=user_api_key_dict.user_id,
|
|
)
|
|
await self.internal_usage_cache.async_set_cache(
|
|
key=file_id,
|
|
value=litellm_managed_file_object.model_dump(),
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
)
|
|
|
|
## STORE MODEL MAPPINGS IN DB
|
|
|
|
db_data = {
|
|
"unified_file_id": file_id,
|
|
"model_mappings": json.dumps(model_mappings),
|
|
"flat_model_file_ids": list(model_mappings.values()),
|
|
"created_by": resolve_resource_owner_id(user_api_key_dict),
|
|
"team_id": user_api_key_dict.team_id,
|
|
"updated_by": user_api_key_dict.user_id,
|
|
}
|
|
update_data = {
|
|
"model_mappings": json.dumps(model_mappings),
|
|
"flat_model_file_ids": list(model_mappings.values()),
|
|
"updated_by": user_api_key_dict.user_id,
|
|
}
|
|
|
|
if file_object is not None:
|
|
file_object_json = file_object.model_dump_json()
|
|
db_data["file_object"] = file_object_json
|
|
update_data["file_object"] = file_object_json
|
|
# Extract storage metadata from hidden params if present
|
|
hidden_params = getattr(file_object, "_hidden_params", {}) or {}
|
|
if "storage_backend" in hidden_params:
|
|
db_data["storage_backend"] = hidden_params["storage_backend"]
|
|
update_data["storage_backend"] = hidden_params["storage_backend"]
|
|
if "storage_url" in hidden_params:
|
|
db_data["storage_url"] = hidden_params["storage_url"]
|
|
update_data["storage_url"] = hidden_params["storage_url"]
|
|
|
|
verbose_logger.debug(
|
|
f"Storage metadata: storage_backend={db_data.get('storage_backend')}, "
|
|
f"storage_url={db_data.get('storage_url')}"
|
|
)
|
|
|
|
result = await _managed_file_table(self.prisma_client).upsert(
|
|
where={"unified_file_id": file_id},
|
|
data={"create": db_data, "update": update_data},
|
|
)
|
|
verbose_logger.debug(f"LiteLLM Managed File object with id={file_id} stored in db: {result}")
|
|
|
|
async def _resolve_creator_org_id(self, user_api_key_dict: UserAPIKeyAuth) -> Optional[str]:
|
|
if user_api_key_dict.org_id:
|
|
return user_api_key_dict.org_id
|
|
if not user_api_key_dict.team_id:
|
|
return None
|
|
from litellm.proxy.auth.auth_checks import get_team_object
|
|
from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
|
|
|
|
try:
|
|
team: Final = await get_team_object(
|
|
team_id=user_api_key_dict.team_id,
|
|
prisma_client=self.prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
parent_otel_span=user_api_key_dict.parent_otel_span,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
return team.organization_id
|
|
except Exception as e:
|
|
verbose_logger.warning(f"could not resolve org for managed object attribution: {e}")
|
|
return None
|
|
|
|
async def store_unified_object_id(
|
|
self,
|
|
unified_object_id: str,
|
|
file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob, "ResponsesAPIResponse"],
|
|
litellm_parent_otel_span: Optional[Span],
|
|
model_object_id: str,
|
|
file_purpose: Literal["batch", "fine-tune", "response"],
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
request_tags: Sequence[str] | None = None,
|
|
persist_attribution: bool = False,
|
|
create_if_missing: bool = True,
|
|
) -> None:
|
|
"""Persist a managed object row, caching it and upserting it in the DB.
|
|
|
|
persist_attribution is set only by the batch create, which is the one caller
|
|
that can speak for the creator; it gates the api_key and request_tags columns
|
|
that CheckBatchCost bills against, so a later poll or retrieve of the same
|
|
batch cannot record itself as the paying key. Like created_by and team_id,
|
|
both are written only in the upsert create branch, never on update.
|
|
|
|
create_if_missing is cleared by callers that observe a batch they did not
|
|
create, such as a poll. They still refresh status and file_object, but a
|
|
row absent from the table is left absent rather than created with the
|
|
observer as its creator, because created_by and team_id are written from
|
|
whoever calls the create branch.
|
|
"""
|
|
verbose_logger.info(f"Storing LiteLLM Managed {file_purpose} object with id={unified_object_id} in cache")
|
|
litellm_managed_object = LiteLLM_ManagedObjectTable(
|
|
unified_object_id=unified_object_id,
|
|
model_object_id=model_object_id,
|
|
file_purpose=file_purpose,
|
|
file_object=file_object,
|
|
)
|
|
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,
|
|
)
|
|
|
|
from prisma import Json
|
|
|
|
api_key = user_api_key_dict.api_key or None
|
|
attribution_columns = (
|
|
{
|
|
**({"api_key": api_key} if api_key is not None else {}),
|
|
**({"request_tags": Json(list(request_tags))} if request_tags else {}),
|
|
}
|
|
if persist_attribution
|
|
else {}
|
|
)
|
|
# FIX: Update status and file_object on every operation to keep state in sync
|
|
update_columns: Final = {
|
|
"file_object": file_object.model_dump_json(),
|
|
"status": file_object.status,
|
|
"updated_by": user_api_key_dict.user_id,
|
|
}
|
|
if not create_if_missing:
|
|
await _managed_object_table(self.prisma_client).update_many(
|
|
where={"unified_object_id": unified_object_id},
|
|
data=update_columns,
|
|
)
|
|
return
|
|
await _managed_object_table(self.prisma_client).upsert(
|
|
where={"unified_object_id": unified_object_id},
|
|
data={
|
|
"create": {
|
|
"unified_object_id": unified_object_id,
|
|
"file_object": file_object.model_dump_json(),
|
|
"model_object_id": model_object_id,
|
|
"file_purpose": file_purpose,
|
|
"created_by": resolve_resource_owner_id(user_api_key_dict),
|
|
"team_id": user_api_key_dict.team_id,
|
|
"org_id": await self._resolve_creator_org_id(user_api_key_dict),
|
|
"updated_by": user_api_key_dict.user_id,
|
|
"status": file_object.status,
|
|
**attribution_columns,
|
|
},
|
|
"update": update_columns,
|
|
},
|
|
)
|
|
|
|
async def get_unified_file_id(
|
|
self, file_id: str, litellm_parent_otel_span: Optional[Span] = None
|
|
) -> Optional[LiteLLM_ManagedFileTable]:
|
|
## CHECK CACHE
|
|
result = cast(
|
|
Optional[dict],
|
|
await self.internal_usage_cache.async_get_cache(
|
|
key=file_id,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
),
|
|
)
|
|
|
|
if result:
|
|
return LiteLLM_ManagedFileTable.model_validate(result)
|
|
|
|
## CHECK DB
|
|
db_object = await _managed_file_table(self.prisma_client).find_first(where={"unified_file_id": file_id})
|
|
|
|
if db_object:
|
|
return LiteLLM_ManagedFileTable.model_validate(db_object.model_dump())
|
|
return None
|
|
|
|
async def delete_unified_file_id(
|
|
self, file_id: str, litellm_parent_otel_span: Optional[Span] = None
|
|
) -> OpenAIFileObject:
|
|
## get old value
|
|
initial_value = await _managed_file_table(self.prisma_client).find_first(where={"unified_file_id": file_id})
|
|
if initial_value is None:
|
|
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
|
|
## delete old value
|
|
await self.internal_usage_cache.async_set_cache(
|
|
key=file_id,
|
|
value=None,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
)
|
|
await _managed_file_table(self.prisma_client).delete(where={"unified_file_id": file_id})
|
|
return initial_value.file_object
|
|
|
|
async def can_user_call_unified_file_id(self, unified_file_id: str, user_api_key_dict: UserAPIKeyAuth) -> bool:
|
|
managed_file = await _managed_file_table(self.prisma_client).find_first(
|
|
where={"unified_file_id": unified_file_id}
|
|
)
|
|
|
|
if managed_file:
|
|
return 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=404,
|
|
detail=f"File not found: {unified_file_id}",
|
|
)
|
|
|
|
async def can_user_call_unified_object_id(self, unified_object_id: str, user_api_key_dict: UserAPIKeyAuth) -> bool:
|
|
managed_object = await _managed_object_table(self.prisma_client).find_first(
|
|
where={"unified_object_id": unified_object_id}
|
|
)
|
|
|
|
if managed_object:
|
|
return 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=404,
|
|
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={"OR": [{"unified_object_id": object_id}, {"model_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={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": 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,
|
|
limit: Optional[int] = None,
|
|
after: Optional[str] = None,
|
|
provider: Optional[str] = None,
|
|
target_model_names: Optional[str] = None,
|
|
llm_router: Optional[Router] = None,
|
|
) -> Dict[str, object]:
|
|
# Provider filtering is not supported for managed batches
|
|
# This is because the encoded object ids stored in the managed objects table do not contain the provider information
|
|
# To support provider filtering, we would need to store the provider information in the encoded object ids
|
|
if provider:
|
|
raise ProxyException(
|
|
message="Filtering by 'provider' is not supported when using managed batches.",
|
|
type="invalid_request_error",
|
|
param="provider",
|
|
code=400,
|
|
)
|
|
|
|
# Model name filtering is not supported for managed batches
|
|
# This is because the encoded object ids stored in the managed objects table do not contain the model name
|
|
# A hash of the model name + litellm_params for the model name is encoded as the model id. This is not sufficient to reliably map the target model names to the model ids.
|
|
if target_model_names:
|
|
raise ProxyException(
|
|
message="Filtering by 'target_model_names' is not supported when using managed batches.",
|
|
type="invalid_request_error",
|
|
param="target_model_names",
|
|
code=400,
|
|
)
|
|
|
|
if limit == 0:
|
|
return build_list_page([])
|
|
|
|
owner_filter = build_owner_filter(user_api_key_dict)
|
|
if owner_filter is None:
|
|
return build_list_page([])
|
|
|
|
where_clause: Dict[str, object] = {"file_purpose": "batch", **owner_filter}
|
|
|
|
if after:
|
|
cursor_row = await _managed_object_table(self.prisma_client).find_first(
|
|
where={**where_clause, "unified_object_id": after}
|
|
)
|
|
if cursor_row is None:
|
|
raise HTTPException(
|
|
status_code=400,
|
|
detail=f"Invalid 'after' cursor: no batch found with id '{after}'.",
|
|
)
|
|
|
|
page_size: Final = min(limit or 20, 100)
|
|
matches: Final = await self._collect_listed_batches(
|
|
where_clause=where_clause,
|
|
after=after,
|
|
wanted=page_size + 1,
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
return build_list_page(list(matches[:page_size]), has_more=len(matches) > page_size)
|
|
|
|
async def _collect_listed_batches(
|
|
self,
|
|
where_clause: Mapping[str, object],
|
|
after: Optional[str],
|
|
wanted: int,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> tuple[LiteLLMBatch, ...]:
|
|
"""Read chunks newest-first until ``wanted`` batches survive parsing and
|
|
file-id resolution or the caller's rows run out, so a run of rows that will
|
|
not parse refills the page instead of emptying it. The first chunk is
|
|
``wanted`` rows, so a healthy page still costs one query; a scan that has to
|
|
continue widens to ``FILE_LIST_CONTINUATION_CHUNK_SIZE`` like ``afile_list``,
|
|
and every chunk advances the keyset cursor, so the walk ends once the
|
|
caller's rows are exhausted."""
|
|
matches: tuple[LiteLLMBatch, ...] = () # rebind-ok: accumulates survivors across chunks
|
|
cursor_id: Optional[str] = after # rebind-ok: keyset cursor advances to each chunk's last row
|
|
chunk_size: int = wanted # rebind-ok: widens once a scan has to continue past the first chunk
|
|
while len(matches) < wanted:
|
|
cursor_args: _CursorPageArgs = {"cursor": {"unified_object_id": cursor_id}, "skip": 1} if cursor_id else {}
|
|
chunk = await _managed_object_table(self.prisma_client).find_many(
|
|
where=where_clause,
|
|
take=chunk_size,
|
|
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
|
|
**cursor_args,
|
|
)
|
|
matches = matches + await self._resolve_listed_rows(
|
|
rows=chunk, wanted=wanted - len(matches), user_api_key_dict=user_api_key_dict
|
|
)
|
|
if len(chunk) < chunk_size:
|
|
break
|
|
cursor_id = chunk[-1].unified_object_id
|
|
chunk_size = max(chunk_size, FILE_LIST_CONTINUATION_CHUNK_SIZE)
|
|
return matches
|
|
|
|
async def _resolve_listed_rows(
|
|
self,
|
|
rows: "Sequence[PrismaManagedObjectRow]",
|
|
wanted: int,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> tuple[LiteLLMBatch, ...]:
|
|
parsed_rows: Final = tuple(
|
|
(row, batch_obj) for row in rows if (batch_obj := _parse_managed_batch_row(row)) is not None
|
|
)
|
|
unified_id_by_raw_id: Final = await map_raw_file_ids_to_unified(
|
|
raw_file_ids=frozenset(
|
|
file_id
|
|
for _, batch_obj in parsed_rows
|
|
for file_id in (batch_obj.input_file_id, batch_obj.output_file_id, batch_obj.error_file_id)
|
|
if file_id and not _is_base64_encoded_unified_file_id(file_id)
|
|
),
|
|
prisma_client=self.prisma_client,
|
|
)
|
|
resolved: Final[list[LiteLLMBatch]] = [] # mutable-ok: resolution stops as soon as the page is full
|
|
for row, batch_obj in parsed_rows:
|
|
if len(resolved) == wanted:
|
|
break
|
|
resolved_batch = await self._resolve_listed_batch(
|
|
row=row,
|
|
batch_obj=batch_obj,
|
|
unified_id_by_raw_id=unified_id_by_raw_id,
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
if resolved_batch is not None:
|
|
resolved.append(resolved_batch)
|
|
return tuple(resolved)
|
|
|
|
async def _resolve_listed_batch(
|
|
self,
|
|
row: "PrismaManagedObjectRow",
|
|
batch_obj: LiteLLMBatch,
|
|
unified_id_by_raw_id: Mapping[str, str],
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> Optional[LiteLLMBatch]:
|
|
apply_unified_file_ids(batch_obj, unified_id_by_raw_id)
|
|
try:
|
|
await ensure_batch_response_managed_file_ids(
|
|
response=batch_obj,
|
|
managed_files_obj=self,
|
|
prisma_client=self.prisma_client,
|
|
verbose_proxy_logger=verbose_logger,
|
|
user_api_key_dict=user_api_key_dict,
|
|
db_batch_object=row,
|
|
unified_batch_id=_is_base64_encoded_unified_file_id(row.unified_object_id),
|
|
)
|
|
except Exception as e:
|
|
verbose_logger.warning(f"Failed to resolve managed file ids for batch {row.unified_object_id}: {e}")
|
|
return None
|
|
return batch_obj
|
|
|
|
async def get_user_created_file_ids(
|
|
self, user_api_key_dict: UserAPIKeyAuth, model_object_ids: List[str]
|
|
) -> List[OpenAIFileObject]:
|
|
"""
|
|
Get all file ids the caller is allowed to see for a list of model
|
|
object ids. Service-account keys (no user_id) are scoped to their
|
|
team via ``team_id``; admins see all matches.
|
|
|
|
Returns:
|
|
- List of OpenAIFileObject's
|
|
"""
|
|
owner_filter = build_owner_filter(user_api_key_dict)
|
|
if owner_filter is None:
|
|
return []
|
|
|
|
file_ids = await _managed_file_table(self.prisma_client).find_many(
|
|
where={
|
|
**owner_filter,
|
|
"flat_model_file_ids": {"hasSome": model_object_ids},
|
|
}
|
|
)
|
|
return [
|
|
parsed_file_object.model_copy(update={"id": row.unified_file_id})
|
|
for row in file_ids
|
|
if (parsed_file_object := _parse_managed_file_object(row.file_object, row.unified_file_id)) is not None
|
|
]
|
|
|
|
async def check_managed_file_id_access(self, data: Dict, user_api_key_dict: UserAPIKeyAuth) -> bool:
|
|
retrieve_file_id = cast(Optional[str], data.get("file_id"))
|
|
potential_file_id = _is_base64_encoded_unified_file_id(retrieve_file_id) if retrieve_file_id else False
|
|
if potential_file_id and retrieve_file_id:
|
|
if await self.can_user_call_unified_file_id(retrieve_file_id, user_api_key_dict):
|
|
return True
|
|
else:
|
|
raise HTTPException(
|
|
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(self, file_ids: List[str], user_api_key_dict: UserAPIKeyAuth) -> None:
|
|
"""
|
|
Check if the user has access to a list of file IDs.
|
|
Only checks managed (unified) file IDs.
|
|
|
|
Args:
|
|
file_ids: List of file IDs to check access for
|
|
user_api_key_dict: User API key authentication details
|
|
|
|
Raises:
|
|
HTTPException: If user doesn't have access to any of the files
|
|
"""
|
|
for file_id in file_ids:
|
|
is_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
|
|
if is_unified_file_id:
|
|
if not await self.can_user_call_unified_file_id(file_id, user_api_key_dict):
|
|
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 async_pre_call_hook(
|
|
self,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
cache: DualCache,
|
|
data: Dict,
|
|
call_type: CallTypesLiteral,
|
|
) -> Union[Exception, str, Dict, None]:
|
|
"""
|
|
- Detect litellm_proxy/ file_id
|
|
- add dictionary of mappings of litellm_proxy/ file_id -> provider_file_id => {litellm_proxy/file_id: {"model_id": id, "file_id": provider_file_id}}
|
|
"""
|
|
### HANDLE FILE ACCESS ### - ensure user has access to the file
|
|
if (
|
|
call_type == CallTypes.afile_content.value
|
|
or call_type == CallTypes.afile_delete.value
|
|
or call_type == CallTypes.afile_retrieve.value
|
|
or call_type == CallTypes.afile_content.value
|
|
):
|
|
await self.check_managed_file_id_access(data, user_api_key_dict)
|
|
|
|
### HANDLE TRANSFORMATIONS ###
|
|
# Check both completion and acompletion call types
|
|
is_completion_call = call_type == CallTypes.completion.value or call_type == CallTypes.acompletion.value
|
|
|
|
if is_completion_call:
|
|
messages = data.get("messages")
|
|
model = data.get("model", "")
|
|
if messages:
|
|
file_ids = self.get_file_ids_from_messages(messages)
|
|
if file_ids:
|
|
# Check user has access to all managed files
|
|
await self.check_file_ids_access(file_ids, user_api_key_dict)
|
|
|
|
# Check if any files are stored in storage backends and need base64 conversion
|
|
# This is needed for Vertex AI/Gemini which requires base64 content
|
|
is_vertex_ai = model and ("vertex_ai" in model or "gemini" in model.lower())
|
|
if is_vertex_ai:
|
|
await self._convert_storage_files_to_base64(
|
|
messages=messages,
|
|
file_ids=file_ids,
|
|
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
|
)
|
|
|
|
model_file_id_mapping = await self.get_model_file_id_mapping(
|
|
file_ids, user_api_key_dict.parent_otel_span
|
|
)
|
|
data["model_file_id_mapping"] = model_file_id_mapping
|
|
elif call_type == CallTypes.aresponses.value or call_type == CallTypes.responses.value:
|
|
# Handle managed files in responses API input and tools
|
|
file_ids = []
|
|
|
|
# Extract file IDs from input parameter
|
|
input_data = data.get("input")
|
|
if input_data:
|
|
file_ids.extend(self.get_file_ids_from_responses_input(input_data))
|
|
|
|
# Extract file IDs from tools parameter (e.g., code_interpreter container)
|
|
tools = data.get("tools")
|
|
if tools:
|
|
file_ids.extend(self.get_file_ids_from_responses_tools(tools))
|
|
|
|
if file_ids:
|
|
# Check user has access to all managed files
|
|
await self.check_file_ids_access(file_ids, user_api_key_dict)
|
|
|
|
model_file_id_mapping = await self.get_model_file_id_mapping(
|
|
file_ids, user_api_key_dict.parent_otel_span
|
|
)
|
|
data["model_file_id_mapping"] = model_file_id_mapping
|
|
|
|
# Check access for file_search vector_store_ids
|
|
if tools:
|
|
unified_vs_ids = self.get_vector_store_ids_from_file_search_tools(tools)
|
|
if unified_vs_ids:
|
|
await self.check_vector_store_ids_access(unified_vs_ids, user_api_key_dict)
|
|
elif call_type == CallTypes.afile_content.value:
|
|
retrieve_file_id = cast(Optional[str], data.get("file_id"))
|
|
potential_file_id = _is_base64_encoded_unified_file_id(retrieve_file_id) if retrieve_file_id else False
|
|
if potential_file_id and "llm_output_file_id," in potential_file_id:
|
|
model_id = self.get_model_id_from_unified_file_id(potential_file_id)
|
|
if model_id:
|
|
data["model"] = model_id
|
|
data["file_id"] = self.get_output_file_id_from_unified_file_id(potential_file_id)
|
|
elif call_type == CallTypes.acreate_batch.value:
|
|
input_file_id = cast(Optional[str], data.get("input_file_id"))
|
|
if input_file_id:
|
|
model_file_id_mapping = await self.get_model_file_id_mapping(
|
|
[input_file_id], user_api_key_dict.parent_otel_span
|
|
)
|
|
|
|
data["model_file_id_mapping"] = model_file_id_mapping
|
|
elif (
|
|
call_type == CallTypes.aretrieve_batch.value
|
|
or call_type == CallTypes.acancel_batch.value
|
|
or call_type == CallTypes.acancel_fine_tuning_job.value
|
|
or call_type == CallTypes.aretrieve_fine_tuning_job.value
|
|
):
|
|
accessor_key: Optional[str] = None
|
|
retrieve_object_id: Optional[str] = None
|
|
if call_type == CallTypes.aretrieve_batch.value or call_type == CallTypes.acancel_batch.value:
|
|
accessor_key = "batch_id"
|
|
elif (
|
|
call_type == CallTypes.acancel_fine_tuning_job.value
|
|
or call_type == CallTypes.aretrieve_fine_tuning_job.value
|
|
):
|
|
accessor_key = "fine_tuning_job_id"
|
|
|
|
if accessor_key:
|
|
retrieve_object_id = cast(Optional[str], data.get(accessor_key))
|
|
|
|
potential_llm_object_id = (
|
|
_is_base64_encoded_unified_file_id(retrieve_object_id) if retrieve_object_id else False
|
|
)
|
|
if potential_llm_object_id and retrieve_object_id:
|
|
## VALIDATE USER HAS ACCESS TO THE OBJECT ##
|
|
if not await self.can_user_call_unified_object_id(retrieve_object_id, user_api_key_dict):
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=f"User {user_api_key_dict.user_id} does not have access to the object {retrieve_object_id}",
|
|
)
|
|
|
|
## for managed batch id - get the model id
|
|
potential_model_id = get_model_id_from_unified_batch_id(potential_llm_object_id)
|
|
if potential_model_id is None:
|
|
raise Exception(
|
|
f"LiteLLM Managed {accessor_key} with id={retrieve_object_id} is invalid - does not contain encoded model_id."
|
|
)
|
|
data["model"] = potential_model_id
|
|
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:
|
|
model_file_id_mapping = await self.get_model_file_id_mapping(
|
|
[input_file_id], user_api_key_dict.parent_otel_span
|
|
)
|
|
|
|
return data
|
|
|
|
async def async_filter_deployments(
|
|
self,
|
|
model: str,
|
|
healthy_deployments: List,
|
|
messages: Optional[List[AllMessageValues]],
|
|
request_kwargs: Optional[Dict] = None,
|
|
parent_otel_span: Optional[Span] = None,
|
|
) -> List[Dict]:
|
|
if request_kwargs is None:
|
|
return healthy_deployments
|
|
|
|
input_file_id = cast(Optional[str], request_kwargs.get("input_file_id"))
|
|
model_file_id_mapping = cast(
|
|
Optional[Dict[str, Dict[str, str]]],
|
|
request_kwargs.get("model_file_id_mapping"),
|
|
)
|
|
allowed_model_ids = []
|
|
if input_file_id and model_file_id_mapping:
|
|
model_id_dict = model_file_id_mapping.get(input_file_id, {})
|
|
allowed_model_ids = list(model_id_dict.keys())
|
|
|
|
if len(allowed_model_ids) == 0:
|
|
return healthy_deployments
|
|
|
|
return [
|
|
deployment
|
|
for deployment in healthy_deployments
|
|
if deployment.get("model_info", {}).get("id") in allowed_model_ids
|
|
]
|
|
|
|
async def async_pre_call_deployment_hook(
|
|
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
|
|
) -> Optional[dict]:
|
|
"""
|
|
Allow modifying the request just before it's sent to the deployment.
|
|
"""
|
|
accessor_key: Optional[str] = None
|
|
if call_type and call_type == CallTypes.acreate_batch:
|
|
accessor_key = "input_file_id"
|
|
elif call_type and call_type == CallTypes.acreate_fine_tuning_job:
|
|
accessor_key = "training_file"
|
|
else:
|
|
return kwargs
|
|
|
|
if accessor_key:
|
|
input_file_id = cast(Optional[str], kwargs.get(accessor_key))
|
|
model_file_id_mapping = cast(Optional[Dict[str, Dict[str, str]]], kwargs.get("model_file_id_mapping"))
|
|
# model_info may be at top-level or nested under litellm_metadata
|
|
# (batch/file operations use litellm_metadata)
|
|
model_id = cast(Optional[str], kwargs.get("model_info", {}).get("id", None))
|
|
if model_id is None:
|
|
model_id = cast(
|
|
Optional[str],
|
|
kwargs.get("litellm_metadata", {}).get("model_info", {}).get("id", None),
|
|
)
|
|
mapped_file_id: Optional[str] = None
|
|
if input_file_id and model_file_id_mapping and model_id:
|
|
mapped_file_id = model_file_id_mapping.get(input_file_id, {}).get(model_id, None)
|
|
if mapped_file_id:
|
|
kwargs[accessor_key] = mapped_file_id
|
|
|
|
return kwargs
|
|
|
|
def get_file_ids_from_messages(self, messages: List[AllMessageValues]) -> List[str]:
|
|
"""
|
|
Gets file ids from messages
|
|
"""
|
|
file_ids = []
|
|
for message in messages:
|
|
if message.get("role") == "user":
|
|
content = message.get("content")
|
|
if content:
|
|
if isinstance(content, str):
|
|
continue
|
|
for c in content:
|
|
if c.get("type") == "file":
|
|
file_object = cast(ChatCompletionFileObject, c)
|
|
file_object_file_field = file_object["file"]
|
|
file_id = file_object_file_field.get("file_id")
|
|
if file_id:
|
|
file_ids.append(file_id)
|
|
return file_ids
|
|
|
|
def get_file_ids_from_responses_input(self, input: Union[str, List[Dict[str, object]]]) -> List[str]:
|
|
"""
|
|
Gets file ids from responses API input.
|
|
|
|
The input can be:
|
|
- A string (no files)
|
|
- A list of input items, where each item can have:
|
|
- type: "input_file" with file_id
|
|
- content: a list that can contain items with type: "input_file" and file_id
|
|
"""
|
|
file_ids: List[str] = []
|
|
|
|
if isinstance(input, str):
|
|
return file_ids
|
|
|
|
if not isinstance(input, list):
|
|
return file_ids
|
|
|
|
for item in input:
|
|
if not isinstance(item, dict):
|
|
continue
|
|
|
|
# Check for direct input_file type
|
|
if item.get("type") == "input_file":
|
|
file_id = item.get("file_id")
|
|
if isinstance(file_id, str) and file_id:
|
|
file_ids.append(file_id)
|
|
|
|
# Check for input_file in content array
|
|
content = item.get("content")
|
|
if isinstance(content, list):
|
|
for content_item in content:
|
|
if isinstance(content_item, dict) and content_item.get("type") == "input_file":
|
|
file_id = content_item.get("file_id")
|
|
if isinstance(file_id, str) and file_id:
|
|
file_ids.append(file_id)
|
|
|
|
return file_ids
|
|
|
|
def get_file_ids_from_responses_tools(self, tools: List[Dict[str, object]]) -> List[str]:
|
|
"""
|
|
Gets file ids from responses API tools parameter.
|
|
|
|
The tools can contain code_interpreter with container.file_ids:
|
|
[
|
|
{
|
|
"type": "code_interpreter",
|
|
"container": {"type": "auto", "file_ids": ["file-123", "file-456"]}
|
|
}
|
|
]
|
|
"""
|
|
file_ids: List[str] = []
|
|
|
|
if not isinstance(tools, list):
|
|
return file_ids
|
|
|
|
for tool in tools:
|
|
if not isinstance(tool, dict):
|
|
continue
|
|
|
|
# Check for code_interpreter with container file_ids
|
|
if tool.get("type") == "code_interpreter":
|
|
container = tool.get("container")
|
|
if isinstance(container, dict):
|
|
container_file_ids = container.get("file_ids")
|
|
if isinstance(container_file_ids, list):
|
|
for file_id in container_file_ids:
|
|
if isinstance(file_id, str):
|
|
file_ids.append(file_id)
|
|
|
|
return file_ids
|
|
|
|
def get_vector_store_ids_from_file_search_tools(self, tools: List[Dict[str, object]]) -> List[str]:
|
|
"""
|
|
Extract unified vector_store_ids from file_search tools.
|
|
|
|
Only returns IDs that are LiteLLM-managed (base64 unified IDs).
|
|
Native provider IDs are skipped — they have no LiteLLM access record.
|
|
"""
|
|
from litellm.llms.base_llm.managed_resources.utils import (
|
|
is_base64_encoded_unified_id,
|
|
)
|
|
|
|
vs_ids: List[str] = []
|
|
if not isinstance(tools, list):
|
|
return vs_ids
|
|
|
|
for tool in tools:
|
|
if not isinstance(tool, dict) or tool.get("type") != "file_search":
|
|
continue
|
|
vector_store_ids = tool.get("vector_store_ids")
|
|
if not isinstance(vector_store_ids, list):
|
|
continue
|
|
for vs_id in vector_store_ids:
|
|
if isinstance(vs_id, str) and is_base64_encoded_unified_id(vs_id):
|
|
vs_ids.append(vs_id)
|
|
|
|
return vs_ids
|
|
|
|
async def check_vector_store_ids_access(
|
|
self,
|
|
vector_store_ids: List[str],
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> None:
|
|
"""
|
|
Verify the caller's team can access each LiteLLM-managed vector store.
|
|
|
|
Batch-fetches vector stores from DB and checks team_id.
|
|
Raises HTTPException(403) on the first access violation.
|
|
Non-managed (native) IDs should already be filtered out before calling this.
|
|
"""
|
|
from litellm.llms.base_llm.managed_resources.utils import (
|
|
extract_unified_uuid_from_unified_id,
|
|
)
|
|
from litellm.proxy.auth.auth_checks import (
|
|
get_managed_vector_store_rows_by_uuids,
|
|
)
|
|
from litellm.proxy.proxy_server import (
|
|
prisma_client,
|
|
proxy_logging_obj,
|
|
user_api_key_cache,
|
|
)
|
|
|
|
if not vector_store_ids or prisma_client is None:
|
|
return
|
|
|
|
# Map each unified ID to its internal UUID for a single batch DB fetch
|
|
uuid_to_unified: Dict[str, str] = {}
|
|
for vs_id in vector_store_ids:
|
|
uuid = extract_unified_uuid_from_unified_id(vs_id)
|
|
if uuid:
|
|
uuid_to_unified[uuid] = vs_id
|
|
|
|
if not uuid_to_unified:
|
|
return
|
|
|
|
rows = await get_managed_vector_store_rows_by_uuids(
|
|
uuids=list(uuid_to_unified.keys()),
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
found_uuids = {row.vector_store_id for row in rows}
|
|
|
|
for uuid, original_id in uuid_to_unified.items():
|
|
if uuid not in found_uuids:
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=f"Vector store '{original_id}' not found or access denied.",
|
|
)
|
|
|
|
caller_team_id = user_api_key_dict.team_id
|
|
for row in rows:
|
|
vs_team_id = getattr(row, "team_id", None)
|
|
if vs_team_id is not None and vs_team_id != caller_team_id:
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=(
|
|
f"Team '{caller_team_id}' does not have access to vector "
|
|
f"store '{row.vector_store_id}'. The store belongs to team "
|
|
f"'{vs_team_id}'."
|
|
),
|
|
)
|
|
|
|
async def get_model_file_id_mapping(self, file_ids: List[str], litellm_parent_otel_span: Span) -> dict:
|
|
"""
|
|
Get model-specific file IDs for a list of proxy file IDs.
|
|
Returns a dictionary mapping litellm_proxy/ file_id -> model_id -> model_file_id
|
|
|
|
1. Get all the litellm_proxy/ file_ids from the messages
|
|
2. For each file_id, search for cache keys matching the pattern file_id:*
|
|
3. Return a dictionary of mappings of litellm_proxy/ file_id -> model_id -> model_file_id
|
|
|
|
Example:
|
|
{
|
|
"litellm_proxy/file_id": {
|
|
"model_id": "model_file_id"
|
|
}
|
|
}
|
|
"""
|
|
|
|
file_id_mapping: Dict[str, Dict[str, str]] = {}
|
|
litellm_managed_file_ids = []
|
|
|
|
for file_id in file_ids:
|
|
## CHECK IF FILE ID IS MANAGED BY LITELM
|
|
is_base64_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
|
|
if is_base64_unified_file_id:
|
|
litellm_managed_file_ids.append(file_id)
|
|
|
|
if litellm_managed_file_ids:
|
|
# Get all cache keys matching the pattern file_id:*
|
|
for file_id in litellm_managed_file_ids:
|
|
# Search for any cache key starting with this file_id
|
|
unified_file_object = await self.get_unified_file_id(file_id, litellm_parent_otel_span)
|
|
|
|
if unified_file_object:
|
|
file_id_mapping[file_id] = unified_file_object.model_mappings
|
|
|
|
return file_id_mapping
|
|
|
|
async def create_file_for_each_model(
|
|
self,
|
|
llm_router: Optional[Router],
|
|
_create_file_request: CreateFileRequest,
|
|
target_model_names_list: List[str],
|
|
litellm_parent_otel_span: Span,
|
|
) -> List[OpenAIFileObject]:
|
|
if llm_router is None:
|
|
raise Exception("LLM Router not initialized. Ensure models added to proxy.")
|
|
responses = []
|
|
for model in target_model_names_list:
|
|
individual_response = await llm_router.acreate_file(model=model, **_create_file_request)
|
|
responses.append(individual_response)
|
|
|
|
return responses
|
|
|
|
async def acreate_file(
|
|
self,
|
|
create_file_request: CreateFileRequest,
|
|
llm_router: Router,
|
|
target_model_names_list: List[str],
|
|
litellm_parent_otel_span: Span,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> OpenAIFileObject:
|
|
responses = await self.create_file_for_each_model(
|
|
llm_router=llm_router,
|
|
_create_file_request=create_file_request,
|
|
target_model_names_list=target_model_names_list,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
)
|
|
response = await _PROXY_LiteLLMManagedFiles.return_unified_file_id(
|
|
file_objects=responses,
|
|
create_file_request=create_file_request,
|
|
internal_usage_cache=self.internal_usage_cache,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
target_model_names_list=target_model_names_list,
|
|
)
|
|
|
|
## STORE MODEL MAPPINGS IN DB
|
|
model_mappings: Dict[str, str] = {}
|
|
|
|
for file_object in responses:
|
|
model_file_id_mapping = file_object._hidden_params.get("model_file_id_mapping")
|
|
if model_file_id_mapping and isinstance(model_file_id_mapping, dict):
|
|
model_mappings.update(model_file_id_mapping)
|
|
|
|
await self.store_unified_file_id(
|
|
file_id=response.id,
|
|
file_object=response,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
model_mappings=model_mappings,
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
|
|
# Emit Prometheus metrics for managed file creation
|
|
prom_logger = self._get_prometheus_logger()
|
|
if prom_logger:
|
|
first_model = target_model_names_list[0] if target_model_names_list else None
|
|
first_provider = ""
|
|
if responses:
|
|
first_provider = getattr(responses[0], "_hidden_params", {}).get("custom_llm_provider") or ""
|
|
prom_logger.record_managed_file_created(
|
|
model=first_model or "",
|
|
api_provider=first_provider,
|
|
user=user_api_key_dict.user_id or "",
|
|
user_email=getattr(user_api_key_dict, "user_email", None) or "",
|
|
api_key_alias=user_api_key_dict.key_alias or "",
|
|
)
|
|
if response.bytes and response.bytes > 0:
|
|
prom_logger.record_managed_file_size(
|
|
size_bytes=response.bytes,
|
|
purpose=response.purpose or "batch",
|
|
file_type="input",
|
|
model=first_model,
|
|
api_provider=first_provider,
|
|
user=user_api_key_dict.user_id,
|
|
)
|
|
|
|
return response
|
|
|
|
@staticmethod
|
|
async def return_unified_file_id(
|
|
file_objects: List[OpenAIFileObject],
|
|
create_file_request: CreateFileRequest,
|
|
internal_usage_cache: InternalUsageCache,
|
|
litellm_parent_otel_span: Span,
|
|
target_model_names_list: List[str],
|
|
) -> OpenAIFileObject:
|
|
## GET THE FILE TYPE FROM THE CREATE FILE REQUEST
|
|
_, file_type = extract_file_metadata(create_file_request["file"])
|
|
|
|
output_file_id = file_objects[0].id
|
|
model_id = file_objects[0]._hidden_params.get("model_id")
|
|
|
|
unified_file_id = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
|
|
file_type,
|
|
str(uuid.uuid4()),
|
|
",".join(target_model_names_list),
|
|
output_file_id,
|
|
model_id,
|
|
)
|
|
|
|
# Convert to URL-safe base64 and strip padding
|
|
base64_unified_file_id = base64.urlsafe_b64encode(unified_file_id.encode()).decode().rstrip("=")
|
|
|
|
## CREATE RESPONSE OBJECT
|
|
|
|
response = OpenAIFileObject(
|
|
id=base64_unified_file_id,
|
|
object="file",
|
|
purpose=create_file_request["purpose"],
|
|
created_at=file_objects[0].created_at,
|
|
bytes=file_objects[0].bytes,
|
|
filename=file_objects[0].filename,
|
|
status="uploaded",
|
|
expires_at=file_objects[0].expires_at,
|
|
)
|
|
|
|
return response
|
|
|
|
def get_unified_generic_response_id(self, model_id: str, generic_response_id: str) -> str:
|
|
unified_generic_response_id = SpecialEnums.LITELLM_MANAGED_GENERIC_RESPONSE_COMPLETE_STR.value.format(
|
|
model_id, generic_response_id
|
|
)
|
|
return base64.urlsafe_b64encode(unified_generic_response_id.encode()).decode().rstrip("=")
|
|
|
|
def get_unified_batch_id(self, batch_id: str, model_id: str) -> str:
|
|
unified_batch_id = SpecialEnums.LITELLM_MANAGED_BATCH_COMPLETE_STR.value.format(model_id, batch_id)
|
|
return base64.urlsafe_b64encode(unified_batch_id.encode()).decode().rstrip("=")
|
|
|
|
def get_unified_output_file_id(self, output_file_id: str, model_id: str, model_name: Optional[str]) -> str:
|
|
deterministic_uuid: Final = uuid5(uuid5(NAMESPACE_URL, model_id), output_file_id)
|
|
unified_output_file_id = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
|
|
"application/json",
|
|
str(deterministic_uuid),
|
|
model_name or "",
|
|
output_file_id,
|
|
model_id,
|
|
)
|
|
return base64.urlsafe_b64encode(unified_output_file_id.encode()).decode().rstrip("=")
|
|
|
|
def get_model_id_from_unified_file_id(self, file_id: str) -> str:
|
|
return file_id.split("llm_output_file_model_id,")[1].split(";")[0]
|
|
|
|
def get_output_file_id_from_unified_file_id(self, file_id: str) -> str:
|
|
marker = "llm_output_file_id,"
|
|
if marker not in file_id:
|
|
raise ValueError(f"Unified id does not contain {marker!r}: {file_id[:80]!r}")
|
|
return file_id.split(marker, 1)[1].split(";")[0]
|
|
|
|
async def async_post_call_success_hook(
|
|
self, data: Dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes
|
|
) -> LLMResponseTypes:
|
|
if isinstance(response, LiteLLMBatch):
|
|
## Check if unified_file_id is in the response
|
|
unified_file_id = response._hidden_params.get("unified_file_id") # managed file id
|
|
unified_batch_id = response._hidden_params.get("unified_batch_id") # managed batch id
|
|
is_batch_create: Final = response._hidden_params.get(BATCH_CREATE_HIDDEN_PARAM) is True
|
|
model_id = cast(Optional[str], response._hidden_params.get("model_id"))
|
|
model_name = cast(Optional[str], response._hidden_params.get("model_name"))
|
|
|
|
resolved_model_name = resolve_managed_output_file_model_name(
|
|
unified_input_file_id=unified_file_id if isinstance(unified_file_id, str) else response.input_file_id,
|
|
fallback_model_name=model_name,
|
|
)
|
|
original_response_id = response.id
|
|
|
|
if (unified_batch_id or unified_file_id) and model_id:
|
|
response.id = self.get_unified_batch_id(batch_id=response.id, model_id=model_id)
|
|
|
|
# Handle both output_file_id and error_file_id
|
|
for file_attr in ["output_file_id", "error_file_id"]:
|
|
file_id_value: str | None = getattr(response, file_attr, None)
|
|
if file_id_value and model_id:
|
|
decoded_output_file_id = _is_base64_encoded_unified_file_id(file_id_value)
|
|
if decoded_output_file_id and "llm_output_file_id," in decoded_output_file_id:
|
|
provider_file_id = self.get_output_file_id_from_unified_file_id(decoded_output_file_id)
|
|
unified_file_id = file_id_value
|
|
elif decoded_output_file_id:
|
|
verbose_logger.warning(
|
|
f"Skipping {file_attr}={file_id_value!r}: unified id is not a managed file output id"
|
|
)
|
|
continue
|
|
else:
|
|
provider_file_id = file_id_value
|
|
unified_file_id = self.get_unified_output_file_id(
|
|
output_file_id=provider_file_id,
|
|
model_id=model_id,
|
|
model_name=resolved_model_name,
|
|
)
|
|
setattr(response, file_attr, unified_file_id)
|
|
|
|
# Use llm_router credentials when available. Without credentials,
|
|
# Azure and other auth-required providers return 500/401.
|
|
file_object = None
|
|
try:
|
|
# Import module and use getattr for better testability with mocks
|
|
import litellm.proxy.proxy_server as proxy_server_module
|
|
|
|
_llm_router = getattr(proxy_server_module, "llm_router", None)
|
|
if _llm_router is not None and model_id:
|
|
_creds = _llm_router.get_deployment_credentials_with_provider(model_id) or {}
|
|
file_object = await litellm.afile_retrieve(
|
|
file_id=provider_file_id,
|
|
**_creds,
|
|
)
|
|
else:
|
|
file_object = await litellm.afile_retrieve(
|
|
custom_llm_provider=model_name.split("/")[0]
|
|
if model_name and "/" in model_name
|
|
else "openai", # type: ignore[arg-type]
|
|
file_id=provider_file_id,
|
|
)
|
|
verbose_logger.debug(
|
|
f"Successfully retrieved file object for {file_attr}={provider_file_id}"
|
|
)
|
|
except Exception as e:
|
|
verbose_logger.warning(
|
|
f"Failed to retrieve file object for {file_attr}={provider_file_id}: {str(e)}. Storing with None and will fetch on-demand."
|
|
)
|
|
|
|
await self.store_unified_file_id(
|
|
file_id=unified_file_id,
|
|
file_object=file_object,
|
|
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
|
model_mappings={model_id: provider_file_id},
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
request_metadata: Final = data.get("litellm_metadata")
|
|
await self.store_unified_object_id(
|
|
unified_object_id=response.id,
|
|
file_object=response,
|
|
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
|
model_object_id=original_response_id,
|
|
file_purpose="batch",
|
|
user_api_key_dict=user_api_key_dict,
|
|
request_tags=request_tags_from_metadata(request_metadata if isinstance(request_metadata, dict) else {}),
|
|
persist_attribution=is_batch_create,
|
|
create_if_missing=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,
|
|
)
|
|
|
|
if is_batch_create:
|
|
prom_logger = self._get_prometheus_logger()
|
|
if prom_logger:
|
|
batch_provider = ""
|
|
if model_name:
|
|
try:
|
|
from litellm.litellm_core_utils.get_llm_provider_logic import (
|
|
get_llm_provider,
|
|
)
|
|
|
|
_, batch_provider, _, _ = get_llm_provider(model=model_name)
|
|
except Exception:
|
|
if "/" in model_name:
|
|
batch_provider = model_name.split("/")[0]
|
|
prom_logger.record_managed_batch_created(
|
|
model=model_name or "",
|
|
api_provider=batch_provider,
|
|
user=user_api_key_dict.user_id or "",
|
|
user_email=getattr(user_api_key_dict, "user_email", None) or "",
|
|
api_key_alias=user_api_key_dict.key_alias or "",
|
|
)
|
|
|
|
elif isinstance(response, LiteLLMFineTuningJob):
|
|
## Check if unified_file_id is in the response
|
|
unified_file_id = response._hidden_params.get("unified_file_id") # managed file id
|
|
unified_finetuning_job_id = response._hidden_params.get(
|
|
"unified_finetuning_job_id"
|
|
) # managed finetuning job id
|
|
model_id = cast(Optional[str], response._hidden_params.get("model_id"))
|
|
model_name = cast(Optional[str], response._hidden_params.get("model_name"))
|
|
original_response_id = response.id
|
|
if (unified_file_id or unified_finetuning_job_id) and model_id:
|
|
response.id = self.get_unified_generic_response_id(model_id=model_id, generic_response_id=response.id)
|
|
await self.store_unified_object_id(
|
|
unified_object_id=response.id,
|
|
file_object=response,
|
|
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
|
model_object_id=original_response_id,
|
|
file_purpose="fine-tune",
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
elif isinstance(response, AsyncCursorPage):
|
|
"""
|
|
For listing files, filter for the ones created by the user
|
|
"""
|
|
## check if file object
|
|
if hasattr(response, "data") and isinstance(response.data, list):
|
|
if all(isinstance(file_object, FileObject) for file_object in response.data):
|
|
## Get all file id's
|
|
## Check which file id's were created by the user
|
|
## Filter the response to only include the files created by the user
|
|
## Return the filtered response
|
|
file_ids = [
|
|
file_object.id
|
|
for file_object in cast(List[FileObject], response.data) # type: ignore
|
|
]
|
|
user_created_file_ids = await self.get_user_created_file_ids(user_api_key_dict, file_ids)
|
|
## Filter the response to only include the files created by the user
|
|
response.data = user_created_file_ids # type: ignore
|
|
self._scope_list_page_cursors(response, user_created_file_ids)
|
|
return response
|
|
return response
|
|
return response
|
|
|
|
@staticmethod
|
|
def _scope_list_page_cursors(response: AsyncCursorPage, data: List[OpenAIFileObject]) -> None:
|
|
"""Rebuild ``first_id`` / ``last_id`` from the caller-scoped page.
|
|
|
|
The upstream cursors point at rows that were just filtered out, so
|
|
leaving them in place discloses other callers' file ids. ``has_more``
|
|
is always cleared because ``after`` is never forwarded upstream, so
|
|
no further page is reachable through the proxy.
|
|
"""
|
|
if hasattr(response, "first_id"):
|
|
response.first_id = data[0].id if data else None
|
|
if hasattr(response, "last_id"):
|
|
response.last_id = data[-1].id if data else None
|
|
if hasattr(response, "has_more"):
|
|
response.has_more = False
|
|
|
|
async def afile_retrieve(
|
|
self, file_id: str, litellm_parent_otel_span: Optional[Span], llm_router: Optional[Router] = None
|
|
) -> OpenAIFileObject:
|
|
stored_file_object = await self.get_unified_file_id(file_id, litellm_parent_otel_span)
|
|
|
|
# Case 1 : This is not a managed file
|
|
if not stored_file_object:
|
|
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
|
|
|
|
# Case 2: Managed file and the file object exists in the database
|
|
# The stored file_object has the raw provider ID. Replace with the unified ID
|
|
# so callers see a consistent ID (matching Case 3 which does response.id = file_id).
|
|
if stored_file_object and stored_file_object.file_object:
|
|
# Use model_copy to ensure the ID update persists (Pydantic v2 compatibility)
|
|
response = stored_file_object.file_object.model_copy(update={"id": file_id})
|
|
return response
|
|
|
|
# Case 3: Managed file exists in the database but not the file object (for. e.g the batch task might not have run)
|
|
# So we fetch the file object from the provider. We deliberately do not store the result to avoid interfering with batch cost tracking code.
|
|
if not llm_router:
|
|
raise Exception(
|
|
f"LiteLLM Managed File object with id={file_id} has no file_object "
|
|
f"and llm_router is required to fetch from provider"
|
|
)
|
|
|
|
try:
|
|
model_id, model_file_id = next(iter(stored_file_object.model_mappings.items()))
|
|
credentials = llm_router.get_deployment_credentials_with_provider(model_id) or {}
|
|
response = await litellm.afile_retrieve(file_id=model_file_id, **credentials)
|
|
response.id = file_id # Replace with unified ID
|
|
return response
|
|
except Exception as e:
|
|
raise Exception(f"Failed to retrieve file {file_id} from provider: {str(e)}") from e
|
|
|
|
async def afile_list(
|
|
self,
|
|
purpose: Optional[str],
|
|
litellm_parent_otel_span: Optional[Span],
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
limit: Optional[int] = None,
|
|
after: Optional[str] = None,
|
|
**data: Dict,
|
|
) -> FileListPage:
|
|
"""List the managed files the caller owns, newest first.
|
|
|
|
Pagination is keyset based on ``unified_file_id`` so a key that owns
|
|
every file on the proxy still reads one bounded page at a time.
|
|
``purpose`` is applied after parsing, because the managed file table
|
|
keeps it inside the ``file_object`` blob instead of a column, and rows
|
|
whose blob will not parse drop out there too, so a chunk of rows can
|
|
yield fewer matches than the page holds. Successive chunks are read
|
|
until the page is full or the caller's rows run out, which keeps
|
|
``data`` non-empty while matches remain and its last id usable as the
|
|
next cursor. A first chunk that fills the page costs one query; once a
|
|
scan has to continue past it, the chunk widens to
|
|
``FILE_LIST_CONTINUATION_CHUNK_SIZE``, so the walk costs one query per
|
|
that many rows instead of one per page. That bound is per query, not
|
|
per request: the work is still linear in the rows the caller owns, and
|
|
a filter matching nothing reads every one of them, with no index
|
|
covering either the owner filter or the sort.
|
|
"""
|
|
validate_file_list_limit(limit)
|
|
validate_file_list_purpose(purpose)
|
|
|
|
owner_filter: Final = build_owner_filter(user_api_key_dict)
|
|
if owner_filter is None:
|
|
return FileListPage(**build_list_page([]))
|
|
|
|
if after:
|
|
cursor_row = await _managed_file_table(self.prisma_client).find_first(
|
|
where={**owner_filter, "unified_file_id": after}
|
|
)
|
|
if cursor_row is None:
|
|
raise ProxyException(
|
|
message=f"Invalid 'after' cursor: no file found with id '{after}'.",
|
|
type="invalid_request_error",
|
|
param="after",
|
|
code=400,
|
|
openai_code="invalid_value",
|
|
)
|
|
|
|
page_size: Final = min(limit or MAX_FILE_LIST_LIMIT, MAX_FILE_LIST_LIMIT)
|
|
matches: Final[List[OpenAIFileObject]] = []
|
|
cursor_id = after
|
|
chunk_size = page_size + 1
|
|
|
|
while len(matches) <= page_size:
|
|
cursor_args: _CursorPageArgs = {"cursor": {"unified_file_id": cursor_id}, "skip": 1} if cursor_id else {}
|
|
chunk = await _managed_file_table(self.prisma_client).find_many(
|
|
where=owner_filter,
|
|
take=chunk_size,
|
|
order=[{"created_at": "desc"}, {"unified_file_id": "desc"}],
|
|
**cursor_args,
|
|
)
|
|
matches.extend(
|
|
parsed_file_object.model_copy(update={"id": row.unified_file_id})
|
|
for row in chunk
|
|
if (parsed_file_object := _parse_managed_file_object(row.file_object, row.unified_file_id)) is not None
|
|
and (purpose is None or parsed_file_object.purpose == purpose)
|
|
)
|
|
if len(chunk) < chunk_size:
|
|
break
|
|
cursor_id = chunk[-1].unified_file_id
|
|
chunk_size = max(chunk_size, FILE_LIST_CONTINUATION_CHUNK_SIZE)
|
|
|
|
return FileListPage(**build_list_page(matches[:page_size], has_more=len(matches) > page_size))
|
|
|
|
def _is_batch_polling_enabled(self) -> bool:
|
|
"""
|
|
Check if batch cost tracking is actually enabled and running.
|
|
Returns:
|
|
bool: True if batch cost tracking is active, False otherwise
|
|
"""
|
|
try:
|
|
# Import here to avoid circular dependencies
|
|
import litellm.proxy.proxy_server as proxy_server_module
|
|
|
|
# Check if the scheduler has the batch cost checking job registered
|
|
scheduler: Final[_SchedulerWithJobLookup | None] = getattr(proxy_server_module, "scheduler", None)
|
|
if scheduler is None:
|
|
return False
|
|
|
|
# Check if the check_batch_cost_job exists in the scheduler
|
|
try:
|
|
job = scheduler.get_job("check_batch_cost_job")
|
|
if job is not None:
|
|
return True
|
|
except Exception:
|
|
# Job not found or scheduler doesn't support get_job
|
|
pass
|
|
|
|
return False
|
|
except Exception as e:
|
|
verbose_logger.warning(f"Error checking batch polling configuration: {e}. Assuming disabled.")
|
|
return False
|
|
|
|
async def _get_batches_referencing_file(self, file_id: str) -> List[Dict[str, object]]:
|
|
"""
|
|
Find batches that reference this file and still need cost tracking.
|
|
Find batches that are in non-terminal state and have not yet been processed by CheckBatchCost.
|
|
Args:
|
|
file_id: The unified file ID to check
|
|
|
|
Returns:
|
|
List of batch objects referencing this file in non-terminal state
|
|
(max 10 for error message display)
|
|
"""
|
|
# Prepare list of file IDs to check (both unified and provider IDs)
|
|
file_ids_to_check = [file_id]
|
|
|
|
# Get model-specific file IDs for this unified file ID if it's a managed file
|
|
try:
|
|
model_file_id_mapping = await self.get_model_file_id_mapping([file_id], litellm_parent_otel_span=None)
|
|
|
|
if model_file_id_mapping and file_id in model_file_id_mapping:
|
|
# Add all provider file IDs for this unified file
|
|
provider_file_ids = list(model_file_id_mapping[file_id].values())
|
|
file_ids_to_check.extend(provider_file_ids)
|
|
except Exception as e:
|
|
verbose_logger.debug(
|
|
f"Could not get model file ID mapping for {file_id}: {e}. Will only check unified file ID."
|
|
)
|
|
MAX_MATCHES_TO_RETURN = 10
|
|
|
|
batches = await _managed_object_table(self.prisma_client).find_many(
|
|
where={
|
|
"file_purpose": "batch",
|
|
"batch_processed": False,
|
|
"status": {"not_in": ["failed", "expired", "cancelled"]},
|
|
},
|
|
take=MAX_MATCHES_TO_RETURN,
|
|
order={"created_at": "desc"},
|
|
)
|
|
|
|
referencing_batches: Final[list[dict[str, object]]] = []
|
|
for batch in batches:
|
|
try:
|
|
# Parse the batch file_object to check for file references
|
|
decoded_file_object = _decode_json_blob(batch.file_object)
|
|
batch_data: Mapping[str, object] = (
|
|
decoded_file_object if isinstance(decoded_file_object, Mapping) else {}
|
|
)
|
|
|
|
# Extract file IDs from batch
|
|
# Batches typically reference the unified file ID in input_file_id
|
|
# Output and error files are generated by the provider
|
|
input_file_id = batch_data.get("input_file_id")
|
|
output_file_id = batch_data.get("output_file_id")
|
|
error_file_id = batch_data.get("error_file_id")
|
|
|
|
referenced_file_ids = [fid for fid in [input_file_id, output_file_id, error_file_id] if fid]
|
|
|
|
# Check if any referenced file ID matches the file we're trying to delete
|
|
if any(ref_id in file_ids_to_check for ref_id in referenced_file_ids):
|
|
referencing_batches.append(
|
|
{
|
|
"batch_id": batch.unified_object_id,
|
|
"status": batch.status,
|
|
"created_at": batch.created_at,
|
|
}
|
|
)
|
|
except Exception as e:
|
|
verbose_logger.warning(f"Error parsing batch object {batch.unified_object_id}: {e}")
|
|
continue
|
|
|
|
return referencing_batches
|
|
|
|
async def _check_file_deletion_allowed(self, file_id: str) -> None:
|
|
"""
|
|
Check if file deletion should be blocked due to batch references.
|
|
|
|
Blocks deletion if:
|
|
1. File is referenced by any batch in non-terminal state, AND
|
|
2. Batch polling is configured (user wants cost tracking)
|
|
|
|
Args:
|
|
file_id: The unified file ID to check
|
|
|
|
Raises:
|
|
HTTPException: If file deletion should be blocked
|
|
"""
|
|
# Check if batch polling is enabled
|
|
if not self._is_batch_polling_enabled():
|
|
# Batch polling not configured, allow deletion
|
|
return
|
|
|
|
# Check if file is referenced by any non-terminal batches
|
|
referencing_batches = await self._get_batches_referencing_file(file_id)
|
|
|
|
if referencing_batches:
|
|
# File is referenced by non-terminal batches and polling is enabled
|
|
MAX_BATCHES_IN_ERROR = 5 # Limit batches shown in error message for readability
|
|
|
|
# Show up to MAX_BATCHES_IN_ERROR in the error message
|
|
batches_to_show = referencing_batches[:MAX_BATCHES_IN_ERROR]
|
|
batch_statuses = [f"{b['batch_id']}: {b['status']}" for b in batches_to_show]
|
|
|
|
# Determine the count message
|
|
count_message = f"{len(referencing_batches)}"
|
|
if len(referencing_batches) >= 10: # MAX_MATCHES_TO_RETURN from _get_batches_referencing_file
|
|
count_message = "10+"
|
|
|
|
error_message = (
|
|
f"Cannot delete file {file_id}. "
|
|
f"The file is referenced by {count_message} batch(es) in non-terminal state"
|
|
)
|
|
|
|
# Add specific batch details if not too many
|
|
if len(referencing_batches) <= MAX_BATCHES_IN_ERROR:
|
|
error_message += f": {', '.join(batch_statuses)}. "
|
|
else:
|
|
error_message += f" (showing {MAX_BATCHES_IN_ERROR} most recent): {', '.join(batch_statuses)}. "
|
|
|
|
error_message += (
|
|
"To delete this file before complete cost tracking, please delete or cancel the referencing batch(es) first. "
|
|
"Alternatively, wait for all batches to complete and for cost to be computed (batch_processed=true)."
|
|
)
|
|
|
|
# Record blocked deletion metric
|
|
prom_logger = self._get_prometheus_logger()
|
|
if prom_logger:
|
|
prom_logger.record_managed_file_deleted(result="blocked")
|
|
|
|
raise HTTPException(
|
|
status_code=400,
|
|
detail=error_message,
|
|
)
|
|
|
|
async def afile_delete(
|
|
self,
|
|
file_id: str,
|
|
litellm_parent_otel_span: Optional[Span],
|
|
llm_router: Router,
|
|
**data: Dict,
|
|
) -> FileDeleted:
|
|
|
|
# Check if file deletion should be blocked due to batch references
|
|
await self._check_file_deletion_allowed(file_id)
|
|
|
|
# file_id = convert_b64_uid_to_unified_uid(file_id)
|
|
model_file_id_mapping = await self.get_model_file_id_mapping([file_id], litellm_parent_otel_span)
|
|
|
|
specific_model_file_id_mapping = model_file_id_mapping.get(file_id)
|
|
if specific_model_file_id_mapping:
|
|
# Remove conflicting keys from data to avoid duplicate keyword arguments
|
|
filtered_data = {k: v for k, v in data.items() if k not in ("model", "file_id")}
|
|
for model_id, model_file_id in specific_model_file_id_mapping.items():
|
|
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
|
|
delete_data = {
|
|
**{k: v for k, v in filtered_data.items() if k != "_litellm_internal_model_credentials"},
|
|
**(
|
|
{"_litellm_internal_model_credentials": MappingProxyType(dict(credentials))}
|
|
if credentials is not None
|
|
else {}
|
|
),
|
|
}
|
|
await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
|
|
|
|
await self.delete_unified_file_id(file_id, litellm_parent_otel_span)
|
|
|
|
prom_logger = self._get_prometheus_logger()
|
|
if prom_logger:
|
|
prom_logger.record_managed_file_deleted(result="success")
|
|
return FileDeleted(id=file_id, object="file", deleted=True)
|
|
|
|
async def afile_content(
|
|
self,
|
|
file_id: str,
|
|
litellm_parent_otel_span: Optional[Span],
|
|
llm_router: Router,
|
|
**data: Dict,
|
|
) -> "HttpxBinaryResponseContent":
|
|
"""
|
|
Get the content of a file from first model that has it
|
|
"""
|
|
model_file_id_mapping = data.pop("model_file_id_mapping", None)
|
|
model_file_id_mapping = model_file_id_mapping or await self.get_model_file_id_mapping(
|
|
[file_id], litellm_parent_otel_span
|
|
)
|
|
|
|
specific_model_file_id_mapping = model_file_id_mapping.get(file_id)
|
|
|
|
if specific_model_file_id_mapping:
|
|
exception_dict = {}
|
|
for model_id, provider_file_id in specific_model_file_id_mapping.items():
|
|
try:
|
|
# Cloud-storage providers (e.g. Bedrock S3) validate file ids
|
|
# against the deployment's configured bucket, which they only
|
|
# trust from this immutable server-side snapshot, never from
|
|
# request params.
|
|
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
|
|
if credentials is not None:
|
|
data["_litellm_internal_model_credentials"] = cast(Dict, MappingProxyType(dict(credentials)))
|
|
else:
|
|
data.pop("_litellm_internal_model_credentials", None)
|
|
return await llm_router.afile_content(model=model_id, file_id=provider_file_id, **data) # type: ignore
|
|
except Exception as e:
|
|
exception_dict[model_id] = str(e)
|
|
raise Exception(
|
|
f"LiteLLM Managed File object with id={file_id} not found. Checked model id's: {specific_model_file_id_mapping.keys()}. Errors: {exception_dict}"
|
|
)
|
|
else:
|
|
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
|
|
|
|
async def _convert_storage_files_to_base64(
|
|
self,
|
|
messages: List[AllMessageValues],
|
|
file_ids: List[str],
|
|
litellm_parent_otel_span: Optional[Span],
|
|
) -> None:
|
|
"""
|
|
Convert files stored in storage backends to base64 format for Vertex AI/Gemini.
|
|
|
|
This method checks if any managed files are stored in storage backends,
|
|
downloads them, and converts them to base64 format in the messages.
|
|
"""
|
|
# Check each file_id to see if it's stored in a storage backend
|
|
for file_id in file_ids:
|
|
# Check if this is a base64 encoded unified file ID
|
|
decoded_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
|
|
|
|
if not decoded_unified_file_id:
|
|
continue
|
|
|
|
# Check database for storage backend info
|
|
# IMPORTANT: The database stores the base64 encoded unified_file_id (not the decoded version)
|
|
# So we query with the original file_id (which is base64 encoded)
|
|
db_file = await _managed_file_table(self.prisma_client).find_first(where={"unified_file_id": file_id})
|
|
|
|
if not db_file or not db_file.storage_backend or not db_file.storage_url:
|
|
continue
|
|
|
|
# File is stored in a storage backend, download and convert to base64
|
|
try:
|
|
from litellm.llms.base_llm.files.storage_backend_factory import (
|
|
get_storage_backend,
|
|
)
|
|
|
|
storage_backend_name = db_file.storage_backend
|
|
storage_url = db_file.storage_url
|
|
|
|
# Get storage backend (uses same env vars as callback)
|
|
try:
|
|
storage_backend = get_storage_backend(storage_backend_name)
|
|
except ValueError as e:
|
|
verbose_logger.warning(
|
|
f"Storage backend '{storage_backend_name}' error for file {file_id}: {str(e)}"
|
|
)
|
|
continue
|
|
|
|
file_content = await storage_backend.download_file(storage_url)
|
|
|
|
# Determine content type from file object
|
|
content_type = self._get_content_type_from_file_object(db_file.file_object)
|
|
|
|
# Convert to base64
|
|
base64_data = base64.b64encode(file_content).decode("utf-8")
|
|
base64_data_uri = f"data:{content_type};base64,{base64_data}"
|
|
|
|
# Update messages to use base64 instead of file_id
|
|
self._update_messages_with_base64_data(messages, file_id, base64_data_uri, content_type)
|
|
except Exception as e:
|
|
verbose_logger.exception(f"Error converting file {file_id} from storage backend to base64: {str(e)}")
|
|
# Continue with other files even if one fails
|
|
continue
|
|
|
|
def _get_content_type_from_file_object(self, file_object: Optional[Any]) -> str:
|
|
"""
|
|
Determine content type from file object.
|
|
|
|
Uses the MIME type utility for consistent detection and normalization.
|
|
|
|
Args:
|
|
file_object: The file object from the database (can be dict, JSON string, or None)
|
|
|
|
Returns:
|
|
str: MIME type (defaults to "application/octet-stream" if cannot be determined)
|
|
"""
|
|
# Use utility function for detection
|
|
content_type = get_content_type_from_file_object(file_object)
|
|
|
|
# Normalize for Gemini/Vertex AI (requires image/jpeg, not image/jpg)
|
|
content_type = normalize_mime_type_for_provider(content_type, provider="gemini")
|
|
|
|
return content_type
|
|
|
|
def _update_messages_with_base64_data(
|
|
self,
|
|
messages: List[AllMessageValues],
|
|
file_id: str,
|
|
base64_data_uri: str,
|
|
content_type: str,
|
|
) -> None:
|
|
"""
|
|
Update messages to replace file_id with base64 data URI.
|
|
|
|
Args:
|
|
messages: List of messages to update
|
|
file_id: The file ID to replace
|
|
base64_data_uri: The base64 data URI to use as replacement
|
|
content_type: The MIME type of the file (e.g., "image/jpeg", "application/pdf")
|
|
"""
|
|
for message in messages:
|
|
if message.get("role") == "user":
|
|
content = message.get("content")
|
|
if content and isinstance(content, list):
|
|
for element in content:
|
|
if element.get("type") == "file":
|
|
file_element = cast(ChatCompletionFileObject, element)
|
|
file_element_file = file_element.get("file", {})
|
|
|
|
if file_element_file.get("file_id") == file_id:
|
|
# Replace file_id with base64 data
|
|
file_element_file["file_data"] = base64_data_uri
|
|
# Set format to help Gemini determine mime type
|
|
file_element_file["format"] = content_type
|
|
# Remove file_id to ensure only file_data is used
|
|
file_element_file.pop("file_id", None)
|
|
|
|
verbose_logger.debug(
|
|
f"Converted file {file_id} from storage backend to base64 with format {content_type}"
|
|
)
|