mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(vector-stores): harden Milvus gRPC transport
This commit is contained in:
parent
c76c5c78c3
commit
2a13c88789
16 changed files with 313 additions and 288 deletions
|
|
@ -3,7 +3,7 @@
|
|||
"limit": 13429
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2198
|
||||
"limit": 2178
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 319
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 3369
|
||||
"limit": 3367
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 5570
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15281
|
||||
"limit": 15270
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -99,19 +99,19 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44094
|
||||
"limit": 43439
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38283
|
||||
"limit": 38255
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19584
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 29829
|
||||
"limit": 29827
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 110
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 4
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 816
|
||||
"limit": 814
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
|
|||
|
|
@ -62,7 +62,7 @@ class AzurePassthroughConfig(BasePassthroughConfig):
|
|||
) -> dict:
|
||||
return BaseAzureLLM._base_validate_azure_environment(
|
||||
headers=headers,
|
||||
litellm_params=GenericLiteLLMParams.model_validate({**litellm_params, "api_key": api_key}),
|
||||
litellm_params=GenericLiteLLMParams(**{**litellm_params, "api_key": api_key}),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -161,13 +161,11 @@ class _MilvusSearchOptions(BaseModel):
|
|||
def result_limit(self) -> int:
|
||||
return self.max_num_results or self.limit
|
||||
|
||||
def output_fields_with_text(
|
||||
self, text_field: str
|
||||
) -> list[str]: # mutable-ok: PyMilvus requires output_fields as a list
|
||||
def output_fields_with_text(self, text_field: str) -> list[str]: # mutable-ok: PyMilvus requires a list
|
||||
output_fields: Final = self.output_fields or ()
|
||||
if "*" in output_fields or text_field in output_fields:
|
||||
return list(output_fields)
|
||||
return [*output_fields, text_field]
|
||||
return list(output_fields) # mutable-ok: PyMilvus requires output_fields as a list
|
||||
return [*output_fields, text_field] # mutable-ok: PyMilvus requires output_fields as a list
|
||||
|
||||
|
||||
class _EmbeddingItem(BaseModel):
|
||||
|
|
|
|||
|
|
@ -112,7 +112,7 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
|
|||
data.pop("custom_llm_provider", None)
|
||||
data.pop("litellm_credential_name", None)
|
||||
managed_data: Final = build_request_data_from_managed_vector_store(vector_store_to_run)
|
||||
request_data: Final = {**data, **managed_data}
|
||||
request_data: Final = {**data, **managed_data} # mutable-ok: request processing requires a mutable payload
|
||||
if user_api_key_dict is not None:
|
||||
assert_proxy_admin_for_user_supplied_vector_store_connection(
|
||||
custom_llm_provider=request_data.get("custom_llm_provider"),
|
||||
|
|
@ -665,4 +665,4 @@ async def index_list(
|
|||
)
|
||||
|
||||
indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db(prisma_client)
|
||||
return IndexListResponse(data=indexes)
|
||||
return IndexListResponse(data=tuple(indexes))
|
||||
|
|
|
|||
|
|
@ -8,8 +8,8 @@ All /vector_store management endpoints
|
|||
/vector_store/list
|
||||
"""
|
||||
|
||||
import copy
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
|
@ -45,7 +45,10 @@ from litellm.types.vector_stores import (
|
|||
VectorStoreInfoRequest,
|
||||
VectorStoreUpdateRequest,
|
||||
)
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
|
||||
from litellm.vector_stores.vector_store_registry import (
|
||||
VectorStoreRegistry,
|
||||
deserialize_litellm_params,
|
||||
)
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
|
@ -110,31 +113,11 @@ def _redact_sensitive_litellm_params(litellm_params: object, _depth: int = 0) ->
|
|||
|
||||
|
||||
def _validated_litellm_params(
|
||||
litellm_params: dict[str, Any], # mutable-ok: provider parameters arrive as a mutable request object
|
||||
) -> dict[str, Any]: # mutable-ok: persistence validation returns a serializable parameter dict
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> Mapping[str, object]:
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
trusted: Final = litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True
|
||||
validated: Final = GenericLiteLLMParams(**litellm_params).model_dump(exclude_none=True)
|
||||
if trusted:
|
||||
validated[MILVUS_ADMIN_CONFIGURED_CONNECTION] = True
|
||||
return validated
|
||||
|
||||
|
||||
def _litellm_params_dict(
|
||||
litellm_params: object,
|
||||
) -> dict[str, Any]: # mutable-ok: update authorization merges a copy of persisted parameters
|
||||
if isinstance(litellm_params, dict):
|
||||
return dict(litellm_params) # mutable-ok: callers require an isolated copy for effective-connection merging
|
||||
if isinstance(litellm_params, str):
|
||||
try:
|
||||
parsed: Final = json.loads(litellm_params)
|
||||
if isinstance(parsed, dict):
|
||||
return dict(parsed) # mutable-ok: parsed persistence data must be copied before merging
|
||||
return {} # mutable-ok: non-object persistence data normalizes to an empty mutable mapping
|
||||
except (TypeError, ValueError):
|
||||
return {} # mutable-ok: invalid persisted parameters normalize to an empty mutable mapping
|
||||
return {} # mutable-ok: absent persisted parameters normalize to an empty mutable mapping
|
||||
return GenericLiteLLMParams.model_validate(litellm_params).model_dump(exclude_none=True)
|
||||
|
||||
|
||||
def _reject_config_vector_store_id(vector_store_id: str) -> None:
|
||||
|
|
@ -191,6 +174,40 @@ async def _check_vector_store_access(
|
|||
return await can_user_access_vector_store(vector_store=vector_store, user_api_key_dict=user_api_key_dict)
|
||||
|
||||
|
||||
def _vector_store_create_data(
|
||||
vector_store_id: str,
|
||||
custom_llm_provider: str,
|
||||
vector_store_name: str | None,
|
||||
vector_store_description: str | None,
|
||||
vector_store_metadata: Mapping[str, object] | None,
|
||||
litellm_params: Mapping[str, object] | None,
|
||||
litellm_credential_name: str | None,
|
||||
team_id: str | None,
|
||||
user_id: str | None,
|
||||
) -> dict[str, object]: # mutable-ok: Prisma create requires a mutable data dict
|
||||
serialized_params: Final = safe_dumps(_validated_litellm_params(litellm_params)) if litellm_params else "{}"
|
||||
return { # mutable-ok: Prisma create requires a mutable data dict
|
||||
"vector_store_id": vector_store_id,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
**{
|
||||
key: value
|
||||
for key, value in (
|
||||
("vector_store_name", vector_store_name),
|
||||
("vector_store_description", vector_store_description),
|
||||
(
|
||||
"vector_store_metadata",
|
||||
safe_dumps(vector_store_metadata) if vector_store_metadata is not None else None,
|
||||
),
|
||||
("litellm_credential_name", litellm_credential_name),
|
||||
("team_id", team_id),
|
||||
("user_id", user_id),
|
||||
)
|
||||
if value is not None
|
||||
},
|
||||
"litellm_params": serialized_params,
|
||||
}
|
||||
|
||||
|
||||
async def create_vector_store_in_db(
|
||||
vector_store_id: str,
|
||||
custom_llm_provider: str,
|
||||
|
|
@ -232,40 +249,17 @@ async def create_vector_store_in_db(
|
|||
detail=f"Vector store with ID {vector_store_id} already exists",
|
||||
)
|
||||
|
||||
# Prepare data for database
|
||||
data_to_create: Final[dict[str, object]] = {
|
||||
"vector_store_id": vector_store_id,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
}
|
||||
|
||||
if vector_store_name is not None:
|
||||
data_to_create["vector_store_name"] = vector_store_name
|
||||
if vector_store_description is not None:
|
||||
data_to_create["vector_store_description"] = vector_store_description
|
||||
if vector_store_metadata is not None:
|
||||
data_to_create["vector_store_metadata"] = safe_dumps(vector_store_metadata)
|
||||
if litellm_credential_name is not None:
|
||||
data_to_create["litellm_credential_name"] = litellm_credential_name
|
||||
if team_id is not None:
|
||||
data_to_create["team_id"] = team_id
|
||||
if user_id is not None:
|
||||
data_to_create["user_id"] = user_id
|
||||
|
||||
# Handle litellm_params - always provide at least an empty dict.
|
||||
# The earlier behaviour resolved ``litellm_embedding_config`` from the
|
||||
# admin-configured router/DB model and persisted the cleartext result
|
||||
# (``api_key``, ``api_base``, ``api_version``) into this row. That
|
||||
# exposed every env-stored embedding-model credential on the
|
||||
# ``/vector_store/{new,info,update,list}`` responses. Keep the user's
|
||||
# raw ``litellm_embedding_model`` reference; each search embeds the
|
||||
# query through the router at request time, so the credentials stay
|
||||
# on the deployment and never reach the database.
|
||||
if litellm_params:
|
||||
litellm_params_dict: Final = _validated_litellm_params(litellm_params)
|
||||
data_to_create["litellm_params"] = safe_dumps(litellm_params_dict)
|
||||
else:
|
||||
# Provide empty dict if no litellm_params provided
|
||||
data_to_create["litellm_params"] = safe_dumps({})
|
||||
data_to_create: Final = _vector_store_create_data(
|
||||
vector_store_id=vector_store_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vector_store_name=vector_store_name,
|
||||
vector_store_description=vector_store_description,
|
||||
vector_store_metadata=vector_store_metadata,
|
||||
litellm_params=litellm_params,
|
||||
litellm_credential_name=litellm_credential_name,
|
||||
team_id=team_id,
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
# Create in database
|
||||
_new_vector_store: Final = await _vector_store_table(prisma_client).create(data=data_to_create)
|
||||
|
|
@ -361,6 +355,54 @@ async def new_vector_store(
|
|||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
def _vector_stores_by_id(
|
||||
vector_stores: Sequence[LiteLLM_ManagedVectorStore],
|
||||
) -> Mapping[str, LiteLLM_ManagedVectorStore]:
|
||||
return { # mutable-ok: registry synchronization requires ID-keyed replacement semantics
|
||||
vector_store_id: vector_store
|
||||
for vector_store in vector_stores
|
||||
if (vector_store_id := vector_store.get("vector_store_id"))
|
||||
}
|
||||
|
||||
|
||||
def _synchronize_vector_store_registry(
|
||||
vector_stores_from_db: Sequence[LiteLLM_ManagedVectorStore],
|
||||
) -> Mapping[str, LiteLLM_ManagedVectorStore]:
|
||||
database_stores: Final = _vector_stores_by_id(vector_stores_from_db)
|
||||
registry: Final = litellm.vector_store_registry
|
||||
if registry is None:
|
||||
return database_stores
|
||||
|
||||
memory_stores: Final = _vector_stores_by_id(registry.vector_stores)
|
||||
config_ids: Final = registry.config_vector_store_ids
|
||||
stale_ids: Final = tuple(
|
||||
vector_store_id
|
||||
for vector_store_id in memory_stores
|
||||
if vector_store_id not in config_ids and vector_store_id not in database_stores
|
||||
)
|
||||
for vector_store_id in stale_ids:
|
||||
registry.delete_vector_store_from_registry(vector_store_id=vector_store_id)
|
||||
verbose_proxy_logger.debug("Removed deleted vector store %s from in-memory registry", vector_store_id)
|
||||
for vector_store_id, vector_store in database_stores.items():
|
||||
if vector_store_id not in config_ids:
|
||||
registry.update_vector_store_in_registry(vector_store_id=vector_store_id, updated_data=vector_store)
|
||||
|
||||
return { # mutable-ok: listing and access filtering require an ID-keyed dict
|
||||
**database_stores,
|
||||
**{
|
||||
vector_store_id: vector_store
|
||||
for vector_store_id, vector_store in memory_stores.items()
|
||||
if vector_store_id in config_ids
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _redact_vector_store(vector_store: LiteLLM_ManagedVectorStore) -> LiteLLM_ManagedVectorStore:
|
||||
redacted: Final = LiteLLM_ManagedVectorStore(**vector_store)
|
||||
redacted["litellm_params"] = _redact_sensitive_litellm_params(vector_store.get("litellm_params"))
|
||||
return redacted
|
||||
|
||||
|
||||
@router.get(
|
||||
"/vector_store/list",
|
||||
tags=["vector store management"],
|
||||
|
|
@ -391,69 +433,17 @@ async def list_vector_stores(
|
|||
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
vector_store_map: Final[dict[str, LiteLLM_ManagedVectorStore]] = {}
|
||||
db_vector_store_ids: Final[set] = set()
|
||||
|
||||
try:
|
||||
# Get vector stores from database first (source of truth)
|
||||
vector_stores_from_db: Final = await VectorStoreRegistry._get_vector_stores_from_db(prisma_client=prisma_client)
|
||||
|
||||
# Build map from database vector stores
|
||||
for vector_store in vector_stores_from_db:
|
||||
vector_store_id = vector_store.get("vector_store_id", None)
|
||||
if vector_store_id:
|
||||
vector_store_map[vector_store_id] = vector_store
|
||||
db_vector_store_ids.add(vector_store_id)
|
||||
|
||||
# Process in-memory vector stores
|
||||
if litellm.vector_store_registry is not None:
|
||||
in_memory_vector_stores: Final = copy.deepcopy(litellm.vector_store_registry.vector_stores)
|
||||
config_vector_store_ids: Final = litellm.vector_store_registry.config_vector_store_ids
|
||||
|
||||
vector_stores_to_delete_from_memory: Final[list[str]] = []
|
||||
|
||||
for vector_store in in_memory_vector_stores:
|
||||
vector_store_id = vector_store.get("vector_store_id", None)
|
||||
if not vector_store_id:
|
||||
continue
|
||||
|
||||
if vector_store_id in config_vector_store_ids:
|
||||
vector_store_map[vector_store_id] = vector_store
|
||||
elif vector_store_id not in db_vector_store_ids:
|
||||
verbose_proxy_logger.info(
|
||||
"Vector store %s exists in memory but not in database - marking for deletion from cache",
|
||||
vector_store_id,
|
||||
)
|
||||
vector_stores_to_delete_from_memory.append(vector_store_id)
|
||||
# If not in our map yet, add it (only in-memory, not in DB)
|
||||
elif vector_store_id not in vector_store_map:
|
||||
vector_store_map[vector_store_id] = vector_store
|
||||
|
||||
# Synchronize in-memory registry with database
|
||||
# 1. Remove deleted vector stores from memory
|
||||
for vs_id in vector_stores_to_delete_from_memory:
|
||||
litellm.vector_store_registry.delete_vector_store_from_registry(vector_store_id=vs_id)
|
||||
verbose_proxy_logger.debug("Removed deleted vector store %s from in-memory registry", vs_id)
|
||||
|
||||
# 2. Update in-memory registry with database versions (for updates)
|
||||
for vector_store in vector_stores_from_db:
|
||||
vector_store_id = vector_store.get("vector_store_id", None)
|
||||
if vector_store_id and vector_store_id not in config_vector_store_ids:
|
||||
litellm.vector_store_registry.update_vector_store_in_registry(
|
||||
vector_store_id=vector_store_id, updated_data=vector_store
|
||||
)
|
||||
|
||||
# Filter vector stores based on access control
|
||||
accessible_vector_stores: Final = []
|
||||
for vs in await filter_listable_vector_stores(vector_store_map.values(), user_api_key_dict):
|
||||
redacted = LiteLLM_ManagedVectorStore(**vs)
|
||||
redacted["litellm_params"] = _redact_sensitive_litellm_params(vs.get("litellm_params"))
|
||||
accessible_vector_stores.append(redacted)
|
||||
vector_store_map: Final = _synchronize_vector_store_registry(vector_stores_from_db)
|
||||
accessible_vector_stores: Final = [ # mutable-ok: response model requires a list
|
||||
_redact_vector_store(vector_store)
|
||||
for vector_store in await filter_listable_vector_stores(vector_store_map.values(), user_api_key_dict)
|
||||
]
|
||||
|
||||
total_count: Final = len(accessible_vector_stores)
|
||||
total_pages: Final = (total_count + page_size - 1) // page_size
|
||||
|
||||
# Format response using LiteLLM_ManagedVectorStoreListResponse
|
||||
response: Final = LiteLLM_ManagedVectorStoreListResponse(
|
||||
object="list",
|
||||
data=accessible_vector_stores,
|
||||
|
|
@ -468,6 +458,23 @@ async def list_vector_stores(
|
|||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
async def _vector_store_delete_target(
|
||||
vector_store_id: str,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> tuple[LiteLLM_ManagedVectorStore, bool, bool]:
|
||||
row: Final = await _vector_store_table(prisma_client).find_unique(where={"vector_store_id": vector_store_id})
|
||||
registry: Final = litellm.vector_store_registry
|
||||
memory_store: Final = (
|
||||
registry.get_litellm_managed_vector_store_from_registry(vector_store_id=vector_store_id)
|
||||
if registry is not None
|
||||
else None
|
||||
)
|
||||
target: Final = _row_to_vector_store(row) if row is not None else memory_store
|
||||
if target is None:
|
||||
raise HTTPException(status_code=404, detail=f"Vector store with ID {vector_store_id} not found")
|
||||
return target, row is not None, memory_store is not None
|
||||
|
||||
|
||||
@router.post(
|
||||
"/vector_store/delete",
|
||||
tags=["vector store management"],
|
||||
|
|
@ -492,50 +499,21 @@ async def delete_vector_store(
|
|||
|
||||
try:
|
||||
_reject_config_vector_store_id(data.vector_store_id)
|
||||
|
||||
# Check if vector store exists in database or in-memory registry
|
||||
db_vector_store_exists = False
|
||||
memory_vector_store_exists = False
|
||||
vector_store_to_check = None
|
||||
|
||||
existing_vector_store: Final = await _vector_store_table(prisma_client).find_unique(
|
||||
where={"vector_store_id": data.vector_store_id}
|
||||
vector_store, database_exists, memory_exists = await _vector_store_delete_target(
|
||||
data.vector_store_id,
|
||||
prisma_client,
|
||||
)
|
||||
if existing_vector_store is not None:
|
||||
db_vector_store_exists = True
|
||||
vector_store_to_check = _row_to_vector_store(existing_vector_store)
|
||||
|
||||
# Check in-memory registry
|
||||
if litellm.vector_store_registry is not None:
|
||||
memory_vector_store: Final = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
|
||||
vector_store_id=data.vector_store_id
|
||||
)
|
||||
if memory_vector_store is not None:
|
||||
memory_vector_store_exists = True
|
||||
if vector_store_to_check is None:
|
||||
vector_store_to_check = memory_vector_store
|
||||
|
||||
# If not found in either location, raise 404
|
||||
if not db_vector_store_exists and not memory_vector_store_exists:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Vector store with ID {data.vector_store_id} not found",
|
||||
)
|
||||
|
||||
# Check access control
|
||||
if vector_store_to_check and not await _check_vector_store_access(vector_store_to_check, user_api_key_dict):
|
||||
if not await _check_vector_store_access(vector_store, user_api_key_dict):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Access denied: You do not have permission to delete this vector store",
|
||||
)
|
||||
|
||||
# Delete from database if exists
|
||||
if db_vector_store_exists:
|
||||
if database_exists:
|
||||
await _vector_store_table(prisma_client).delete(where={"vector_store_id": data.vector_store_id})
|
||||
|
||||
# Delete from in-memory registry if exists
|
||||
if memory_vector_store_exists and litellm.vector_store_registry is not None:
|
||||
litellm.vector_store_registry.delete_vector_store_from_registry(vector_store_id=data.vector_store_id)
|
||||
if memory_exists:
|
||||
registry: Final = litellm.vector_store_registry
|
||||
if registry is not None:
|
||||
registry.delete_vector_store_from_registry(vector_store_id=data.vector_store_id)
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
|
|
@ -654,7 +632,7 @@ async def update_vector_store(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
existing_litellm_params: Final = _litellm_params_dict(existing_vector_store.get("litellm_params"))
|
||||
existing_litellm_params: Final = deserialize_litellm_params(existing_vector_store.get("litellm_params"))
|
||||
effective_provider: Final = update_data.get("custom_llm_provider") or existing_vector_store.get(
|
||||
"custom_llm_provider"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from types import MappingProxyType
|
|||
from typing import Final, Literal
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -21,6 +22,7 @@ from litellm.types.utils import LlmProviders
|
|||
from litellm.types.vector_stores import (
|
||||
MILVUS_ADMIN_CONFIGURED_CONNECTION,
|
||||
LiteLLM_ManagedVectorStore,
|
||||
VectorStoreIndexEndpoints,
|
||||
)
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
|
|
@ -36,6 +38,7 @@ MILVUS_MANAGED_CONFIGURATION_FIELDS: Final = frozenset(
|
|||
"milvus_text_field",
|
||||
}
|
||||
)
|
||||
_MANAGED_VECTOR_STORE_ADAPTER: Final = TypeAdapter(LiteLLM_ManagedVectorStore)
|
||||
|
||||
|
||||
def _normalize_litellm_params(
|
||||
|
|
@ -43,7 +46,7 @@ def _normalize_litellm_params(
|
|||
) -> LiteLLM_ManagedVectorStore:
|
||||
litellm_params: Final = vector_store.get("litellm_params")
|
||||
if isinstance(litellm_params, str):
|
||||
normalized: Final = LiteLLM_ManagedVectorStore(**dict(vector_store))
|
||||
normalized: Final = _MANAGED_VECTOR_STORE_ADAPTER.validate_python(vector_store)
|
||||
try:
|
||||
parsed: Final = json.loads(litellm_params)
|
||||
normalized["litellm_params"] = parsed if isinstance(parsed, dict) else {}
|
||||
|
|
@ -128,7 +131,7 @@ def prepare_milvus_connection_for_persistence(
|
|||
litellm_credential_name: object | None = None,
|
||||
existing_litellm_credential_name: object | None = None,
|
||||
litellm_credential_name_supplied: bool = False,
|
||||
) -> dict[str, Any]: # mutable-ok: persistence requires a serializable effective-connection dict
|
||||
) -> dict[str, object]: # mutable-ok: persistence requires a serializable effective-connection dict
|
||||
existing: Final = existing_litellm_params if isinstance(existing_litellm_params, dict) else MappingProxyType({})
|
||||
supplied: Final = litellm_params if isinstance(litellm_params, dict) else MappingProxyType({})
|
||||
effective: Final = { # mutable-ok: the validated connection must be JSON-serializable for database persistence
|
||||
|
|
@ -139,6 +142,11 @@ def prepare_milvus_connection_for_persistence(
|
|||
}
|
||||
previous_is_grpc: Final = is_milvus_grpc_connection(existing_custom_llm_provider, existing)
|
||||
effective_is_grpc: Final = is_milvus_grpc_connection(custom_llm_provider, effective)
|
||||
if not previous_is_grpc and not effective_is_grpc:
|
||||
return ( # mutable-ok: persistence requires an isolated JSON-serializable dict
|
||||
dict(supplied) if isinstance(litellm_params, dict) else dict(existing)
|
||||
)
|
||||
|
||||
is_create: Final = existing_custom_llm_provider is None
|
||||
provider_changed: Final = not is_create and custom_llm_provider != existing_custom_llm_provider
|
||||
managed_configuration_changed: Final = any(
|
||||
|
|
@ -150,10 +158,8 @@ def prepare_milvus_connection_for_persistence(
|
|||
missing_marker: Final = effective_is_grpc and existing.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is not True
|
||||
|
||||
if (
|
||||
(previous_is_grpc or effective_is_grpc)
|
||||
and (is_create or provider_changed or managed_configuration_changed or credential_changed or missing_marker)
|
||||
and not _is_proxy_admin(user_api_key_dict)
|
||||
):
|
||||
is_create or provider_changed or managed_configuration_changed or credential_changed or missing_marker
|
||||
) and not _is_proxy_admin(user_api_key_dict):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Only proxy admins can configure vector store connections. Contact your LiteLLM administrator.",
|
||||
|
|
@ -516,6 +522,28 @@ def check_vector_store_permission(
|
|||
return False
|
||||
|
||||
|
||||
def _index_lifecycle_operation(request_method: str) -> Literal["create", "delete", "update"]:
|
||||
if request_method == "DELETE":
|
||||
return "delete"
|
||||
if request_method in ("PUT", "PATCH"):
|
||||
return "update"
|
||||
return "create"
|
||||
|
||||
|
||||
def _matching_index_permission(
|
||||
endpoints: VectorStoreIndexEndpoints,
|
||||
request_method: str,
|
||||
request_path: str,
|
||||
) -> Literal["read", "write"] | None:
|
||||
if any(
|
||||
request_method == method and _does_endpoint_match(path, request_path) for method, path in endpoints["write"]
|
||||
):
|
||||
return "write"
|
||||
if any(request_method == method and _does_endpoint_match(path, request_path) for method, path in endpoints["read"]):
|
||||
return "read"
|
||||
return None
|
||||
|
||||
|
||||
def is_allowed_to_call_vector_store_endpoint(
|
||||
provider: LlmProviders,
|
||||
index_name: str,
|
||||
|
|
@ -554,32 +582,17 @@ def is_allowed_to_call_vector_store_endpoint(
|
|||
request_path=request_route,
|
||||
index_name=index_name,
|
||||
):
|
||||
operation_label: Literal["create", "delete", "update"] = "create"
|
||||
if request.method == "DELETE":
|
||||
operation_label = "delete"
|
||||
elif request.method in ("PUT", "PATCH"):
|
||||
operation_label = "update"
|
||||
assert_proxy_admin_for_vector_store_index_management(
|
||||
user_api_key_dict,
|
||||
operation=operation_label,
|
||||
operation=_index_lifecycle_operation(request.method),
|
||||
)
|
||||
return True
|
||||
|
||||
# Writes are classified before reads so a path matching both patterns
|
||||
# requires the stronger grant (e.g. the azure batch write on an index
|
||||
# named "analyze*" also contains the "/analyze" read fragment)
|
||||
permission_type = None
|
||||
for endpoint in provider_vector_store_endpoints["write"]:
|
||||
if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route):
|
||||
permission_type = "write"
|
||||
break
|
||||
|
||||
if permission_type is None:
|
||||
for endpoint in provider_vector_store_endpoints["read"]:
|
||||
if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route):
|
||||
permission_type = "read"
|
||||
break
|
||||
|
||||
permission_type: Final = _matching_index_permission(
|
||||
provider_vector_store_endpoints,
|
||||
request.method,
|
||||
request_route,
|
||||
)
|
||||
if permission_type is None:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
|
|
|
|||
|
|
@ -8643,8 +8643,12 @@ class Router:
|
|||
if ptu_error is not None and is_ptu_cost_attribution_enabled():
|
||||
raise ValueError(ptu_error)
|
||||
zeroed_pricing: Final = zeroed_ptu_pricing(_model_info, _litellm_params) if config_sourced else None
|
||||
litellm_params: Final = LiteLLM_Params.model_validate(
|
||||
_litellm_params if zeroed_pricing is None else MappingProxyType({**_litellm_params, **zeroed_pricing})
|
||||
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(
|
||||
**(
|
||||
_litellm_params
|
||||
if zeroed_pricing is None
|
||||
else MappingProxyType({**_litellm_params, **zeroed_pricing})
|
||||
)
|
||||
)
|
||||
warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params)
|
||||
deployment = Deployment(
|
||||
|
|
|
|||
|
|
@ -64,7 +64,7 @@ class VectorStoreUpdateRequest(BaseModel):
|
|||
vector_store_description: str | None = None
|
||||
vector_store_metadata: dict | None = None
|
||||
litellm_credential_name: str | None = None
|
||||
litellm_params: dict[str, object] | None = None
|
||||
litellm_params: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
class VectorStoreDeleteRequest(BaseModel):
|
||||
|
|
@ -103,9 +103,9 @@ class VectorStoreSearchResponse(TypedDict, total=False):
|
|||
class VectorStoreSearchOptionalRequestParams(TypedDict, total=False):
|
||||
"""TypedDict for Optional parameters supported by the vector store search API."""
|
||||
|
||||
filters: dict[str, object] | None
|
||||
filters: Mapping[str, object] | None
|
||||
max_num_results: int | None
|
||||
ranking_options: dict[str, object] | None
|
||||
ranking_options: Mapping[str, object] | None
|
||||
rewrite_query: bool | None
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -8908,11 +8908,27 @@ class ProviderConfigManager:
|
|||
return BedrockVectorStore.get_initialized_custom_logger()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _get_milvus_vector_stores_config(
|
||||
transport: Literal["rest", "grpc"] | None,
|
||||
) -> BaseVectorStoreConfig:
|
||||
if transport == "grpc":
|
||||
from litellm.llms.milvus.vector_stores.grpc_transformation import (
|
||||
MilvusGRPCVectorStoreConfig,
|
||||
)
|
||||
|
||||
return MilvusGRPCVectorStoreConfig()
|
||||
from litellm.llms.milvus.vector_stores.transformation import (
|
||||
MilvusVectorStoreConfig,
|
||||
)
|
||||
|
||||
return MilvusVectorStoreConfig()
|
||||
|
||||
@staticmethod
|
||||
def get_provider_vector_stores_config(
|
||||
provider: LlmProviders,
|
||||
api_type: str | None = None,
|
||||
transport: object | None = None,
|
||||
transport: Literal["rest", "grpc"] | None = None,
|
||||
) -> BaseVectorStoreConfig | None:
|
||||
"""
|
||||
v2 vector store config, use this for new vector store integrations
|
||||
|
|
@ -8961,17 +8977,7 @@ class ProviderConfigManager:
|
|||
|
||||
return AzureAIVectorStoreConfig()
|
||||
elif litellm.LlmProviders.MILVUS == provider:
|
||||
if transport == "grpc":
|
||||
from litellm.llms.milvus.vector_stores.grpc_transformation import (
|
||||
MilvusGRPCVectorStoreConfig,
|
||||
)
|
||||
|
||||
return MilvusGRPCVectorStoreConfig()
|
||||
from litellm.llms.milvus.vector_stores.transformation import (
|
||||
MilvusVectorStoreConfig,
|
||||
)
|
||||
|
||||
return MilvusVectorStoreConfig()
|
||||
return ProviderConfigManager._get_milvus_vector_stores_config(transport)
|
||||
elif litellm.LlmProviders.GEMINI == provider:
|
||||
from litellm.llms.gemini.vector_stores.transformation import (
|
||||
GeminiVectorStoreConfig,
|
||||
|
|
|
|||
|
|
@ -33,18 +33,19 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
PrismaClient = Any
|
||||
|
||||
_LITELLM_PARAMS_ADAPTER: Final = TypeAdapter(dict[str, object] | None)
|
||||
_LITELLM_PARAMS_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_MANAGED_VECTOR_STORE_ADAPTER: Final = TypeAdapter(LiteLLM_ManagedVectorStore)
|
||||
|
||||
|
||||
def _deserialize_litellm_params(
|
||||
def deserialize_litellm_params(
|
||||
value: object,
|
||||
) -> dict[str, object] | None: # mutable-ok: managed vector store rows expose JSON objects as dicts
|
||||
) -> dict[str, object]: # mutable-ok: managed vector store rows expose JSON objects as dicts
|
||||
try:
|
||||
if isinstance(value, str):
|
||||
return _LITELLM_PARAMS_ADAPTER.validate_json(value)
|
||||
return _LITELLM_PARAMS_ADAPTER.validate_python(value)
|
||||
except ValidationError:
|
||||
return {}
|
||||
return {} # mutable-ok: managed vector store rows expose JSON objects as dicts
|
||||
|
||||
|
||||
class VectorStoreIndexRegistry:
|
||||
|
|
@ -276,6 +277,63 @@ class VectorStoreRegistry:
|
|||
return vector_store
|
||||
return None
|
||||
|
||||
def _cached_vector_store(self, vector_store_id: str) -> LiteLLM_ManagedVectorStore | None:
|
||||
return next(
|
||||
(
|
||||
vector_store
|
||||
for vector_store in self.vector_stores
|
||||
if vector_store.get("vector_store_id") == vector_store_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
async def _verified_cached_vector_store(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
vector_store: LiteLLM_ManagedVectorStore | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
) -> LiteLLM_ManagedVectorStore | None:
|
||||
if vector_store is None or prisma_client is None or vector_store_id in self.config_vector_store_ids:
|
||||
return vector_store
|
||||
try:
|
||||
database_store: Final = await ManagedVectorStoresRepository(prisma_client).table.find_unique(
|
||||
where=cast( # cast-ok: every value is already an object, only the popped id is stub-untyped
|
||||
"Mapping[str, object]", {"vector_store_id": vector_store_id}
|
||||
)
|
||||
)
|
||||
except Exception as error:
|
||||
verbose_logger.debug("Error verifying vector store %s in database: %s", vector_store_id, error)
|
||||
return vector_store
|
||||
if database_store is not None:
|
||||
return vector_store
|
||||
verbose_logger.debug(
|
||||
"Vector store %s found in memory but deleted from database, removing from cache",
|
||||
vector_store_id,
|
||||
)
|
||||
self.delete_vector_store_from_registry(vector_store_id=vector_store_id)
|
||||
return None
|
||||
|
||||
async def _resolve_vector_store(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
prisma_client: PrismaClient | None,
|
||||
) -> LiteLLM_ManagedVectorStore | None:
|
||||
cached: Final = await self._verified_cached_vector_store(
|
||||
vector_store_id,
|
||||
self._cached_vector_store(vector_store_id),
|
||||
prisma_client,
|
||||
)
|
||||
if cached is not None or prisma_client is None:
|
||||
return cached
|
||||
try:
|
||||
return await self.get_litellm_managed_vector_store_from_registry_or_db(
|
||||
vector_store_id=vector_store_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
except Exception as error:
|
||||
verbose_logger.debug("Error fetching vector store %s from database: %s", vector_store_id, error)
|
||||
return None
|
||||
|
||||
def pop_vector_stores_to_run(
|
||||
self, non_default_params: dict, tools: list[dict] | None = None
|
||||
) -> list[LiteLLM_ManagedVectorStore]:
|
||||
|
|
@ -348,47 +406,7 @@ class VectorStoreRegistry:
|
|||
vector_stores_to_run: Final[list[LiteLLM_ManagedVectorStore]] = []
|
||||
|
||||
for vector_store_id in vector_store_ids:
|
||||
vector_store = None
|
||||
|
||||
# First check in-memory registry
|
||||
for vs in self.vector_stores:
|
||||
if vs.get("vector_store_id") == vector_store_id:
|
||||
vector_store = vs
|
||||
break
|
||||
|
||||
# Verify vector store still exists in database (if we have DB access)
|
||||
# This ensures deleted vector stores are removed from cache
|
||||
if (
|
||||
vector_store is not None
|
||||
and prisma_client is not None
|
||||
and vector_store_id not in self.config_vector_store_ids
|
||||
):
|
||||
try:
|
||||
# Check if it still exists in database
|
||||
db_vector_store = await ManagedVectorStoresRepository(prisma_client).table.find_unique(
|
||||
where=cast( # cast-ok: every value is already an object, only the popped id is stub-untyped
|
||||
"Mapping[str, object]", {"vector_store_id": vector_store_id}
|
||||
)
|
||||
)
|
||||
if db_vector_store is None:
|
||||
# Vector store was deleted from database, remove from cache
|
||||
verbose_logger.debug(
|
||||
"Vector store %s found in memory but deleted from database, removing from cache",
|
||||
vector_store_id,
|
||||
)
|
||||
self.delete_vector_store_from_registry(vector_store_id=vector_store_id)
|
||||
vector_store = None
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Error verifying vector store %s in database: %s", vector_store_id, e)
|
||||
|
||||
# Fall back to database if not found in memory (or was deleted)
|
||||
if vector_store is None and prisma_client is not None:
|
||||
try:
|
||||
vector_store = await self.get_litellm_managed_vector_store_from_registry_or_db(
|
||||
vector_store_id=vector_store_id, prisma_client=prisma_client
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Error fetching vector store %s from database: %s", vector_store_id, e)
|
||||
vector_store = await self._resolve_vector_store(vector_store_id, prisma_client)
|
||||
|
||||
if vector_store is not None:
|
||||
# Create a copy to avoid modifying the registry
|
||||
|
|
@ -426,7 +444,9 @@ class VectorStoreRegistry:
|
|||
# cast to VectorStoreConfig
|
||||
litellm_vector_store_config = LiteLLM_VectorStoreConfig(**vector_store_config)
|
||||
vector_store_name = litellm_vector_store_config.get("vector_store_name")
|
||||
vector_store_litellm_params: dict[str, Any] = dict(litellm_vector_store_config.get("litellm_params") or {})
|
||||
vector_store_litellm_params = dict( # mutable-ok: config trust is marked on an isolated copy
|
||||
_LITELLM_PARAMS_ADAPTER.validate_python(litellm_vector_store_config.get("litellm_params") or {})
|
||||
)
|
||||
|
||||
vector_store_id = vector_store_litellm_params.get("vector_store_id")
|
||||
if not isinstance(vector_store_id, str) or not vector_store_id:
|
||||
|
|
@ -442,15 +462,17 @@ class VectorStoreRegistry:
|
|||
if custom_llm_provider == "milvus" and vector_store_litellm_params.get("milvus_transport") == "grpc":
|
||||
vector_store_litellm_params[MILVUS_ADMIN_CONFIGURED_CONNECTION] = True
|
||||
|
||||
litellm_managed_vector_store = LiteLLM_ManagedVectorStore(
|
||||
vector_store_id=vector_store_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=vector_store_litellm_params,
|
||||
vector_store_name=vector_store_name,
|
||||
vector_store_description=vector_store_litellm_params.get("vector_store_description"),
|
||||
vector_store_metadata=vector_store_litellm_params.get("vector_store_metadata"),
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
litellm_managed_vector_store = _MANAGED_VECTOR_STORE_ADAPTER.validate_python(
|
||||
{ # mutable-ok: Pydantic validates the config mapping into a managed vector store
|
||||
"vector_store_id": vector_store_id,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"litellm_params": vector_store_litellm_params,
|
||||
"vector_store_name": vector_store_name,
|
||||
"vector_store_description": vector_store_litellm_params.get("vector_store_description"),
|
||||
"vector_store_metadata": vector_store_litellm_params.get("vector_store_metadata"),
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
}
|
||||
)
|
||||
self.vector_stores.append(litellm_managed_vector_store)
|
||||
self.config_vector_store_ids = self.config_vector_store_ids.union((vector_store_id,))
|
||||
|
|
@ -530,7 +552,7 @@ class VectorStoreRegistry:
|
|||
)
|
||||
for vector_store in _vector_stores_from_db:
|
||||
_dict_vector_store = dict(vector_store)
|
||||
_dict_vector_store["litellm_params"] = _deserialize_litellm_params(
|
||||
_dict_vector_store["litellm_params"] = deserialize_litellm_params(
|
||||
_dict_vector_store.get("litellm_params")
|
||||
)
|
||||
_litellm_managed_vector_store = LiteLLM_ManagedVectorStore(**_dict_vector_store)
|
||||
|
|
|
|||
|
|
@ -21,6 +21,6 @@
|
|||
"limit": 117
|
||||
},
|
||||
"TQ008": {
|
||||
"limit": 10993
|
||||
"limit": 10984
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1056,7 +1056,7 @@ class TestNumericFormFields:
|
|||
read_only_not_required: ReadOnly[NotRequired[int]]
|
||||
read_only_required: ReadOnly[Required[float]]
|
||||
|
||||
assert dict(numeric_form_fields(get_type_hints(Schema, include_extras=True))) == {
|
||||
assert dict(numeric_form_fields(get_type_hints(Schema))) == {
|
||||
"plain": int,
|
||||
"optional": int,
|
||||
"piped": int,
|
||||
|
|
|
|||
|
|
@ -1017,6 +1017,18 @@ def test_admin_persistence_strips_forged_marker_and_adds_server_marker():
|
|||
assert params[MILVUS_ADMIN_CONFIGURED_CONNECTION] is True
|
||||
|
||||
|
||||
def test_non_grpc_connection_update_keeps_replacement_semantics():
|
||||
params = prepare_milvus_connection_for_persistence(
|
||||
custom_llm_provider="openai",
|
||||
litellm_params={"api_key": "new-key"},
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
existing_custom_llm_provider="openai",
|
||||
existing_litellm_params={"api_key": "old-key", "api_base": "https://old.example"},
|
||||
)
|
||||
|
||||
assert params == {"api_key": "new-key"}
|
||||
|
||||
|
||||
class TestCheckVectorStorePermission:
|
||||
"""Test suite for check_vector_store_permission function."""
|
||||
|
||||
|
|
|
|||
|
|
@ -1,14 +1,8 @@
|
|||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22181
|
||||
"limit": 22165
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26745
|
||||
"limit": 26733
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 261
|
||||
|
|
@ -27,7 +27,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16464
|
||||
"limit": 16442
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5506
|
||||
|
|
|
|||
|
|
@ -30,14 +30,12 @@ interface PluginModeContextValue {
|
|||
activePlugin: Plugin | null;
|
||||
}
|
||||
|
||||
const defaultPluginModeContext: PluginModeContextValue = {
|
||||
const PluginModeContext = createContext<PluginModeContextValue>({
|
||||
mode: "ai-gateway",
|
||||
setMode: () => {},
|
||||
plugins: [],
|
||||
activePlugin: null,
|
||||
};
|
||||
|
||||
const PluginModeContext = createContext(defaultPluginModeContext);
|
||||
});
|
||||
|
||||
const STORAGE_KEY = "litellm_plugin_mode";
|
||||
const pluginApiClient = createApiClient({ getBaseUrl: () => getProxyBaseUrl() ?? "" });
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue