mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(vector-stores): block caller connection fields for nested milvus grpc stores
A managed store whose top-level custom_llm_provider is not milvus but whose litellm_params carry custom_llm_provider milvus with the grpc transport let a non-admin caller pass connection fields in the search body. The managed connection field check now looks at both the top-level and nested provider
This commit is contained in:
parent
7ffa8eb1d0
commit
b87186a717
3 changed files with 41 additions and 7 deletions
|
|
@ -39,13 +39,17 @@ def _normalize_provider(custom_llm_provider: object) -> str | None:
|
|||
return custom_llm_provider.split("/", 1)[0]
|
||||
|
||||
|
||||
def _targets_milvus(custom_llm_provider: object, litellm_params: object) -> bool:
|
||||
return _normalize_provider(custom_llm_provider) == "milvus" or (
|
||||
isinstance(litellm_params, Mapping)
|
||||
and _normalize_provider(litellm_params.get("custom_llm_provider")) == "milvus"
|
||||
)
|
||||
|
||||
|
||||
def _is_grpc_connection(custom_llm_provider: object, litellm_params: object) -> bool:
|
||||
return (
|
||||
isinstance(litellm_params, dict)
|
||||
and (
|
||||
_normalize_provider(custom_llm_provider) == "milvus"
|
||||
or _normalize_provider(litellm_params.get("custom_llm_provider")) == "milvus"
|
||||
)
|
||||
and _targets_milvus(custom_llm_provider, litellm_params)
|
||||
and litellm_params.get("milvus_transport") == "grpc"
|
||||
)
|
||||
|
||||
|
|
@ -116,9 +120,9 @@ def prepare_connection_for_persistence(
|
|||
)
|
||||
|
||||
|
||||
def managed_connection_fields(custom_llm_provider: object) -> frozenset[str]:
|
||||
def managed_connection_fields(custom_llm_provider: object, litellm_params: object) -> frozenset[str]:
|
||||
return frozenset((MILVUS_ADMIN_CONFIGURED_CONNECTION, "custom_llm_provider", "litellm_credential_name")) | (
|
||||
MILVUS_MANAGED_CONFIGURATION_FIELDS if _normalize_provider(custom_llm_provider) == "milvus" else frozenset()
|
||||
MILVUS_MANAGED_CONFIGURATION_FIELDS if _targets_milvus(custom_llm_provider, litellm_params) else frozenset()
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -103,7 +103,9 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
|
|||
vector_store=vector_store_to_run,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
blocked_fields: Final = managed_connection_fields(vector_store_to_run.get("custom_llm_provider"))
|
||||
blocked_fields: Final = managed_connection_fields(
|
||||
vector_store_to_run.get("custom_llm_provider"), vector_store_to_run.get("litellm_params")
|
||||
)
|
||||
managed_data: Final = build_request_data_from_managed_vector_store(vector_store_to_run)
|
||||
request_data: Final = {
|
||||
**{key: value for key, value in data.items() if key not in blocked_fields},
|
||||
|
|
|
|||
|
|
@ -764,6 +764,34 @@ async def test_managed_milvus_uses_only_persisted_connection_for_non_admin():
|
|||
assert managed_vector_store == original_store
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("blocked_key", ["api_base", "api_key", "milvus_db_name", "milvus_partition_names"])
|
||||
async def test_nested_milvus_grpc_store_drops_caller_connection_fields(blocked_key: str) -> None:
|
||||
managed_vector_store: LiteLLM_ManagedVectorStore = {
|
||||
"vector_store_id": "nested",
|
||||
"custom_llm_provider": "openai",
|
||||
"litellm_params": {
|
||||
"custom_llm_provider": "milvus",
|
||||
"milvus_transport": "grpc",
|
||||
"litellm_embedding_model": "team-embedding-alias",
|
||||
MILVUS_ADMIN_CONFIGURED_CONNECTION: True,
|
||||
},
|
||||
}
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = managed_vector_store
|
||||
|
||||
with patch.object( # test-quality-ok: the helper reads the process-wide registry directly
|
||||
litellm, "vector_store_registry", mock_registry
|
||||
):
|
||||
result = await _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
data={"query": "safe", blocked_key: "attacker-choice"},
|
||||
vector_store_id="nested",
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
|
||||
assert blocked_key not in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unmarked_managed_milvus_connection_requires_admin_resave():
|
||||
managed_vector_store: LiteLLM_ManagedVectorStore = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue