Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 28 additions & 4 deletions enterprise/litellm_enterprise/proxy/hooks/managed_files.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
from uuid import NAMESPACE_URL, uuid5

from fastapi import HTTPException
from pydantic import ValidationError

import litellm
from litellm import Router, verbose_logger
Expand Down Expand Up @@ -74,6 +75,26 @@
PrismaClient = Any


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 ValidationError as e:
verbose_logger.warning(
f"Failed to parse managed file object {unified_file_id}: "
f"{e.errors(include_input=False, include_url=False, include_context=False)}"
)
return None
except Exception as e:
verbose_logger.warning(
f"Failed to parse managed file object {unified_file_id}: {type(e).__name__}"
)
Comment thread
greptile-apps[bot] marked this conversation as resolved.
return None


class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Class variables or attributes
def __init__(
Expand Down Expand Up @@ -384,11 +405,14 @@ async def get_user_created_file_ids(
}
)
return [
OpenAIFileObject.model_validate(row.file_object).model_copy(
update={"id": row.unified_file_id}
)
parsed_file_object.model_copy(update={"id": row.unified_file_id})
for row in file_ids
if row.file_object is not None
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(
Expand Down
42 changes: 42 additions & 0 deletions tests/test_litellm/enterprise/proxy/test_managed_files_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
import asyncio
import base64
import json
import logging

import pytest
from typing import Optional
Expand Down Expand Up @@ -189,6 +190,47 @@ async def test_get_user_created_file_ids_remaps_stored_raw_provider_id_to_unifie
assert files[0].purpose == raw_provider_object.purpose


@pytest.mark.asyncio
async def test_parse_managed_file_object_warning_omits_rejected_values(caplog):
from litellm_enterprise.proxy.hooks.managed_files import (
_parse_managed_file_object,
)

with caplog.at_level(logging.WARNING):
parsed = _parse_managed_file_object(
{"id": "file-corrupt", "object": "file", "filename": "confidential.jsonl"},
"unified-corrupt",
)

assert parsed is None
assert "unified-corrupt" in caplog.text
assert "bytes" in caplog.text
assert "confidential.jsonl" not in caplog.text


@pytest.mark.asyncio
async def test_get_user_created_file_ids_skips_unparseable_rows():
managed_files = _make_managed_files_instance()
managed_files.prisma_client.db.litellm_managedfiletable.find_many = AsyncMock(
return_value=[
MagicMock(
file_object={"id": "file-corrupt", "object": "file"},
unified_file_id="unified-corrupt",
),
MagicMock(
file_object=_make_file_object().model_dump(),
unified_file_id="unified-valid",
),
]
)

files = await managed_files.get_user_created_file_ids(
_make_user_api_key_dict(), ["file-output-abc"]
)

assert [file.id for file in files] == ["unified-valid"]


@pytest.mark.asyncio
async def test_should_fallback_when_no_router():
"""
Expand Down
Loading