From 8efb4a21f6ebb6a2c4f71e0f422ff9dbc9318972 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 21:28:50 -0700 Subject: [PATCH] fix(vector_stores): return managed file ids from vector store file list (#43800) * fix(vector_stores): return managed file ids from vector store file list Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(vector_stores): cover managed file list route Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(vector_stores): only map round-trippable managed ids and index flat file ids Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy-extras): build managed file gin index concurrently Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(proxy-extras): move the managed file gin index migration after main's newest * fix(vector_stores): satisfy lint gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(proxy): drop the stale no-index note on the raw file-id guard * test(vector_stores): cover managed file ids on the vector store file list end to end Integration cells for GET /v1/vector_stores/{vs}/files mapping provider file ids back to the caller's owner-scoped managed ids and decoding managed after and before cursors: raw httpx, the OpenAI SDK sync and async pagers, the three credential routing modes, the owner filter branches, raw and unmappable cursors, provider errors, duplicate and non-string ids, a provider outage mid-burst, a worker SIGKILL mid-burst, and the GIN index migration applied by the migration entrypoint and by db push --------- Co-authored-by: shivam Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../proxy/hooks/managed_files.py | 47 +- .../migration.sql | 12 + .../litellm_proxy_extras/schema.prisma | 1 + .../openai_files_endpoints/common_utils.py | 11 +- .../managed_id_rewriter.py | 7 +- litellm/proxy/schema.prisma | 1 + .../vector_store_files_endpoints/endpoints.py | 116 ++- schema.prisma | 1 + .../test_managed_file_flat_ids_index.py | 184 +++++ ...test_vector_store_file_list_managed_ids.py | 753 ++++++++++++++++++ .../proxy/test_managed_files_hook.py | 87 +- .../test_vector_store_tenant_guard.py | 70 +- .../test_endpoints.py | 157 +++- 13 files changed, 1423 insertions(+), 24 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20261003000000_add_managed_file_flat_ids_gin_index/migration.sql create mode 100644 tests/integration/database/test_managed_file_flat_ids_index.py create mode 100644 tests/integration/management/test_vector_store_file_list_managed_ids.py diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 21bf7abdc2e..7129573c30c 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -3,7 +3,7 @@ import base64 import json -from collections.abc import Mapping, Sequence +from collections.abc import Iterator, Mapping, Sequence from types import MappingProxyType from typing import ( TYPE_CHECKING, @@ -54,6 +54,7 @@ from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.openai_files_endpoints.common_utils import ( BATCH_CREATE_HIDDEN_PARAM, FILE_LIST_CONTINUATION_CHUNK_SIZE, + ManagedFileIdResolver, _is_base64_encoded_unified_file_id, apply_unified_file_ids, decode_model_from_file_id, @@ -144,6 +145,7 @@ def _parse_managed_file_object(raw_file_object: object, unified_file_id: str) -> 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] @@ -201,6 +203,16 @@ def _managed_file_table(prisma_client: PrismaClient) -> _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: PrismaClient) -> _ManagedObjectTableActions: return prisma_client.db.litellm_managedobjecttable @@ -710,6 +722,39 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): 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]: diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261003000000_add_managed_file_flat_ids_gin_index/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261003000000_add_managed_file_flat_ids_gin_index/migration.sql new file mode 100644 index 00000000000..b222cc57dab --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261003000000_add_managed_file_flat_ids_gin_index/migration.sql @@ -0,0 +1,12 @@ +-- CreateIndex (CONCURRENTLY) +-- +-- Disclaimer: +-- - CREATE INDEX CONCURRENTLY cannot run inside a transaction. This migration must stay a +-- single statement so Prisma Migrate on PostgreSQL can apply it outside a transaction. +-- - Builds are slower and use more I/O than a blocking CREATE INDEX; if the build is +-- interrupted, Postgres may leave an INVALID index that must be dropped and recreated. +-- - Do not edit this file after it has been applied to any database: Prisma checksums +-- migrations; add a new migration instead. +-- - Requires PostgreSQL that supports CONCURRENTLY with IF NOT EXISTS (use a new migration +-- without IF NOT EXISTS if you must support older versions). +CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_ManagedFileTable_flat_model_file_ids_idx" ON "LiteLLM_ManagedFileTable" USING GIN ("flat_model_file_ids"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index e75a9b71c7c..cf76b764350 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1144,6 +1144,7 @@ model LiteLLM_ManagedFileTable { updated_by String? @@index([unified_file_id]) + @@index([flat_model_file_ids], type: Gin) @@index([team_id, created_at(sort: Desc)]) } diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 40478e75d7c..4840b28cb39 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -1,7 +1,7 @@ import base64 import mimetypes import re -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from dataclasses import dataclass, field from types import MappingProxyType from typing import ( @@ -95,6 +95,15 @@ class ManagedResourceAccessChecker(Protocol): ) -> bool: ... +@runtime_checkable +class ManagedFileIdResolver(Protocol): + async def get_unified_file_ids_for_provider_file_ids( + self, + provider_file_ids: Sequence[str], + user_api_key_dict: "UserAPIKeyAuth", + ) -> Mapping[str, str]: ... + + def _is_base64_encoded_unified_file_id(b64_uid: str) -> str | Literal[False]: # Ensure b64_uid is a string and not a mock object if not isinstance(b64_uid, str): diff --git a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py index 8e1dba928af..94a75a9802e 100644 --- a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py +++ b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py @@ -188,10 +188,9 @@ _OBJECT_PREFIXES: Final[frozenset[str]] = frozenset({"batch_", "resp_"}) _MAX_BODY_REWRITE_DEPTH: Final = 64 # Caps the distinct raw-provider-id guard lookups issued per request. A raw -# file-id guard is an unindexed array-containment scan over -# LiteLLM_ManagedFileTable (flat_model_file_ids has no index), so a body packed -# with id-shaped strings could otherwise amplify one request into thousands of -# full-table scans. Legitimate callers reference managed IDs (resolved via an +# file-id guard is an array-containment lookup over LiteLLM_ManagedFileTable, +# so a body packed with id-shaped strings could otherwise amplify one request +# into thousands of lookups. Legitimate callers reference managed IDs (resolved via an # indexed lookup, never the guard), so guarding more raw ids than this only # happens under abuse; the request is rejected rather than skipping the guard. _MAX_RAW_ID_GUARD_LOOKUPS: Final = 100 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index e75a9b71c7c..cf76b764350 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1144,6 +1144,7 @@ model LiteLLM_ManagedFileTable { updated_by String? @@index([unified_file_id]) + @@index([flat_model_file_ids], type: Gin) @@index([team_id, created_at(sort: Desc)]) } diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py index 8a8e43abd79..e44aa022334 100644 --- a/litellm/proxy/vector_store_files_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py @@ -1,9 +1,13 @@ -from typing import TYPE_CHECKING, Final, Optional +import re +from collections.abc import Mapping +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Optional, cast from fastapi import APIRouter, Depends, Request, Response from fastapi.responses import ORJSONResponse import litellm +from litellm.llms.base_llm.managed_resources.utils import is_base64_encoded_unified_id from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing @@ -13,6 +17,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_query, ) from litellm.proxy.openai_files_endpoints.common_utils import ( + ManagedFileIdResolver, authorize_model_for_key, get_credentials_for_model, handle_model_based_routing, @@ -24,6 +29,10 @@ from litellm.proxy.vector_store_endpoints.utils import ( is_allowed_to_call_vector_store_files_endpoint, ) from litellm.types.utils import LlmProviders +from litellm.types.vector_store_files import ( + VectorStoreFileListResponse, + VectorStoreFileObject, +) from litellm.types.vector_stores import LiteLLM_ManagedVectorStore if TYPE_CHECKING: @@ -32,6 +41,93 @@ if TYPE_CHECKING: router: Final = APIRouter() +def _provider_file_id_from_managed_id(managed_file_id: str | None) -> str | None: + if managed_file_id is None: + return None + + decoded_id: Final = is_base64_encoded_unified_id(managed_file_id) + if not decoded_id: + return managed_file_id + + match: Final = re.search(r"(?:^|;)llm_output_file_id,([^;]+)", decoded_id) + return match.group(1).strip() if match else managed_file_id + + +def _with_provider_file_id_cursors( + query_params: Mapping[str, str], +) -> Mapping[str, str | None]: + return MappingProxyType( + { + key: (_provider_file_id_from_managed_id(value) if key in {"after", "before"} else value) + for key, value in query_params.items() + } + ) + + +def _managed_file_id_or_original( + file_id: str | None, + id_map: Mapping[str, str], +) -> str | None: + return id_map.get(file_id, file_id) if file_id is not None else None + + +def _with_managed_file_id( + file: VectorStoreFileObject, + id_map: Mapping[str, str], +) -> VectorStoreFileObject: + file_id: Final = file.get("id") + if not isinstance(file_id, str) or file_id not in id_map: + return file + managed_file: Final[VectorStoreFileObject] = {**file, "id": id_map[file_id]} + return managed_file + + +def _with_managed_file_ids( + response: VectorStoreFileListResponse, + id_map: Mapping[str, str], +) -> VectorStoreFileListResponse: + data: Final = response.get("data") + if not data: + return response + + first_id: Final = response.get("first_id") + last_id: Final = response.get("last_id") + mapped_data: Final = [_with_managed_file_id(file, id_map) for file in data] + mapped_response: Final[VectorStoreFileListResponse] = { + **response, + "data": mapped_data, + "first_id": _managed_file_id_or_original(first_id, id_map), + "last_id": _managed_file_id_or_original(last_id, id_map), + } + return mapped_response + + +async def _with_managed_file_list_ids( + response: VectorStoreFileListResponse, + managed_files_obj: object | None, + user_api_key_dict: UserAPIKeyAuth, +) -> VectorStoreFileListResponse: + data: Final = response.get("data") + if not data or not isinstance(managed_files_obj, ManagedFileIdResolver): + return response + + provider_file_ids: Final = tuple( + dict.fromkeys(provider_file_id for file in data if isinstance(provider_file_id := file.get("id"), str)) + ) + id_map: Final = await managed_files_obj.get_unified_file_ids_for_provider_file_ids( + provider_file_ids=provider_file_ids, + user_api_key_dict=user_api_key_dict, + ) + round_trippable_id_map: Final = MappingProxyType( + { + provider_file_id: managed_file_id + for provider_file_id, managed_file_id in id_map.items() + if _provider_file_id_from_managed_id(managed_file_id) == provider_file_id + } + ) + return _with_managed_file_ids(response, round_trippable_id_map) + + async def _update_request_data_with_managed_file_id( data: dict, file_id: str, @@ -62,11 +158,8 @@ async def _update_request_data_with_managed_file_id( Tuple of (updated request data, original_managed_file_id) - original_managed_file_id is the original file_id if it was managed/encoded, None otherwise """ - import re - from litellm import verbose_logger from litellm.llms.base_llm.managed_resources.utils import ( - is_base64_encoded_unified_id, parse_unified_id, ) from litellm.proxy.openai_files_endpoints.common_utils import ( @@ -591,7 +684,7 @@ async def vector_store_file_list( version, ) - query_params: Final = dict(request.query_params) + query_params: Final = _with_provider_file_id_cursors(request.query_params) data: dict[str, str | None] = {"vector_store_id": vector_store_id} data.update(query_params) data["vector_store_id"] = vector_store_id @@ -628,7 +721,7 @@ async def vector_store_file_list( processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + response: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -646,6 +739,17 @@ async def vector_store_file_list( user_api_base=user_api_base, version=version, ) + if not isinstance(response, dict): + return response + managed_files_obj: Final[object | None] = proxy_logging_obj.get_proxy_hook("managed_files") + return await _with_managed_file_list_ids( + response=cast( # cast-ok: [LIT006] this route returns the provider's file-list response shape + VectorStoreFileListResponse, + response, + ), + managed_files_obj=managed_files_obj, + user_api_key_dict=user_api_key_dict, + ) except Exception as e: # noqa: BLE001 raise await processor._handle_llm_api_exception( e=e, diff --git a/schema.prisma b/schema.prisma index e75a9b71c7c..cf76b764350 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1144,6 +1144,7 @@ model LiteLLM_ManagedFileTable { updated_by String? @@index([unified_file_id]) + @@index([flat_model_file_ids], type: Gin) @@index([team_id, created_at(sort: Desc)]) } diff --git a/tests/integration/database/test_managed_file_flat_ids_index.py b/tests/integration/database/test_managed_file_flat_ids_index.py new file mode 100644 index 00000000000..1d1708f0b90 --- /dev/null +++ b/tests/integration/database/test_managed_file_flat_ids_index.py @@ -0,0 +1,184 @@ +import os +import re +import shutil +import subprocess +import sys +import uuid +from collections.abc import Callable +from pathlib import Path +from typing import Final +from urllib.parse import urlsplit + +import pytest +from integration._support.client import Gateway, object_value, string_value +from integration._support.database import read_rows, scratch_database +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +REPO_ROOT: Final = Path(__file__).resolve().parents[3] +PRISMA_DIR: Final = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras" +GIN_MIGRATION: Final = "20261003000000_add_managed_file_flat_ids_gin_index" +INDEX_NAME: Final = "LiteLLM_ManagedFileTable_flat_model_file_ids_idx" +SHIPPED_MIGRATIONS: Final = tuple(sorted(path.name for path in (PRISMA_DIR / "migrations").iterdir() if path.is_dir())) +INDEX_ROW: Final = ( + "SELECT i.indexdef, x.indisvalid FROM pg_indexes i " + "JOIN pg_class c ON c.relname = i.indexname JOIN pg_index x ON x.indexrelid = c.oid WHERE i.indexname = %s" +) +APPLIED_MIGRATIONS: Final = ( + 'SELECT migration_name FROM "_prisma_migrations" ' + "WHERE finished_at IS NOT NULL AND rolled_back_at IS NULL AND migration_name <> %s ORDER BY migration_name" +) +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +UPLOAD_FILENAME: Final = re.compile(rb'filename="([^"]+)"') + + +def _index_rows(database_url: str | None = None) -> list[dict[str, JsonValue]]: + return read_rows(INDEX_ROW, (INDEX_NAME,), database_url=database_url) + + +def _assert_valid_gin_index(rows: list[dict[str, JsonValue]]) -> None: + assert len(rows) == 1, rows + definition: Final = string_value(rows[0]["indexdef"]) + assert "USING gin" in definition, definition + assert '"LiteLLM_ManagedFileTable"' in definition, definition + assert "flat_model_file_ids" in definition, definition + assert rows[0]["indisvalid"] is True, rows + + +def _applied_migrations(database_url: str) -> tuple[str, ...]: + rows: Final = read_rows(APPLIED_MIGRATIONS, ("",), database_url=database_url) + return tuple(string_value(row["migration_name"]) for row in rows) + + +def _leg_python_path() -> str: + return os.pathsep.join( + ( + str(REPO_ROOT), + str(REPO_ROOT / "litellm-proxy-extras"), + str(REPO_ROOT / "enterprise"), + os.environ.get("PYTHONPATH", ""), + ) + ) + + +def _deploy_schema_before(database_url: str, directory: Path, migration: str) -> None: + older: Final = directory / "older-release" + (older / "migrations").mkdir(parents=True) + shutil.copy(PRISMA_DIR / "schema.prisma", older / "schema.prisma") + shutil.copy(PRISMA_DIR / "migrations" / "migration_lock.toml", older / "migrations" / "migration_lock.toml") + for name in (name for name in SHIPPED_MIGRATIONS if name < migration): + shutil.copytree(PRISMA_DIR / "migrations" / name, older / "migrations" / name) + subprocess.run( + [sys.executable, "-I", "-m", "prisma", "migrate", "deploy", "--schema", str(older / "schema.prisma")], + check=True, + capture_output=True, + text=True, + timeout=600, + env={**os.environ, "DATABASE_URL": database_url}, + ) + + +def _run_migration_entrypoint(database_url: str) -> subprocess.CompletedProcess[str]: + return subprocess.run( + [sys.executable, "-m", "litellm.proxy.prisma_migration"], + capture_output=True, + text=True, + timeout=600, + cwd=REPO_ROOT, + env={**os.environ, "DATABASE_URL": database_url, "PYTHONPATH": _leg_python_path()}, + ) + + +def _provider(store: str, provider_file_id: str) -> Callable[[Request], Reply]: + page: Final[dict[str, JsonValue]] = { + "object": "list", + "data": [ + {"id": provider_file_id, "object": "vector_store.file", "vector_store_id": store, "status": "completed"} + ], + "first_id": provider_file_id, + "last_id": provider_file_id, + "has_more": False, + } + file_object: Final[dict[str, JsonValue]] = { + "id": provider_file_id, + "object": "file", + "bytes": 6, + "created_at": 1700000000, + "filename": "a.txt", + "purpose": "user_data", + "status": "processed", + } + + def respond(request: Request) -> Reply: + path: Final = urlsplit(request.target).path + if request.method == "POST" and path == "/v1/files" and UPLOAD_FILENAME.search(request.body): + return Reply(body=JSON_OBJECT.dump_json(file_object)) + if request.method == "GET" and path == f"/v1/vector_stores/{store}/files": + return Reply(body=JSON_OBJECT.dump_json(page)) + return Reply(status=404, body=b'{"error": {"message": "unscripted"}}') + + return respond + + +def _listed_ids(gateway: Gateway, store: str, model: str) -> tuple[JsonValue, ...]: + listed: Final = gateway.request("GET", f"/v1/vector_stores/{store}/files", params={"model": model}) + assert listed.status_code == 200, listed.text + page: Final = JSON_OBJECT.validate_json(listed.content) + data: Final = page["data"] + assert isinstance(data, list), listed.text + ids: Final = tuple(object_value(entry)["id"] for entry in data) + assert (page["first_id"], page["last_id"]) == (ids[0], ids[-1]), listed.text + return ids + + +@pytest.mark.timeout(900) +def test_migration_entrypoint_adds_the_gin_index_and_the_upgraded_proxy_maps_managed_ids( + gateway: Gateway, tmp_path: Path +) -> None: + with scratch_database() as database_url: + _deploy_schema_before(database_url, tmp_path, GIN_MIGRATION) + assert _index_rows(database_url) == [] + assert _applied_migrations(database_url) == tuple(name for name in SHIPPED_MIGRATIONS if name < GIN_MIGRATION) + entrypoint: Final = _run_migration_entrypoint(database_url) + assert entrypoint.returncode == 0, entrypoint.stdout + entrypoint.stderr + _assert_valid_gin_index(_index_rows(database_url)) + assert GIN_MIGRATION in _applied_migrations(database_url), entrypoint.stdout + store: Final = "vs_" + uuid.uuid4().hex + provider_file_id: Final = "file-" + uuid.uuid4().hex[:16] + upgraded_environment: Final = {"DATABASE_URL": database_url, "DISABLE_SCHEMA_UPDATE": "true"} + with ( + wire_server(_provider(store, provider_file_id)) as wire, + owned_proxy(gateway, tmp_path, upgraded_environment) as upgraded, + ): + model: Final = f"integration-{uuid.uuid4().hex}" + upgraded.post( + "/model/new", + { + "model_name": model, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "upgraded-provider-key", + "api_base": wire.url + "/v1", + }, + "model_info": {}, + }, + ) + uploaded: Final = upgraded.request_multipart( + "/v1/files", + {"purpose": "user_data", "target_model_names": model}, + {"file": ("a.txt", b"notes\n", "text/plain")}, + ) + assert uploaded.status_code == 200, uploaded.text + managed: Final = string_value(JSON_OBJECT.validate_json(uploaded.content)["id"]) + assert read_rows( + 'SELECT flat_model_file_ids FROM "LiteLLM_ManagedFileTable" WHERE unified_file_id = %s', + (managed,), + database_url=database_url, + ) == [{"flat_model_file_ids": [provider_file_id]}] + assert _listed_ids(upgraded, store, model) == (managed,) + + +def test_db_push_creates_a_valid_gin_index_on_the_flat_provider_file_ids(gateway: Gateway) -> None: + assert gateway.request("GET", "/health/liveliness").status_code == 200 + _assert_valid_gin_index(_index_rows()) diff --git a/tests/integration/management/test_vector_store_file_list_managed_ids.py b/tests/integration/management/test_vector_store_file_list_managed_ids.py new file mode 100644 index 00000000000..7ad68e884e6 --- /dev/null +++ b/tests/integration/management/test_vector_store_file_list_managed_ids.py @@ -0,0 +1,753 @@ +import base64 +import hashlib +import json +import re +import signal +import threading +import uuid +from collections.abc import Callable, Generator, Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from dataclasses import dataclass +from functools import partial +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final +from urllib.parse import parse_qs, urlsplit + +import httpx +import psutil +import pytest +from integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from openai import AsyncOpenAI, OpenAI +from pydantic import JsonValue, TypeAdapter + +MANAGED_PREFIX: Final = "litellm_proxy:" +CARRIED_PROVIDER_FILE_ID: Final = re.compile(r"(?:^|;)llm_output_file_id,([^;]+)") +UPLOAD_FILENAME: Final = re.compile(rb'filename="([^"]+)"') +FILE_PATH: Final = re.compile(r"^/v1/files/([^/]+)$") +STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +MANAGED_FILE_ROW: Final = ( + 'SELECT flat_model_file_ids, created_by, team_id FROM "LiteLLM_ManagedFileTable" WHERE unified_file_id = %s' +) + +Listing = Callable[[Request], Reply] + + +def _provider_file_id(bearer: str, filename: str) -> str: + return "file-" + hashlib.sha256(f"{bearer}:{filename}".encode()).hexdigest()[:16] + + +def _bearer(request: Request) -> str: + return request.headers.get("authorization", "").removeprefix("Bearer ") + + +def _query(request: Request) -> dict[str, list[str]]: + return parse_qs(urlsplit(request.target).query, keep_blank_values=True) + + +def _json(response: httpx.Response) -> dict[str, JsonValue]: + return JSON_OBJECT.validate_json(response.content) + + +def _json_reply(body: Mapping[str, JsonValue], status: int = 200) -> Reply: + return Reply(status=status, body=json.dumps(body).encode()) + + +def _file_object(file_id: str) -> dict[str, JsonValue]: + return { + "id": file_id, + "object": "file", + "bytes": 12, + "created_at": 1700000000, + "filename": "notes.txt", + "purpose": "user_data", + "status": "processed", + } + + +def _store_file(store: str, file_id: JsonValue) -> dict[str, JsonValue]: + return { + "id": file_id, + "object": "vector_store.file", + "usage_bytes": 123, + "created_at": 1700000001, + "vector_store_id": store, + "status": "completed", + "last_error": None, + "chunking_strategy": {"type": "static", "static": {"max_chunk_size_tokens": 800, "chunk_overlap_tokens": 400}}, + "attributes": {}, + } + + +def _page(store: str, file_ids: tuple[JsonValue, ...], *, has_more: bool = False) -> dict[str, JsonValue]: + return { + "object": "list", + "data": [_store_file(store, file_id) for file_id in file_ids], + "first_id": file_ids[0] if file_ids else None, + "last_id": file_ids[-1] if file_ids else None, + "has_more": has_more, + } + + +def _constant_listing(store: str, *file_ids: JsonValue) -> Listing: + return lambda _: _json_reply(_page(store, file_ids)) + + +def _paged_listing(store: str, first: str, second: str) -> Listing: + def listing(request: Request) -> Reply: + if _query(request).get("after") == [first]: + return _json_reply(_page(store, (second,))) + return _json_reply(_page(store, (first,), has_more=True)) + + return listing + + +def _provider_error(status: int, message: str) -> dict[str, JsonValue]: + return {"error": {"message": message, "type": "provider_error", "code": str(status)}} + + +def _error_listing(status: int, message: str) -> Callable[[str, str], Listing]: + return lambda _store, _bearer: lambda _: _json_reply(_provider_error(status, message), status) + + +def _html_listing() -> Callable[[str, str], Listing]: + return lambda _store, _bearer: lambda _: Reply(body=b"upstream maintenance", content_type="text/html") + + +def _two_pages(store: str, bearer: str) -> Listing: + return _paged_listing(store, _provider_file_id(bearer, "a.txt"), _provider_file_id(bearer, "b.txt")) + + +def _raw_then_uploaded(raw_id: str) -> Callable[[str, str], Listing]: + return lambda store, bearer: _constant_listing(store, raw_id, _provider_file_id(bearer, "a.txt")) + + +def _uploaded_then_integer(store: str, bearer: str) -> Listing: + return _constant_listing(store, _provider_file_id(bearer, "a.txt"), 7) + + +def _provider(store: str, listing: Listing) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + path: Final = urlsplit(request.target).path + if request.method == "POST" and path == "/v1/files": + filename: Final = UPLOAD_FILENAME.search(request.body) + assert filename is not None, request.body[:200] + return _json_reply(_file_object(_provider_file_id(_bearer(request), filename.group(1).decode()))) + if request.method == "POST" and path == f"/v1/vector_stores/{store}/files": + return _json_reply(_store_file(store, JSON_OBJECT.validate_json(request.body)["file_id"])) + if request.method == "GET" and path == f"/v1/vector_stores/{store}/files": + return listing(request) + file: Final = FILE_PATH.match(path) + if request.method == "GET" and file: + return _json_reply(_file_object(file.group(1))) + if request.method == "DELETE" and file: + return _json_reply({"id": file.group(1), "object": "file", "deleted": True}) + return _json_reply({"error": {"message": f"unscripted {request.method} {request.target}"}}, 404) + + return respond + + +def _decoded(managed_file_id: str) -> str: + decoded: Final = base64.urlsafe_b64decode(managed_file_id + "=" * (-len(managed_file_id) % 4)).decode() + assert decoded.startswith(MANAGED_PREFIX), decoded + return decoded + + +def _carried_provider_file_id(managed_file_id: str) -> str: + carried: Final = CARRIED_PROVIDER_FILE_ID.search(_decoded(managed_file_id)) + assert carried is not None, managed_file_id + return carried.group(1) + + +def _upload(gateway: Gateway, key: str, target_model_names: str, filename: str) -> str: + uploaded: Final = gateway.request_multipart( + "/v1/files", + {"purpose": "user_data", "target_model_names": target_model_names}, + {"file": (filename, f"notes in {filename}\n".encode(), "text/plain")}, + key=key, + ) + assert uploaded.status_code == 200, uploaded.text + return string_value(_json(uploaded)["id"]) + + +def _listed(response: httpx.Response) -> dict[str, JsonValue]: + assert response.status_code == 200, response.text + return _json(response) + + +def _ids(page: Mapping[str, JsonValue]) -> tuple[JsonValue, ...]: + data: Final = page["data"] + assert isinstance(data, list), page + return tuple(object_value(entry)["id"] for entry in data) + + +def _sdk_base_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + "/v1" + + +def _models_over_a_fresh_connection(gateway: Gateway, _: int) -> frozenset[str]: + with httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False) as client: + listed: Final = client.get("/v1/models", headers={"Authorization": f"Bearer {gateway.key}"}) + assert listed.status_code == 200, listed.text + data: Final = _json(listed)["data"] + assert isinstance(data, list), listed.text + return frozenset(string_value(object_value(entry)["id"]) for entry in data) + + +def _every_worker_serves(gateway: Gateway, model: str) -> bool: + with ThreadPoolExecutor(max_workers=16) as pool: + rounds: Final = tuple( + tuple(pool.map(partial(_models_over_a_fresh_connection, gateway), range(16))) for _ in range(2) + ) + return all(model in seen for round_ in rounds for seen in round_) + + +def _wait_until_every_worker_serves(gateway: Gateway, model: str) -> None: + eventually(lambda: _every_worker_serves(gateway, model), lambda served: served, seconds=90) + + +@dataclass(frozen=True, slots=True) +class _Member: + team: str + user: str + key: str + + +def _member(scenario: Scenario, *models: str) -> _Member: + team: Final = scenario.team(models=list(models)) + user: Final = scenario.member(team) + return _Member(team, user, scenario.key(team_id=team, user_id=user)) + + +@dataclass(frozen=True, slots=True) +class _Rig: + gateway: Gateway + scenario: Scenario + wire: Wire + store: str + bearer: str + model: str + + def file_id(self, filename: str) -> str: + return _provider_file_id(self.bearer, filename) + + def upload(self, key: str, filename: str) -> str: + managed: Final = _upload(self.gateway, key, self.model, filename) + assert _carried_provider_file_id(managed) == self.file_id(filename), _decoded(managed) + return managed + + def list( + self, + key: str, + params: Mapping[str, str] | None = None, + headers: Mapping[str, str] | None = None, + *, + query: str | None = None, + ) -> httpx.Response: + suffix: Final = "" if query is None else f"?{query}" + return self.gateway.request( + "GET", f"/v1/vector_stores/{self.store}/files{suffix}", key=key, params=params, headers=headers + ) + + def listed(self, key: str, params: Mapping[str, str] | None = None) -> dict[str, JsonValue]: + return _listed(self.list(key, params if params is not None else {"model": self.model})) + + def list_requests(self) -> tuple[Request, ...]: + return tuple( + request + for request in self.wire.drain() + if (request.method, urlsplit(request.target).path) == ("GET", f"/v1/vector_stores/{self.store}/files") + ) + + def single_list_request(self) -> Request: + (request,) = self.list_requests() + return request + + +@contextmanager +def _rig(gateway: Gateway, *filenames: str, listing: Callable[[str, str], Listing] | None = None) -> Generator[_Rig]: + store: Final = "vs_" + uuid.uuid4().hex + bearer: Final = "provider-key-" + uuid.uuid4().hex[:8] + served: Final = ( + listing(store, bearer) + if listing is not None + else _constant_listing(store, *(_provider_file_id(bearer, filename) for filename in filenames)) + ) + with gateway.scenario() as scenario, wire_server(_provider(store, served)) as wire: + model: Final = scenario.model(api_base=wire.url + "/v1", api_key=bearer) + _wait_until_every_worker_serves(gateway, model) + yield _Rig(gateway, scenario, wire, store, bearer, model) + + +def test_raw_httpx_list_returns_the_uploaders_managed_ids(gateway: Gateway) -> None: + with _rig(gateway, "a.txt", "b.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + managed_b: Final = rig.upload(member.key, "b.txt") + assert read_rows(MANAGED_FILE_ROW, (managed_a,)) == [ + {"flat_model_file_ids": [rig.file_id("a.txt")], "created_by": member.user, "team_id": member.team} + ] + page: Final = rig.listed(member.key) + assert _ids(page) == (managed_a, managed_b), page + assert (page["first_id"], page["last_id"]) == (managed_a, managed_b), page + assert page["has_more"] is False, page + listed: Final = rig.single_list_request() + assert _query(listed) == {}, listed.target + assert listed.headers["authorization"] == f"Bearer {rig.bearer}", listed.headers + + +def test_attach_by_managed_id_sends_the_provider_file_id_and_lists_it_back_managed(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + attached: Final = rig.gateway.request( + "POST", f"/v1/vector_stores/{rig.store}/files", {"file_id": managed_a}, key=member.key + ) + assert attached.status_code == 200, attached.text + assert _json(attached)["id"] == managed_a, attached.text + attach_path: Final = f"/v1/vector_stores/{rig.store}/files" + attach_bodies: Final = [ + JSON_OBJECT.validate_json(request.body) + for request in rig.wire.drain() + if (request.method, request.target) == ("POST", attach_path) + ] + assert attach_bodies == [{"file_id": rig.file_id("a.txt")}], attach_bodies + assert _ids(rig.listed(member.key)) == (managed_a,) + + +def test_openai_sdk_sync_auto_pager_walks_pages_with_managed_cursors(gateway: Gateway) -> None: + with _rig(gateway, listing=_two_pages) as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + managed_b: Final = rig.upload(member.key, "b.txt") + with OpenAI(base_url=_sdk_base_url(gateway), api_key=member.key, max_retries=0) as client: + first: Final = client.vector_stores.files.list(rig.store, limit=1, extra_query={"model": rig.model}) + assert [file.id for file in first.data] == [managed_a], first.model_dump_json() + assert first.has_more is True, first.model_dump_json() + second: Final = first.get_next_page() + assert [file.id for file in second.data] == [managed_b], second.model_dump_json() + assert second.has_more is False, second.model_dump_json() + queries: Final = [_query(request) for request in rig.list_requests()] + assert queries == [{"limit": ["1"]}, {"after": [rig.file_id("a.txt")], "limit": ["1"]}], queries + + +async def test_openai_sdk_async_auto_pager_walks_pages_with_managed_cursors(gateway: Gateway) -> None: + with _rig(gateway, listing=_two_pages) as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + managed_b: Final = rig.upload(member.key, "b.txt") + async with AsyncOpenAI(base_url=_sdk_base_url(gateway), api_key=member.key, max_retries=0) as client: + first: Final = await client.vector_stores.files.list(rig.store, limit=1, extra_query={"model": rig.model}) + assert [file.id for file in first.data] == [managed_a], first.model_dump_json() + second: Final = await first.get_next_page() + assert [file.id for file in second.data] == [managed_b], second.model_dump_json() + queries: Final = [_query(request) for request in rig.list_requests()] + assert queries == [{"limit": ["1"]}, {"after": [rig.file_id("a.txt")], "limit": ["1"]}], queries + + +def test_after_cursor_with_a_managed_id_reaches_the_provider_decoded(gateway: Gateway) -> None: + with _rig(gateway, "b.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + managed_b: Final = rig.upload(member.key, "b.txt") + assert _ids(rig.listed(member.key, {"model": rig.model, "after": managed_a})) == (managed_b,) + assert _query(rig.single_list_request()) == {"after": [rig.file_id("a.txt")]} + + +def test_before_cursor_with_a_managed_id_reaches_the_provider_decoded(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + managed_b: Final = rig.upload(member.key, "b.txt") + assert _ids(rig.listed(member.key, {"model": rig.model, "before": managed_b})) == (managed_a,) + assert _query(rig.single_list_request()) == {"before": [rig.file_id("b.txt")]} + + +def test_after_cursor_with_a_raw_provider_id_is_forwarded_verbatim(gateway: Gateway) -> None: + with _rig(gateway, "b.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_b: Final = rig.upload(member.key, "b.txt") + raw_cursor: Final = "file-" + uuid.uuid4().hex[:16] + page: Final = rig.listed(member.key, {"model": rig.model, "after": raw_cursor}) + assert _query(rig.single_list_request()) == {"after": [raw_cursor]} + assert _ids(page) == (managed_b,), page + + +def test_model_header_routing_returns_managed_ids(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + page: Final = _listed(rig.list(member.key, {}, {"x-litellm-model": rig.model})) + assert _ids(page) == (managed_a,), page + assert _query(rig.single_list_request()) == {} + + +def test_managed_vector_store_registry_routing_returns_managed_ids(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + registry_bearer: Final = "registry-key-" + uuid.uuid4().hex[:8] + gateway.post( + "/vector_store/new", + { + "vector_store_id": rig.store, + "custom_llm_provider": "openai", + "vector_store_name": "managed-ids-registry", + "litellm_params": {"api_base": rig.wire.url + "/v1", "api_key": registry_bearer}, + }, + ) + rig.scenario.cleanups.callback(gateway.post, "/vector_store/delete", {"vector_store_id": rig.store}) + page: Final = _listed(rig.list(member.key, {})) + assert _ids(page) == (managed_a,), page + listed: Final = rig.single_list_request() + assert listed.headers["authorization"] == f"Bearer {registry_bearer}", listed.headers + assert _query(listed) == {}, listed.target + + +def test_team_model_fallback_routing_returns_managed_ids(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + page: Final = _listed(rig.list(member.key, {})) + assert _ids(page) == (managed_a,), page + listed: Final = rig.single_list_request() + assert listed.headers["authorization"] == f"Bearer {rig.bearer}", listed.headers + + +def test_teammate_sees_the_uploaders_managed_ids(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + uploader: Final = _member(rig.scenario, rig.model) + teammate_user: Final = rig.scenario.member(uploader.team) + teammate_key: Final = rig.scenario.key(team_id=uploader.team, user_id=teammate_user) + managed_a: Final = rig.upload(uploader.key, "a.txt") + assert _ids(rig.listed(teammate_key)) == (managed_a,) + + +def test_proxy_admin_sees_every_managed_id(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + uploader: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(uploader.key, "a.txt") + assert _ids(rig.listed(gateway.key)) == (managed_a,) + + +def test_stranger_in_another_team_sees_raw_provider_ids(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + uploader: Final = _member(rig.scenario, rig.model) + stranger: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(uploader.key, "a.txt") + assert len(read_rows(MANAGED_FILE_ROW, (managed_a,))) == 1 + page: Final = rig.listed(stranger.key) + assert _ids(page) == (rig.file_id("a.txt"),), page + assert (page["first_id"], page["last_id"]) == (rig.file_id("a.txt"), rig.file_id("a.txt")), page + + +def test_service_account_upload_is_shared_with_its_team_only(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + teammate: Final = _member(rig.scenario, rig.model) + service_account: Final = rig.scenario.key(team_id=teammate.team) + stranger: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(service_account, "a.txt") + assert read_rows(MANAGED_FILE_ROW, (managed_a,)) == [ + {"flat_model_file_ids": [rig.file_id("a.txt")], "created_by": None, "team_id": teammate.team} + ] + assert _ids(rig.listed(service_account)) == (managed_a,) + assert _ids(rig.listed(teammate.key)) == (managed_a,) + assert _ids(rig.listed(stranger.key)) == (rig.file_id("a.txt"),) + + +def test_key_without_user_or_team_owns_its_upload_alone(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + owner: Final = rig.scenario.key() + sibling: Final = rig.scenario.key() + managed_a: Final = rig.upload(owner, "a.txt") + assert read_rows(MANAGED_FILE_ROW, (managed_a,)) == [ + { + "flat_model_file_ids": [rig.file_id("a.txt")], + "created_by": f"key:{hashlib.sha256(owner.encode()).hexdigest()}", + "team_id": None, + } + ] + assert _ids(rig.listed(owner)) == (managed_a,) + assert _ids(rig.listed(sibling)) == (rig.file_id("a.txt"),) + + +def test_file_attached_by_raw_provider_id_stays_raw_beside_a_managed_one(gateway: Gateway) -> None: + raw_id: Final = "file-raw-" + uuid.uuid4().hex[:12] + with _rig(gateway, listing=_raw_then_uploaded(raw_id)) as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + attached: Final = rig.gateway.request( + "POST", f"/v1/vector_stores/{rig.store}/files", {"file_id": raw_id, "model": rig.model}, key=member.key + ) + assert attached.status_code == 200, attached.text + assert _json(attached)["id"] == raw_id, attached.text + assert _ids(rig.listed(member.key)) == (raw_id, managed_a) + + +def test_multi_model_upload_maps_only_the_provider_id_the_managed_id_carries(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as first, _rig(gateway, "a.txt") as second: + member: Final = _member(first.scenario, first.model, second.model) + managed_a: Final = _upload(gateway, member.key, f"{first.model},{second.model}", "a.txt") + (row,) = read_rows(MANAGED_FILE_ROW, (managed_a,)) + flat_ids: Final = row["flat_model_file_ids"] + assert isinstance(flat_ids, list), row + assert sorted(string_value(value) for value in flat_ids) == sorted( + (first.file_id("a.txt"), second.file_id("a.txt")) + ), row + carried: Final = _carried_provider_file_id(managed_a) + assert carried in {first.file_id("a.txt"), second.file_id("a.txt")}, carried + first_ids: Final = _ids(first.listed(member.key)) + second_ids: Final = _ids(second.listed(member.key)) + assert first_ids == ((managed_a,) if carried == first.file_id("a.txt") else (first.file_id("a.txt"),)) + assert second_ids == ((managed_a,) if carried == second.file_id("a.txt") else (second.file_id("a.txt"),)) + + +def test_deleting_the_managed_file_makes_its_listing_raw_again(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + assert _ids(rig.listed(member.key)) == (managed_a,) + deleted: Final = gateway.request("DELETE", f"/v1/files/{managed_a}", key=member.key) + assert deleted.status_code == 200, deleted.text + assert _json(deleted)["deleted"] is True, deleted.text + assert read_rows(MANAGED_FILE_ROW, (managed_a,)) == [] + assert _ids(rig.listed(member.key)) == (rig.file_id("a.txt"),) + deletes: Final = [request.target for request in rig.wire.drain() if request.method == "DELETE"] + assert deletes == [f"/v1/files/{rig.file_id('a.txt')}"], deletes + + +def test_empty_page_is_returned_unchanged(gateway: Gateway) -> None: + with _rig(gateway) as rig: + member: Final = _member(rig.scenario, rig.model) + rig.upload(member.key, "a.txt") + assert rig.listed(member.key) == _page(rig.store, ()) + + +def test_duplicate_provider_ids_in_one_page_are_both_mapped(gateway: Gateway) -> None: + with _rig(gateway, "a.txt", "a.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + page: Final = rig.listed(member.key) + assert _ids(page) == (managed_a, managed_a), page + assert (page["first_id"], page["last_id"]) == (managed_a, managed_a), page + + +def test_mixed_page_maps_only_the_managed_entries_and_the_matching_edge_ids(gateway: Gateway) -> None: + raw_id: Final = "file-raw-" + uuid.uuid4().hex[:12] + with _rig(gateway, listing=_raw_then_uploaded(raw_id)) as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + page: Final = rig.listed(member.key) + assert _ids(page) == (raw_id, managed_a), page + assert (page["first_id"], page["last_id"]) == (raw_id, managed_a), page + + +def test_repeated_identical_lists_each_reach_the_provider_and_each_map(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + assert _ids(rig.listed(member.key)) == (managed_a,) + assert _ids(rig.listed(member.key)) == (managed_a,) + targets: Final = [request.target for request in rig.list_requests()] + assert targets == [f"/v1/vector_stores/{rig.store}/files"] * 2, targets + + +def test_duplicated_managed_after_cursor_reaches_the_provider_once_decoded(gateway: Gateway) -> None: + with _rig(gateway, "b.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + managed_b: Final = rig.upload(member.key, "b.txt") + page: Final = _listed(rig.list(member.key, query=f"model={rig.model}&after={managed_a}&after={managed_a}")) + assert _ids(page) == (managed_b,), page + assert _query(rig.single_list_request()) == {"after": [rig.file_id("a.txt")]} + + +def _unpadded(raw: bytes) -> str: + return base64.urlsafe_b64encode(raw).decode().rstrip("=") + + +@pytest.mark.parametrize( + "cursor", + ( + pytest.param("12345", id="integer-like"), + pytest.param("", id="empty"), + pytest.param("x" * 5000, id="five-kilobyte"), + pytest.param(_unpadded(b"litellm_proxy:text/plain;unified_id,abc"), id="managed-without-provider-id"), + pytest.param(_unpadded(b"\xff\xfe\xfd\xfc"), id="non-utf8-base64"), + ), +) +def test_unmappable_after_cursors_are_forwarded_verbatim(gateway: Gateway, cursor: str) -> None: + with _rig(gateway, "b.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_b: Final = rig.upload(member.key, "b.txt") + page: Final = _listed(rig.list(member.key, {"model": rig.model, "after": cursor})) + assert _query(rig.single_list_request()) == {"after": [cursor]} + liveliness: Final = gateway.request("GET", "/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + assert _ids(page) == (managed_b,), page + + +def test_two_different_after_values_forward_the_last_one(gateway: Gateway) -> None: + with _rig(gateway, "b.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_b: Final = rig.upload(member.key, "b.txt") + page: Final = _listed(rig.list(member.key, query=f"model={rig.model}&after=first-value&after=second-value")) + assert _query(rig.single_list_request()) == {"after": ["second-value"]} + assert _ids(page) == (managed_b,), page + + +@pytest.mark.parametrize("status", (401, 404, 500)) +def test_provider_errors_reach_the_caller_and_other_models_keep_mapping(gateway: Gateway, status: int) -> None: + message: Final = f"provider refused listing {uuid.uuid4().hex[:8]}" + with _rig(gateway, listing=_error_listing(status, message)) as failing, _rig(gateway, "a.txt") as healthy: + member: Final = _member(failing.scenario, failing.model, healthy.model) + managed_a: Final = healthy.upload(member.key, "a.txt") + failed: Final = failing.list(member.key, {"model": failing.model}) + assert _json(failed) == _provider_error(status, message), failed.text + assert len(failing.list_requests()) == 1 + assert _ids(healthy.listed(member.key)) == (managed_a,) + liveliness: Final = gateway.request("GET", "/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + + +def test_non_json_provider_body_is_an_error_response_and_other_models_keep_mapping(gateway: Gateway) -> None: + with _rig(gateway, listing=_html_listing()) as failing, _rig(gateway, "a.txt") as healthy: + member: Final = _member(failing.scenario, failing.model, healthy.model) + managed_a: Final = healthy.upload(member.key, "a.txt") + failed: Final = failing.list(member.key, {"model": failing.model}) + assert failed.status_code == 500, failed.text + assert string_value(object_value(_json(failed)["error"])["message"]), failed.text + assert _ids(healthy.listed(member.key)) == (managed_a,) + liveliness: Final = gateway.request("GET", "/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + + +def test_non_string_ids_in_a_page_are_left_alone_while_strings_map(gateway: Gateway) -> None: + with _rig(gateway, listing=_uploaded_then_integer) as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + page: Final = rig.listed(member.key) + assert _ids(page) == (managed_a, 7), page + assert (page["first_id"], page["last_id"]) == (managed_a, 7), page + + +def test_retrieving_the_managed_file_still_resolves_to_the_provider_file(gateway: Gateway) -> None: + with _rig(gateway, "a.txt") as rig: + member: Final = _member(rig.scenario, rig.model) + managed_a: Final = rig.upload(member.key, "a.txt") + retrieved: Final = gateway.request("GET", f"/v1/files/{managed_a}", key=member.key) + assert retrieved.status_code == 200, retrieved.text + file: Final = _json(retrieved) + assert (file["id"], file["object"], file["purpose"]) == (managed_a, "file", "user_data"), retrieved.text + + +def _burst(gateway: Gateway, store: str, key: str, model: str, size: int) -> tuple[httpx.Response, ...]: + def one(_: int) -> httpx.Response: + return gateway.request("GET", f"/v1/vector_stores/{store}/files", key=key, params={"model": model}) + + with ThreadPoolExecutor(max_workers=size) as pool: + return tuple(pool.map(one, range(size))) + + +@pytest.mark.timeout(180) +def test_provider_outage_mid_burst_fails_loudly_and_mapping_resumes_after_recovery(gateway: Gateway) -> None: + store: Final = "vs_" + uuid.uuid4().hex + bearer: Final = "provider-key-" + uuid.uuid4().hex[:8] + provider_a: Final = _provider_file_id(bearer, "a.txt") + respond: Final = _provider(store, _constant_listing(store, provider_a)) + with gateway.scenario() as scenario: + with wire_server(respond) as wire: + model: Final = scenario.model(api_base=wire.url + "/v1", api_key=bearer) + _wait_until_every_worker_serves(gateway, model) + member: Final = _member(scenario, model) + managed_a: Final = _upload(gateway, member.key, model, "a.txt") + assert _carried_provider_file_id(managed_a) == provider_a + served: Final = _burst(gateway, store, member.key, model, 40) + assert [_ids(_listed(response)) for response in served] == [(managed_a,)] * 40 + assert sum(1 for request in wire.drain() if request.method == "GET") == 40 + port: Final = urlsplit(wire.url).port + assert port is not None + failed: Final = _burst(gateway, store, member.key, model, 20) + assert [response.status_code for response in failed] == [500] * 20, [r.text for r in failed[:3]] + for response in failed: + assert string_value(object_value(_json(response)["error"])["message"]), response.text + liveliness: Final = gateway.request("GET", "/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + with wire_server(respond, port=port) as revived: + recovered: Final = _burst(gateway, store, member.key, model, 40) + assert [_ids(_listed(response)) for response in recovered] == [(managed_a,)] * 40 + assert sum(1 for request in revived.drain() if request.method == "GET") == 40 + + +def _open_connections_to(pid: int, url: str) -> int: + port: Final = urlsplit(url).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +def _tolerant_list(gateway: Gateway, store: str, key: str, model: str) -> httpx.Response | None: + try: + return gateway.request("GET", f"/v1/vector_stores/{store}/files", key=key, params={"model": model}) + except httpx.HTTPError: + return None + + +@pytest.mark.timeout(300) +def test_worker_sigkill_mid_burst_leaves_the_sibling_mapping_ids(gateway: Gateway, tmp_path: Path) -> None: + store: Final = "vs_" + uuid.uuid4().hex + bearer: Final = "provider-key-" + uuid.uuid4().hex[:8] + provider_a: Final = _provider_file_id(bearer, "a.txt") + release: Final = threading.Event() + held: Final[SimpleQueue[str]] = SimpleQueue() + + def held_listing(request: Request) -> Reply: + held.put(request.target) + assert release.wait(timeout=120), "The burst was never released" + return _json_reply(_page(store, (provider_a,))) + + with gateway.scenario() as scenario, wire_server(_provider(store, held_listing)) as wire: + model: Final = scenario.model(api_base=wire.url + "/v1", api_key=bearer) + _wait_until_every_worker_serves(gateway, model) + member: Final = _member(scenario, model) + managed_a: Final = _upload(gateway, member.key, model, "a.txt") + with owned_proxy_process(gateway, tmp_path, {}, workers=2) as owned: + candidate: Final = owned.gateway + workers: Final = eventually( + lambda: tuple(int(match.group(1)) for match in STARTED_WORKER.finditer(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=60, + ) + with ThreadPoolExecutor(max_workers=20) as pool: + burst: Final = tuple( + pool.submit(_tolerant_list, candidate, store, member.key, model) for _ in range(20) + ) + eventually(held.qsize, lambda size: size == 20, seconds=60) + held_by: Final = MappingProxyType({pid: _open_connections_to(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + served: Final = tuple(result for future in burst if (result := future.result()) is not None) + assert held_by[survivor_pid] >= 1, held_by + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + for response in served: + assert _ids(_listed(response)) == (managed_a,) + assert psutil.Process(survivor_pid).is_running() + follow_up: Final = eventually( + lambda: _tolerant_list(candidate, store, member.key, model), + lambda response: response is not None and response.status_code == 200, + seconds=60, + ) + assert follow_up is not None + assert _ids(_listed(follow_up)) == (managed_a,) diff --git a/tests/unit/enterprise/proxy/test_managed_files_hook.py b/tests/unit/enterprise/proxy/test_managed_files_hook.py index 74bd67efaf2..d99ea5ab445 100644 --- a/tests/unit/enterprise/proxy/test_managed_files_hook.py +++ b/tests/unit/enterprise/proxy/test_managed_files_hook.py @@ -9,9 +9,10 @@ import asyncio import base64 import json import logging +from types import MappingProxyType import pytest -from typing import Optional +from typing import Final, Optional from unittest.mock import AsyncMock, MagicMock, patch from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth @@ -299,6 +300,90 @@ 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_provider_file_id_resolver_returns_owned_mappings_with_owner_scoped_filter() -> ( + None +): + managed_files: Final = _make_managed_files_instance() + managed_row: Final = MagicMock( + unified_file_id="unified-file-id", + flat_model_file_ids=["file-provider-1", "file-provider-2"], + ) + find_many: Final = AsyncMock(return_value=[managed_row]) + managed_files.prisma_client.db.litellm_managedfiletable.find_many = find_many + + unified_file_ids: Final = ( + await managed_files.get_unified_file_ids_for_provider_file_ids( + provider_file_ids=( + "file-provider-1", + "file-provider-2", + "file-unmanaged-2", + "file-provider-1", + ), + user_api_key_dict=_make_team_member_api_key_dict(), + ) + ) + + assert unified_file_ids == { + "file-provider-1": "unified-file-id", + "file-provider-2": "unified-file-id", + } + assert isinstance(unified_file_ids, MappingProxyType) + find_many.assert_awaited_once_with( + where={ + "OR": [{"created_by": "test-user"}, {"team_id": "test-team"}], + "flat_model_file_ids": { + "hasSome": ["file-provider-1", "file-provider-2", "file-unmanaged-2"], + }, + } + ) + + +@pytest.mark.asyncio +async def test_provider_file_id_resolver_denies_unowned_callers_without_database_query() -> ( + None +): + managed_files: Final = _make_managed_files_instance() + find_many: Final = AsyncMock() + managed_files.prisma_client.db.litellm_managedfiletable.find_many = find_many + no_owner: Final = UserAPIKeyAuth( + api_key=None, + token=None, + user_id=None, + team_id=None, + parent_otel_span=None, + ) + + unified_file_ids: Final = ( + await managed_files.get_unified_file_ids_for_provider_file_ids( + provider_file_ids=("file-provider-1",), + user_api_key_dict=no_owner, + ) + ) + + assert unified_file_ids == {} + assert isinstance(unified_file_ids, MappingProxyType) + find_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_provider_file_id_resolver_skips_database_query_for_empty_input() -> None: + managed_files: Final = _make_managed_files_instance() + find_many: Final = AsyncMock() + managed_files.prisma_client.db.litellm_managedfiletable.find_many = find_many + + unified_file_ids: Final = ( + await managed_files.get_unified_file_ids_for_provider_file_ids( + provider_file_ids=(), + user_api_key_dict=_make_user_api_key_dict(), + ) + ) + + assert unified_file_ids == {} + assert isinstance(unified_file_ids, MappingProxyType) + find_many.assert_not_awaited() + + @pytest.mark.asyncio async def test_afile_list_returns_owner_scoped_managed_files(): managed_files = _make_managed_files_instance() diff --git a/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index b1bd7ccbf0f..268000517d3 100644 --- a/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -1,3 +1,5 @@ +import base64 +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -5,6 +7,12 @@ from fastapi import HTTPException, Request, Response import litellm from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable, UserAPIKeyAuth +from litellm.types.utils import SpecialEnums +from litellm.types.vector_store_files import ( + VectorStoreFileListResponse, + VectorStoreFileObject, + VectorStoreFileStatus, +) def _mock_request() -> MagicMock: @@ -107,18 +115,54 @@ async def test_vector_store_file_create_forces_path_id_over_body_id(): @pytest.mark.asyncio -async def test_vector_store_file_list_resolves_managed_vector_store_before_team_fallback(): - import base64 - +async def test_vector_store_file_list_resolves_managed_ids_and_cursors(): from litellm.proxy.vector_store_files_endpoints.endpoints import ( vector_store_file_list, ) captured_data = {} + provider_file_id: Final = "file-list-owned" + managed_file_data: Final = ( + SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format( + "application/json", + "unified-file", + "managed-deployment", + provider_file_id, + "managed-deployment-id", + ) + ) + managed_file_id: Final = ( + base64.urlsafe_b64encode(managed_file_data.encode()).decode().rstrip("=") + ) + user_api_key_dict: Final = UserAPIKeyAuth(team_models=["team-openai"]) + managed_file: Final[VectorStoreFileObject] = { + "id": provider_file_id, + "object": "vector_store.file", + "created_at": 1700000000, + "usage_bytes": 100, + "vector_store_id": "vs_provider_native", + "status": VectorStoreFileStatus.COMPLETED, + "last_error": None, + "chunking_strategy": {"type": "auto"}, + "attributes": {"source": "test"}, + } + provider_response: Final[VectorStoreFileListResponse] = { + "object": "list", + "data": [managed_file], + "first_id": provider_file_id, + "last_id": provider_file_id, + "has_more": False, + } + expected_response: Final[VectorStoreFileListResponse] = { + **provider_response, + "data": [{**managed_file, "id": managed_file_id}], + "first_id": managed_file_id, + "last_id": managed_file_id, + } async def fake_base_process(self, **kwargs): captured_data.update(self.data) - return {"ok": True} + return provider_response raw_vector_store_id = ( "litellm_proxy:vector_store;" @@ -133,7 +177,7 @@ async def test_vector_store_file_list_resolves_managed_vector_store_before_team_ request = _mock_request() request.method = "GET" - request.query_params = {"limit": "10"} + request.query_params = {"after": managed_file_id, "limit": "10"} request.url.path = f"/v1/vector_stores/{vector_store_id}/files" llm_router = MagicMock() @@ -147,6 +191,11 @@ async def test_vector_store_file_list_resolves_managed_vector_store_before_team_ } llm_router.get_deployment_credentials_with_provider.side_effect = get_credentials + managed_files_obj = MagicMock() + resolver = AsyncMock(return_value={provider_file_id: managed_file_id}) + managed_files_obj.get_unified_file_ids_for_provider_file_ids = resolver + proxy_logging_obj = MagicMock() + proxy_logging_obj.get_proxy_hook.return_value = managed_files_obj with ( patch( @@ -154,6 +203,7 @@ async def test_vector_store_file_list_resolves_managed_vector_store_before_team_ new=AsyncMock(return_value=None), ), patch("litellm.proxy.proxy_server.llm_router", llm_router), + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj), patch( "litellm.proxy.vector_store_files_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", new=fake_base_process, @@ -163,16 +213,22 @@ async def test_vector_store_file_list_resolves_managed_vector_store_before_team_ vector_store_id=vector_store_id, request=request, fastapi_response=Response(), - user_api_key_dict=UserAPIKeyAuth(team_models=["team-openai"]), + user_api_key_dict=user_api_key_dict, ) - assert response == {"ok": True} + assert response == expected_response + assert captured_data["after"] == provider_file_id assert captured_data["vector_store_id"] == "vs_provider_native" assert captured_data["api_key"] == "sk-managed-deployment" assert captured_data["model"] == "openai/managed-deployment" llm_router.get_deployment_credentials_with_provider.assert_called_once_with( model_id="managed-deployment" ) + proxy_logging_obj.get_proxy_hook.assert_called_once_with("managed_files") + resolver.assert_awaited_once_with( + provider_file_ids=(provider_file_id,), + user_api_key_dict=user_api_key_dict, + ) @pytest.mark.asyncio diff --git a/tests/unit/proxy/vector_store_files_endpoints/test_endpoints.py b/tests/unit/proxy/vector_store_files_endpoints/test_endpoints.py index 4cb3a3d4c7f..271deab7b36 100644 --- a/tests/unit/proxy/vector_store_files_endpoints/test_endpoints.py +++ b/tests/unit/proxy/vector_store_files_endpoints/test_endpoints.py @@ -10,9 +10,11 @@ is attached to a vector store or read back under shared provider credentials. """ import base64 +from collections.abc import Mapping, Sequence +from copy import deepcopy from dataclasses import dataclass -from typing import Literal -from unittest.mock import MagicMock, patch +from typing import Final, Literal +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -23,8 +25,15 @@ import litellm from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.vector_store_files_endpoints.endpoints import ( _update_request_data_with_managed_file_id, + _with_managed_file_list_ids, + _with_provider_file_id_cursors, ) from litellm.types.utils import SpecialEnums +from litellm.types.vector_store_files import ( + VectorStoreFileListResponse, + VectorStoreFileObject, + VectorStoreFileStatus, +) RAW_FILE_ID = "file-victim-abc123" CALLER = UserAPIKeyAuth(api_key="sk-test", user_id="attacker-user", team_id="team-b") @@ -51,13 +60,46 @@ class ManagedResourceAccessCheckerStub: return False -def _unified_file_id() -> str: +@dataclass(frozen=True) +class ManagedFileIdResolverStub: + resolver: AsyncMock + + async def get_unified_file_ids_for_provider_file_ids( + self, + provider_file_ids: Sequence[str], + user_api_key_dict: UserAPIKeyAuth, + ) -> Mapping[str, str]: + return await self.resolver( + provider_file_ids=provider_file_ids, + user_api_key_dict=user_api_key_dict, + ) + + +def _unified_file_id(provider_file_id: str = RAW_FILE_ID) -> str: unified = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format( - "application/json", "victim-unified-id", "gpt-4o-mini", RAW_FILE_ID, "gpt-4o-mini-id" + "application/json", + "victim-unified-id", + "gpt-4o-mini", + provider_file_id, + "gpt-4o-mini-id", ) return base64.urlsafe_b64encode(unified.encode()).decode().rstrip("=") +def _vector_store_file_row(file_id: str) -> VectorStoreFileObject: + return { + "id": file_id, + "object": "vector_store.file", + "created_at": 1700000000, + "usage_bytes": 100, + "vector_store_id": "vs-test", + "status": VectorStoreFileStatus.COMPLETED, + "last_error": None, + "chunking_strategy": {"type": "auto"}, + "attributes": {"source": "test"}, + } + + async def _resolve( file_id: str, file_access: Literal["allow", "deny", "missing"] = "allow", @@ -72,6 +114,113 @@ async def _resolve( ) +@pytest.mark.parametrize( + "provider_ids", + [ + (RAW_FILE_ID, "file-unmanaged-123"), + ("file-unmanaged-123", RAW_FILE_ID), + ], +) +@pytest.mark.asyncio +async def test_vector_store_file_list_maps_owned_ids_and_preserves_raw_ids( + provider_ids: tuple[str, str], +) -> None: + managed_file_id: Final = _unified_file_id() + expected_provider_ids: Final = tuple( + managed_file_id if provider_file_id == RAW_FILE_ID else provider_file_id + for provider_file_id in provider_ids + ) + provider_response: Final[VectorStoreFileListResponse] = { + "object": "list", + "data": [ + _vector_store_file_row(provider_file_id) + for provider_file_id in provider_ids + ], + "first_id": provider_ids[0], + "last_id": provider_ids[1], + "has_more": True, + } + original_response: Final = deepcopy(provider_response) + resolver: Final = AsyncMock(return_value={RAW_FILE_ID: managed_file_id}) + managed_files_obj: Final = ManagedFileIdResolverStub(resolver=resolver) + + response: Final = await _with_managed_file_list_ids( + response=provider_response, + managed_files_obj=managed_files_obj, + user_api_key_dict=CALLER, + ) + + expected_response: Final[VectorStoreFileListResponse] = { + "object": "list", + "data": [ + _vector_store_file_row(provider_file_id) + for provider_file_id in expected_provider_ids + ], + "first_id": expected_provider_ids[0], + "last_id": expected_provider_ids[1], + "has_more": True, + } + assert response == expected_response + assert provider_response == original_response + resolver.assert_awaited_once_with( + provider_file_ids=tuple(dict.fromkeys(provider_ids)), + user_api_key_dict=CALLER, + ) + + +@pytest.mark.asyncio +async def test_vector_store_file_list_only_maps_round_trippable_ids() -> None: + managed_file_id: Final = _unified_file_id("file-model-a") + provider_response: Final[VectorStoreFileListResponse] = { + "object": "list", + "data": [ + _vector_store_file_row("file-model-a"), + _vector_store_file_row("file-model-b"), + ], + "first_id": "file-model-a", + "last_id": "file-model-b", + "has_more": False, + } + resolver: Final = AsyncMock( + return_value={ + "file-model-a": managed_file_id, + "file-model-b": managed_file_id, + } + ) + managed_files_obj: Final = ManagedFileIdResolverStub(resolver=resolver) + + response: Final = await _with_managed_file_list_ids( + response=provider_response, + managed_files_obj=managed_files_obj, + user_api_key_dict=CALLER, + ) + + expected_response: Final[VectorStoreFileListResponse] = { + "object": "list", + "data": [ + _vector_store_file_row(managed_file_id), + _vector_store_file_row("file-model-b"), + ], + "first_id": managed_file_id, + "last_id": "file-model-b", + "has_more": False, + } + assert response == expected_response + + +def test_vector_store_file_list_translates_managed_cursors_and_preserves_raw_after() -> ( + None +): + managed_file_id: Final = _unified_file_id() + + assert _with_provider_file_id_cursors( + {"after": managed_file_id, "before": managed_file_id} + ) == {"after": RAW_FILE_ID, "before": RAW_FILE_ID} + assert _with_provider_file_id_cursors({"after": RAW_FILE_ID}) == { + "after": RAW_FILE_ID + } + + @pytest.mark.asyncio async def test_raw_file_id_rejected_when_managed_files_required(): with patch.object(litellm, "require_managed_files", True):