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 <shivam@berri.ai>
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>
This commit is contained in:
devin-ai-integration[bot] 2026-10-02 21:28:50 -07:00 • committed by GitHub
parent 0d5ea45c86
commit 8efb4a21f6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 1423 additions and 24 deletions

View file

@ -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]:

View file

@ -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");

View file

@ -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)])
}

View file

@ -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):

View file

@ -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

View file

@ -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)])
}

View file

@ -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,

View file

@ -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)])
}

View file

@ -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())

View file

@ -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"<html>upstream maintenance</html>", 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,)

View file

@ -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()

View file

@ -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

View file

@ -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):