mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
* fix(proxy): register bedrock batch output files with a file object so they list Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): guard output file size when content is missing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover managed batch output file listings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): persist metadata for batch output files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover batch output file listing and retrieval end to end Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): preserve provider metadata for managed batch files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): update managed batch output registration mocks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): inject fake prisma client in managed batch output file tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): retry provider file details and refresh fallback batch output entries Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover fallback batch output file details refresh after provider recovers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover model-name managed batch file retrieval Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(managed_files): bound the batch output file lookup so a slow provider cannot stall GET /v1/batches * fix(managed_files): skip the provider lookup for a fallback entry written moments ago * fix(proxy): safely refresh managed batch file details Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep batch listing provider-free Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(bedrock): return metadata for empty S3 objects Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): strengthen batch fallback regression assertions Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(proxy): describe provider reads and DB writes for batch output files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): sanitize caller-derived ids in managed file logs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(proxy): simplify batch and file endpoint docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(ui): regenerate API types from proxy OpenAPI spec Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(test): resolve managed files lint violations Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): save refreshed managed file details through ManagedFileRepository Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(bedrock): classify retrieved file purpose by its own bucket prefix Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(proxy): move batch and file endpoint behavior notes to litellm-docs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(bedrock): resolve strict lint violations in file retrieval Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(bedrock): mark retrieval request metadata handoffs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): resolve managed batch type-check diagnostics Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): avoid duplicate final response bindings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mrinal <mrinal@berri.ai> Co-authored-by: Yucheng He <yucheng@berri.ai>
2446 lines
103 KiB
Python
2446 lines
103 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 asyncio
|
|
import base64
|
|
import json
|
|
import time
|
|
from collections.abc import Awaitable, Callable, Iterator, 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
|
|
|
|
import httpx
|
|
from fastapi import HTTPException
|
|
from pydantic import ValidationError
|
|
from typing_extensions import ReadOnly, Unpack
|
|
|
|
import litellm
|
|
from litellm import Router, verbose_logger
|
|
from litellm._internal_context import with_service_target
|
|
from litellm._uuid import uuid
|
|
from litellm.caching.caching import DualCache
|
|
from litellm.constants import (
|
|
BATCH_OUTPUT_FILE_FALLBACK_RETRY_AFTER_SECONDS,
|
|
BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS,
|
|
MAX_FILE_LIST_LIMIT,
|
|
)
|
|
from litellm.files.types import FileRetrieveCallOptions, FileRetrieveProvider
|
|
from litellm.integrations.custom_logger import CustomLogger
|
|
from litellm.litellm_core_utils.hidden_params import HIDDEN_PARAMS_ATTR
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|
extract_file_metadata,
|
|
)
|
|
from openai import APIConnectionError, AsyncOpenAI
|
|
from openai.types.file_deleted import FileDeleted
|
|
|
|
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
|
|
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.litellm_pre_call_utils import (
|
|
LiteLLMProxyRequestSetup,
|
|
sanitize_for_log,
|
|
)
|
|
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,
|
|
is_litellm_executed_batch,
|
|
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.repositories.managed_file_repository import ManagedFileRepository
|
|
from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccessIssue]
|
|
AllMessageValues,
|
|
AsyncCursorPage,
|
|
ChatCompletionFileObject,
|
|
CreateFileRequest,
|
|
FileListPage,
|
|
FileObject,
|
|
HttpxBinaryResponseContent,
|
|
OpenAIFileObject,
|
|
ResponsesAPIResponse,
|
|
)
|
|
from litellm.types.utils import (
|
|
CallTypesLiteral,
|
|
LiteLLMBatch,
|
|
LiteLLMFineTuningJob,
|
|
LLMResponseTypes,
|
|
SpecialEnums,
|
|
)
|
|
|
|
|
|
class _ManagedFileRetrieve(Protocol):
|
|
async def __call__(
|
|
self,
|
|
*,
|
|
file_id: str,
|
|
_litellm_internal_model_credentials: Mapping[str, object] | None = None,
|
|
**kwargs: Unpack[FileRetrieveCallOptions],
|
|
) -> OpenAIFileObject: ...
|
|
|
|
|
|
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
|
|
|
|
|
|
def _batch_output_file_object(
|
|
unified_file_id: str, raw_file_id: str, size_bytes: int, *, fallback: bool
|
|
) -> OpenAIFileObject:
|
|
filename: Final = raw_file_id.rsplit("/", 1)[-1] or raw_file_id
|
|
return OpenAIFileObject(
|
|
id=unified_file_id,
|
|
object="file",
|
|
purpose="batch_output",
|
|
filename=filename,
|
|
created_at=int(time.time()),
|
|
bytes=size_bytes,
|
|
status="processed",
|
|
litellm_details_fallback=True if fallback else None,
|
|
)
|
|
|
|
|
|
def _public_file_object(file_object: OpenAIFileObject, unified_file_id: str) -> OpenAIFileObject:
|
|
return file_object.model_copy(update={"id": unified_file_id, "litellm_details_fallback": None})
|
|
|
|
|
|
def _is_transient_file_retrieve_error(error: Exception) -> bool:
|
|
status_code: Final = getattr(error, "status_code", None)
|
|
if isinstance(status_code, int):
|
|
return status_code in {408, 429} or status_code >= 500
|
|
return isinstance(error, (httpx.TransportError, APIConnectionError, asyncio.TimeoutError))
|
|
|
|
|
|
def _proxy_llm_router() -> Router | None:
|
|
import litellm.proxy.proxy_server as proxy_server_module
|
|
|
|
return cast(Router | None, getattr(proxy_server_module, "llm_router", None))
|
|
|
|
|
|
def _provider_file_retrieve_credentials(
|
|
*,
|
|
llm_router: Router | None,
|
|
model_id: str | None,
|
|
) -> Mapping[str, object] | None:
|
|
if llm_router is None or model_id is None:
|
|
return None
|
|
try:
|
|
credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id)
|
|
except Exception as error:
|
|
verbose_logger.warning(
|
|
"Failed to retrieve credentials for provider file "
|
|
f"model_id={sanitize_for_log(model_id)}: {sanitize_for_log(error)}"
|
|
)
|
|
return None
|
|
return cast(Mapping[str, object], credentials) if credentials else None
|
|
|
|
|
|
_PROVIDER_FILE_RETRIEVE_PROVIDERS: Final[frozenset[str]] = frozenset(
|
|
{
|
|
"openai",
|
|
"azure",
|
|
"gemini",
|
|
"vertex_ai",
|
|
"bedrock",
|
|
"hosted_vllm",
|
|
"litellm_proxy",
|
|
"manus",
|
|
"anthropic",
|
|
"mistral",
|
|
"xai",
|
|
}
|
|
)
|
|
|
|
|
|
def _model_name_file_retrieve_provider(model_name: str | None) -> FileRetrieveProvider | None:
|
|
if model_name is None:
|
|
return None
|
|
provider, separator, _ = model_name.partition("/")
|
|
if not separator or provider not in _PROVIDER_FILE_RETRIEVE_PROVIDERS:
|
|
return None
|
|
return cast(FileRetrieveProvider, provider)
|
|
|
|
|
|
def _has_provider_file_retrieve_route(
|
|
*,
|
|
router_credentials: Mapping[str, object] | None,
|
|
model_name: str | None,
|
|
) -> bool:
|
|
return bool(router_credentials) or _model_name_file_retrieve_provider(model_name) is not None
|
|
|
|
|
|
class _ManagedFileRow(Protocol):
|
|
unified_file_id: str
|
|
file_object: OpenAIFileObject
|
|
flat_model_file_ids: Sequence[str]
|
|
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 update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
|
|
|
|
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 _ManagedResourceDatabase(Protocol):
|
|
@property
|
|
def litellm_managedfiletable(self) -> _ManagedFileTableActions: ...
|
|
|
|
@property
|
|
def litellm_managedobjecttable(self) -> _ManagedObjectTableActions: ...
|
|
|
|
|
|
class _ManagedResourcePrismaClient(Protocol):
|
|
@property
|
|
def db(self) -> _ManagedResourceDatabase: ...
|
|
|
|
|
|
class _SchedulerWithJobLookup(Protocol):
|
|
def get_job(self, job_id: str) -> object: ...
|
|
|
|
|
|
class _CursorPageArgs(TypedDict, total=False):
|
|
cursor: Mapping[str, str]
|
|
skip: int
|
|
|
|
|
|
class _RouterFileCallKwargs(TypedDict, total=False):
|
|
client: ReadOnly[AsyncOpenAI | None]
|
|
custom_llm_provider: ReadOnly[str | None]
|
|
|
|
|
|
def _managed_file_table(prisma_client: _ManagedResourcePrismaClient) -> _ManagedFileTableActions:
|
|
return prisma_client.db.litellm_managedfiletable
|
|
|
|
|
|
def _iter_provider_file_id_pairs(
|
|
rows: Sequence[_ManagedFileRow],
|
|
requested_provider_file_ids: frozenset[str],
|
|
) -> Iterator[tuple[str, str]]:
|
|
for row in rows:
|
|
for provider_file_id in row.flat_model_file_ids:
|
|
if provider_file_id in requested_provider_file_ids:
|
|
yield provider_file_id, row.unified_file_id
|
|
|
|
|
|
def _managed_object_table(prisma_client: _ManagedResourcePrismaClient) -> _ManagedObjectTableActions:
|
|
return prisma_client.db.litellm_managedobjecttable
|
|
|
|
|
|
def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, str]:
|
|
hidden_params: Final = cast( # cast-ok: _hidden_params is an untyped attribute the upload path sets
|
|
"Mapping[str, object]", getattr(file_object, "_hidden_params", None) or {}
|
|
)
|
|
return MappingProxyType(
|
|
{
|
|
key: value
|
|
for key in ("storage_backend", "storage_url")
|
|
if isinstance(value := hidden_params.get(key), str)
|
|
}
|
|
)
|
|
|
|
|
|
_MANAGED_FILES_TARGET: Final = "managed_files"
|
|
_PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS: Final = (0.5, 1.0, 2.0)
|
|
|
|
|
|
class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|
# Class variables or attributes
|
|
def __init__(
|
|
self,
|
|
internal_usage_cache: InternalUsageCache,
|
|
prisma_client: PrismaClient,
|
|
*,
|
|
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
|
|
):
|
|
self.internal_usage_cache = internal_usage_cache
|
|
self.prisma_client = prisma_client
|
|
self._sleep = sleep
|
|
|
|
@staticmethod
|
|
def _get_prometheus_logger():
|
|
"""Find PrometheusLogger from litellm.callbacks, if registered."""
|
|
from litellm.integrations.prometheus import PrometheusLogger
|
|
|
|
return PrometheusLogger.get_instance()
|
|
|
|
@with_service_target(_MANAGED_FILES_TARGET)
|
|
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")
|
|
storage_metadata: Final = _storage_metadata_of(file_object)
|
|
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,
|
|
storage_backend=storage_metadata.get("storage_backend"),
|
|
storage_url=storage_metadata.get("storage_url"),
|
|
)
|
|
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
|
|
db_data.update(storage_metadata)
|
|
update_data.update(storage_metadata)
|
|
|
|
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
|
|
|
|
@with_service_target(_MANAGED_FILES_TARGET)
|
|
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,
|
|
batch_processed: bool = False,
|
|
) -> 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.
|
|
|
|
batch_processed is set by callers that have already billed the batch
|
|
themselves, so CheckBatchCost skips the row instead of billing it twice.
|
|
It is written only in the upsert 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 = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict) 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,
|
|
"batch_processed": batch_processed,
|
|
},
|
|
"update": update_columns,
|
|
},
|
|
)
|
|
|
|
@with_service_target(_MANAGED_FILES_TARGET)
|
|
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
|
|
|
|
@with_service_target(_MANAGED_FILES_TARGET)
|
|
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 _managed_object_table(self.prisma_client).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 _managed_file_table(self.prisma_client).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 (
|
|
response.output_file_id,
|
|
response.error_file_id,
|
|
)
|
|
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 _managed_object_table(self.prisma_client).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_batch_output_file(
|
|
unified_file_id=file_id,
|
|
provider_file_id=raw_file_id,
|
|
model_id=model_name or None,
|
|
model_name=model_name,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
owner=owner_identity,
|
|
)
|
|
|
|
async def _afile_retrieve_with_retries(
|
|
self,
|
|
*,
|
|
provider_file_id: str,
|
|
call_options: Mapping[str, object],
|
|
internal_model_credentials: Mapping[str, object] | None = None,
|
|
) -> OpenAIFileObject:
|
|
"""Retrieve provider file details with transient retries and SDK retries disabled."""
|
|
retrieve_file: Final = cast(_ManagedFileRetrieve, litellm.afile_retrieve)
|
|
retrieve_options: Final = cast(
|
|
FileRetrieveCallOptions,
|
|
{**call_options, "max_retries": 0},
|
|
)
|
|
for attempt in range(len(_PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS) + 1):
|
|
if attempt > 0:
|
|
await self._sleep(_PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS[attempt - 1])
|
|
try:
|
|
if internal_model_credentials is None:
|
|
return await retrieve_file(
|
|
file_id=provider_file_id,
|
|
**retrieve_options,
|
|
)
|
|
return await retrieve_file(
|
|
file_id=provider_file_id,
|
|
_litellm_internal_model_credentials=MappingProxyType(dict(internal_model_credentials)),
|
|
**retrieve_options,
|
|
)
|
|
except Exception as error:
|
|
if not _is_transient_file_retrieve_error(error) or attempt == len(
|
|
_PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS
|
|
):
|
|
raise
|
|
raise RuntimeError("Provider file retrieve retry loop ended without a result")
|
|
|
|
async def _fetch_provider_file_object(
|
|
self,
|
|
*,
|
|
unified_file_id: str,
|
|
provider_file_id: str,
|
|
model_id: str | None,
|
|
model_name: str | None,
|
|
llm_router: Router | None = None,
|
|
raise_on_failure: bool = False,
|
|
allow_default_provider: bool = False,
|
|
) -> tuple[OpenAIFileObject | None, bool]:
|
|
"""Fetch provider file details through the configured route under a total timeout."""
|
|
route_llm_router: Final = llm_router if llm_router is not None else _proxy_llm_router()
|
|
router_credentials: Final = _provider_file_retrieve_credentials(
|
|
llm_router=route_llm_router,
|
|
model_id=model_id,
|
|
)
|
|
model_name_provider: Final = _model_name_file_retrieve_provider(model_name)
|
|
fetch_route_available: Final = _has_provider_file_retrieve_route(
|
|
router_credentials=router_credentials,
|
|
model_name=model_name,
|
|
)
|
|
default_provider_route: Final = allow_default_provider and route_llm_router is not None
|
|
if not fetch_route_available and not default_provider_route:
|
|
return None, False
|
|
|
|
try:
|
|
if router_credentials is not None:
|
|
provider_file_object: Final = await asyncio.wait_for(
|
|
self._afile_retrieve_with_retries(
|
|
provider_file_id=provider_file_id,
|
|
call_options=router_credentials,
|
|
internal_model_credentials=router_credentials,
|
|
),
|
|
timeout=BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS,
|
|
)
|
|
return provider_file_object.model_copy(update={"id": unified_file_id}), True
|
|
if model_name_provider is not None:
|
|
provider_file_object_by_model_name: Final = await asyncio.wait_for(
|
|
self._afile_retrieve_with_retries(
|
|
provider_file_id=provider_file_id,
|
|
call_options={"custom_llm_provider": model_name_provider},
|
|
),
|
|
timeout=BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS,
|
|
)
|
|
return (
|
|
provider_file_object_by_model_name.model_copy(update={"id": unified_file_id}),
|
|
True,
|
|
)
|
|
if default_provider_route:
|
|
provider_file_object_by_default_route: Final = await asyncio.wait_for(
|
|
self._afile_retrieve_with_retries(
|
|
provider_file_id=provider_file_id,
|
|
call_options=router_credentials or {},
|
|
),
|
|
timeout=BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS,
|
|
)
|
|
return (
|
|
provider_file_object_by_default_route.model_copy(update={"id": unified_file_id}),
|
|
True,
|
|
)
|
|
return None, False
|
|
except Exception as error:
|
|
verbose_logger.warning(
|
|
"Failed to retrieve batch file object for "
|
|
f"provider_file_id={sanitize_for_log(provider_file_id)}: "
|
|
f"{type(error).__name__} {sanitize_for_log(error)}"
|
|
)
|
|
if raise_on_failure:
|
|
if isinstance(error, TimeoutError) and not str(error):
|
|
raise TimeoutError(
|
|
"Provider file retrieve timed out "
|
|
f"after {BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS} seconds"
|
|
) from error
|
|
raise
|
|
return None, True
|
|
|
|
async def _save_refreshed_file_object(
|
|
self,
|
|
stored: LiteLLM_ManagedFileTable,
|
|
file_object: OpenAIFileObject,
|
|
) -> None:
|
|
if not await ManagedFileRepository(self.prisma_client).update_file_object(stored.unified_file_id, file_object):
|
|
return
|
|
refreshed_row: Final = stored.model_copy(update={"file_object": file_object})
|
|
await self.internal_usage_cache.async_set_cache(
|
|
key=stored.unified_file_id,
|
|
value=refreshed_row.model_dump(),
|
|
litellm_parent_otel_span=None,
|
|
)
|
|
|
|
async def store_batch_output_file(
|
|
self,
|
|
*,
|
|
unified_file_id: str,
|
|
provider_file_id: str,
|
|
model_id: str | None,
|
|
model_name: str | None = None,
|
|
owner: UserAPIKeyAuth,
|
|
litellm_parent_otel_span: Span | None,
|
|
size_bytes: int | None = None,
|
|
fetch_provider_details: bool = True,
|
|
) -> None:
|
|
"""Register batch output or error file metadata, optionally fetching provider details."""
|
|
stored_file: Final = await self.get_unified_file_id(unified_file_id, litellm_parent_otel_span)
|
|
stored_object: Final = stored_file.file_object if stored_file is not None else None
|
|
if not fetch_provider_details:
|
|
if stored_file is not None:
|
|
return
|
|
router_credentials: Final = _provider_file_retrieve_credentials(
|
|
llm_router=_proxy_llm_router(),
|
|
model_id=model_id,
|
|
)
|
|
file_object_without_provider_details: Final = _batch_output_file_object(
|
|
unified_file_id,
|
|
provider_file_id,
|
|
size_bytes or 0,
|
|
fallback=_has_provider_file_retrieve_route(
|
|
router_credentials=router_credentials,
|
|
model_name=model_name,
|
|
),
|
|
)
|
|
await self.store_unified_file_id(
|
|
file_id=unified_file_id,
|
|
file_object=file_object_without_provider_details,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
model_mappings={model_id: provider_file_id} if model_id else {},
|
|
user_api_key_dict=owner,
|
|
)
|
|
return
|
|
if stored_object is not None and not stored_object.litellm_details_fallback:
|
|
return
|
|
|
|
fallback_written_recently: Final = (
|
|
stored_object is not None
|
|
and time.time() - stored_object.created_at < BATCH_OUTPUT_FILE_FALLBACK_RETRY_AFTER_SECONDS
|
|
)
|
|
provider_fetch_result: Final = (
|
|
(None, True)
|
|
if fallback_written_recently
|
|
else await self._fetch_provider_file_object(
|
|
unified_file_id=unified_file_id,
|
|
provider_file_id=provider_file_id,
|
|
model_id=model_id,
|
|
model_name=model_name,
|
|
)
|
|
)
|
|
provider_object, fetch_route_available = provider_fetch_result
|
|
if (
|
|
stored_file is not None
|
|
and provider_object is None
|
|
and (size_bytes is None or (stored_object is not None and stored_object.bytes == size_bytes))
|
|
):
|
|
return
|
|
|
|
file_object: Final = (
|
|
provider_object
|
|
if provider_object is not None
|
|
else (
|
|
stored_object.model_copy(update={"bytes": size_bytes})
|
|
if stored_object is not None and size_bytes is not None
|
|
else _batch_output_file_object(
|
|
unified_file_id,
|
|
provider_file_id,
|
|
size_bytes or 0,
|
|
fallback=fetch_route_available,
|
|
)
|
|
)
|
|
)
|
|
|
|
if stored_file is not None:
|
|
await self._save_refreshed_file_object(stored_file, file_object)
|
|
return
|
|
|
|
await self.store_unified_file_id(
|
|
file_id=unified_file_id,
|
|
file_object=file_object,
|
|
litellm_parent_otel_span=litellm_parent_otel_span,
|
|
model_mappings={model_id: provider_file_id} if model_id else {},
|
|
user_api_key_dict=owner,
|
|
)
|
|
|
|
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),
|
|
fetch_provider_details=False,
|
|
)
|
|
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_unified_file_ids_for_provider_file_ids(
|
|
self,
|
|
provider_file_ids: Sequence[str],
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
) -> Mapping[str, str]:
|
|
if not provider_file_ids:
|
|
return MappingProxyType({})
|
|
|
|
unique_provider_file_ids: Final = tuple(dict.fromkeys(provider_file_ids))
|
|
owner_filter: Final = build_owner_filter(user_api_key_dict)
|
|
if owner_filter is None:
|
|
return MappingProxyType({})
|
|
|
|
provider_file_ids_list: Final = [ # mutable-ok: Prisma hasSome requires a list
|
|
provider_file_id for provider_file_id in unique_provider_file_ids
|
|
]
|
|
rows: Final = await _managed_file_table(self.prisma_client).find_many(
|
|
where={ # mutable-ok: Prisma requires a plain dictionary for where
|
|
**owner_filter,
|
|
"flat_model_file_ids": { # mutable-ok: Prisma requires a plain filter dictionary
|
|
"hasSome": provider_file_ids_list,
|
|
},
|
|
}
|
|
)
|
|
return MappingProxyType(
|
|
dict(
|
|
_iter_provider_file_id_pairs(
|
|
rows,
|
|
frozenset(unique_provider_file_ids),
|
|
)
|
|
)
|
|
)
|
|
|
|
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 [
|
|
_public_file_object(parsed_file_object, 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:
|
|
file_hidden_params = cast( # cast-ok: preserve mapping operations on dynamic file metadata
|
|
dict[str, object], getattr(file_object, HIDDEN_PARAMS_ATTR)
|
|
)
|
|
model_file_id_mapping = file_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
|
|
file_hidden_params: Final = cast(dict[str, object], getattr(file_objects[0], HIDDEN_PARAMS_ATTR))
|
|
model_id = file_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):
|
|
decoded_batch_id: Final = _is_base64_encoded_unified_file_id(response.id)
|
|
if decoded_batch_id and is_litellm_executed_batch(decoded_batch_id):
|
|
return response
|
|
## Check if unified_file_id is in the response
|
|
response_hidden_params: Final = cast(dict[str, object], getattr(response, HIDDEN_PARAMS_ATTR))
|
|
unified_file_id = response_hidden_params.get("unified_file_id")
|
|
unified_batch_id = response_hidden_params.get("unified_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)
|
|
|
|
await self.store_batch_output_file(
|
|
unified_file_id=unified_file_id,
|
|
provider_file_id=provider_file_id,
|
|
model_id=model_id,
|
|
model_name=resolved_model_name,
|
|
owner=user_api_key_dict,
|
|
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
|
)
|
|
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
|
|
finetuning_response_hidden_params: Final = cast( # cast-ok: preserve dynamic mapping behavior
|
|
dict[str, object], getattr(response, HIDDEN_PARAMS_ATTR)
|
|
)
|
|
unified_file_id = finetuning_response_hidden_params.get("unified_file_id")
|
|
unified_finetuning_job_id = finetuning_response_hidden_params.get("unified_finetuning_job_id")
|
|
model_id = cast(Optional[str], finetuning_response_hidden_params.get("model_id"))
|
|
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:
|
|
"""Return public details for a managed file ID, refreshing a basic entry when possible."""
|
|
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.file_object is not None:
|
|
file_object: Final = stored_file_object.file_object
|
|
if file_object.litellm_details_fallback and stored_file_object.model_mappings:
|
|
try:
|
|
model_id, provider_file_id = next(iter(stored_file_object.model_mappings.items()))
|
|
refreshed_file_object, _ = await self._fetch_provider_file_object(
|
|
unified_file_id=file_id,
|
|
provider_file_id=provider_file_id,
|
|
model_id=model_id,
|
|
model_name=model_id,
|
|
llm_router=llm_router,
|
|
)
|
|
if refreshed_file_object is None:
|
|
return _public_file_object(file_object, file_id)
|
|
await self._save_refreshed_file_object(stored_file_object, refreshed_file_object)
|
|
return _public_file_object(refreshed_file_object, file_id)
|
|
except Exception as error:
|
|
verbose_logger.warning(
|
|
"Failed to refresh batch file object for "
|
|
f"file_id={sanitize_for_log(file_id)}: {sanitize_for_log(error)}"
|
|
)
|
|
return _public_file_object(file_object, file_id)
|
|
|
|
# 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.
|
|
model_mapping: Final = next(iter(stored_file_object.model_mappings.items()), None)
|
|
if model_mapping is None:
|
|
raise Exception(
|
|
f"LiteLLM Managed File object with id={file_id} has no file_object and no provider route to fetch it"
|
|
)
|
|
|
|
model_id, model_file_id = model_mapping
|
|
try:
|
|
response, fetch_route_available = await self._fetch_provider_file_object(
|
|
unified_file_id=file_id,
|
|
provider_file_id=model_file_id,
|
|
model_id=model_id,
|
|
model_name=model_id,
|
|
llm_router=llm_router,
|
|
raise_on_failure=True,
|
|
allow_default_provider=True,
|
|
)
|
|
except Exception as e:
|
|
raise Exception(f"Failed to retrieve file {file_id} from provider: {str(e)}") from e
|
|
if not fetch_route_available:
|
|
raise Exception(
|
|
f"LiteLLM Managed File object with id={file_id} has no file_object and no provider route to fetch it"
|
|
)
|
|
if response is None:
|
|
raise ValueError("Provider file details could not be retrieved")
|
|
return _public_file_object(response, file_id)
|
|
|
|
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.model_validate(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(
|
|
_public_file_object(parsed_file_object, 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.model_validate(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)
|
|
|
|
managed_file: Final = await self.get_unified_file_id(file_id, litellm_parent_otel_span)
|
|
if managed_file is not None and managed_file.storage_backend and managed_file.storage_url:
|
|
await self._delete_storage_backend_content(managed_file.storage_backend, managed_file.storage_url)
|
|
else:
|
|
await self._delete_provider_files(file_id, litellm_parent_otel_span, llm_router, 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 _delete_storage_backend_content(self, storage_backend_name: str, storage_url: str) -> None:
|
|
try:
|
|
storage_backend: Final = get_storage_backend(storage_backend_name, prisma_client=self.prisma_client)
|
|
except ValueError as e:
|
|
raise HTTPException(status_code=400, detail=f"Cannot delete the stored file content: {e}") from e
|
|
await storage_backend.delete_file(storage_url)
|
|
|
|
async def _delete_provider_files(
|
|
self,
|
|
file_id: str,
|
|
litellm_parent_otel_span: Span | None,
|
|
llm_router: Router,
|
|
data: Mapping[str, object],
|
|
) -> None:
|
|
model_file_id_mapping: Final = await self.get_model_file_id_mapping([file_id], litellm_parent_otel_span)
|
|
specific_model_file_id_mapping: Final = model_file_id_mapping.get(file_id)
|
|
if not specific_model_file_id_mapping:
|
|
return
|
|
filtered_data: Final = {
|
|
k: v for k, v in data.items() if k not in ("model", "file_id", "_litellm_internal_model_credentials")
|
|
}
|
|
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 = {
|
|
**filtered_data,
|
|
**(
|
|
{"_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,
|
|
**cast(_RouterFileCallKwargs, delete_data),
|
|
)
|
|
|
|
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
|
|
"""
|
|
managed_file: Final = await self.get_unified_file_id(file_id, litellm_parent_otel_span)
|
|
if managed_file is not None and managed_file.storage_backend and managed_file.storage_url:
|
|
return await self._storage_backend_content(managed_file.storage_backend, managed_file.storage_url)
|
|
|
|
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 cast(
|
|
HttpxBinaryResponseContent,
|
|
await llm_router.afile_content(
|
|
model=model_id,
|
|
file_id=provider_file_id,
|
|
**cast(_RouterFileCallKwargs, data),
|
|
),
|
|
)
|
|
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 _storage_backend_content(self, storage_backend_name: str, storage_url: str) -> HttpxBinaryResponseContent:
|
|
storage_backend: Final = get_storage_backend(storage_backend_name, prisma_client=self.prisma_client)
|
|
content: Final = await storage_backend.download_file(storage_url)
|
|
return HttpxBinaryResponseContent(response=httpx.Response(status_code=httpx.codes.OK, content=content))
|
|
|
|
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:
|
|
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, prisma_client=self.prisma_client)
|
|
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}"
|
|
)
|
|
PROXY_LiteLLMManagedFiles = _PROXY_LiteLLMManagedFiles
|