mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
0d5ea45c86
commit
8efb4a21f6
13 changed files with 1423 additions and 24 deletions
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
@ -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)])
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)])
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)])
|
||||
}
|
||||
|
||||
|
|
|
|||
184
tests/integration/database/test_managed_file_flat_ids_index.py
Normal file
184
tests/integration/database/test_managed_file_flat_ids_index.py
Normal 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())
|
||||
|
|
@ -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,)
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue