mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(vector-stores): admin-gate every connection change on update and refresh registry rows from the DB sync
Any update that changes a managed vector store's provider, credential name, or litellm_params now needs a proxy admin, whatever the provider. Before this, only Milvus gRPC connections were gated, so a key holding the update route could point an OpenAI, Bedrock, Azure AI Search, or REST Milvus store at another host while keeping the stored key The periodic DB sync now replaces registry entries that already exist instead of skipping them, so an admin re-save on one worker reaches every worker within one config reload interval
This commit is contained in:
parent
f950274f93
commit
43a77d7076
5 changed files with 168 additions and 38 deletions
|
|
@ -48,12 +48,20 @@ def _targets_milvus(custom_llm_provider: object, litellm_params: object) -> bool
|
|||
|
||||
def _is_grpc_connection(custom_llm_provider: object, litellm_params: object) -> bool:
|
||||
return (
|
||||
isinstance(litellm_params, dict)
|
||||
isinstance(litellm_params, Mapping)
|
||||
and _targets_milvus(custom_llm_provider, litellm_params)
|
||||
and litellm_params.get("milvus_transport") == "grpc"
|
||||
)
|
||||
|
||||
|
||||
def _connection_fields(litellm_params: object) -> Mapping[str, object]:
|
||||
if not isinstance(litellm_params, Mapping):
|
||||
return MappingProxyType({})
|
||||
return MappingProxyType(
|
||||
{key: value for key, value in litellm_params.items() if key != MILVUS_ADMIN_CONFIGURED_CONNECTION}
|
||||
)
|
||||
|
||||
|
||||
def connection_rejection(
|
||||
custom_llm_provider: object,
|
||||
litellm_params: object,
|
||||
|
|
@ -64,7 +72,7 @@ def connection_rejection(
|
|||
if not _is_grpc_connection(custom_llm_provider, litellm_params):
|
||||
return None
|
||||
if managed:
|
||||
if isinstance(litellm_params, dict) and litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True:
|
||||
if isinstance(litellm_params, Mapping) and litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True:
|
||||
return None
|
||||
return MilvusConnectionRejection.ADMIN_SAVE_REQUIRED
|
||||
if is_proxy_admin:
|
||||
|
|
@ -82,44 +90,32 @@ def prepare_connection_for_persistence(
|
|||
litellm_credential_name: object | None = None,
|
||||
existing_litellm_credential_name: object | None = None,
|
||||
litellm_credential_name_supplied: bool = False,
|
||||
) -> dict[str, object] | MilvusConnectionRejection: # mutable-ok: persistence requires a serializable 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
|
||||
key: value
|
||||
for params in (existing, supplied)
|
||||
for key, value in params.items()
|
||||
if key != MILVUS_ADMIN_CONFIGURED_CONNECTION
|
||||
}
|
||||
) -> Mapping[str, object] | MilvusConnectionRejection:
|
||||
existing: Final = _connection_fields(existing_litellm_params)
|
||||
supplied: Final = _connection_fields(litellm_params)
|
||||
merged: Final = MappingProxyType({**existing, **supplied})
|
||||
previous_is_grpc: Final = _is_grpc_connection(existing_custom_llm_provider, existing)
|
||||
effective_is_grpc: Final = _is_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
|
||||
key: value
|
||||
for key, value in (supplied if isinstance(litellm_params, dict) else existing).items()
|
||||
if key != MILVUS_ADMIN_CONFIGURED_CONNECTION
|
||||
}
|
||||
|
||||
effective_is_grpc: Final = _is_grpc_connection(custom_llm_provider, merged)
|
||||
effective: Final = (
|
||||
merged
|
||||
if previous_is_grpc or effective_is_grpc
|
||||
else supplied
|
||||
if isinstance(litellm_params, Mapping)
|
||||
else 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(
|
||||
existing.get(field) != effective.get(field) for field in MILVUS_MANAGED_CONFIGURATION_FIELDS
|
||||
)
|
||||
credential_changed: Final = litellm_credential_name_supplied and (
|
||||
litellm_credential_name != existing_litellm_credential_name
|
||||
)
|
||||
missing_marker: Final = effective_is_grpc and existing.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is not True
|
||||
|
||||
if (
|
||||
is_create or provider_changed or managed_configuration_changed or credential_changed or missing_marker
|
||||
) and not is_proxy_admin:
|
||||
return MilvusConnectionRejection.ADMIN_REQUIRED
|
||||
|
||||
return (
|
||||
{**effective, MILVUS_ADMIN_CONFIGURED_CONNECTION: True} # mutable-ok: persisted JSON carries the server marker
|
||||
if effective_is_grpc
|
||||
else effective
|
||||
connection_changed: Final = not is_create and (provider_changed or credential_changed or effective != existing)
|
||||
missing_marker: Final = effective_is_grpc and (
|
||||
not isinstance(existing_litellm_params, Mapping)
|
||||
or existing_litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is not True
|
||||
)
|
||||
if (connection_changed or missing_marker) and not is_proxy_admin:
|
||||
return MilvusConnectionRejection.ADMIN_REQUIRED
|
||||
return MappingProxyType({**effective, MILVUS_ADMIN_CONFIGURED_CONNECTION: True}) if effective_is_grpc else effective
|
||||
|
||||
|
||||
def managed_connection_fields(custom_llm_provider: object, litellm_params: object) -> frozenset[str]:
|
||||
|
|
|
|||
|
|
@ -7763,7 +7763,10 @@ class ProxyConfig:
|
|||
litellm.vector_store_registry = VectorStoreRegistry(vector_stores=vector_stores)
|
||||
else:
|
||||
for vector_store in vector_stores:
|
||||
litellm.vector_store_registry.add_vector_store_to_registry(vector_store=vector_store)
|
||||
if (vector_store_id := vector_store.get("vector_store_id")) is not None:
|
||||
litellm.vector_store_registry.update_vector_store_in_registry(
|
||||
vector_store_id=vector_store_id, updated_data=vector_store
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.proxy_server.py::ProxyConfig:_init_vector_stores_in_db - %s", e
|
||||
|
|
|
|||
|
|
@ -92,7 +92,7 @@ def prepare_vector_store_connection_for_persistence(
|
|||
litellm_credential_name: object | None = None,
|
||||
existing_litellm_credential_name: object | None = None,
|
||||
litellm_credential_name_supplied: bool = False,
|
||||
) -> dict[str, object]: # mutable-ok: persistence requires a serializable effective-connection dict
|
||||
) -> Mapping[str, object]:
|
||||
result: Final = prepare_connection_for_persistence(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
|
|
|
|||
|
|
@ -12781,3 +12781,39 @@ async def test_update_general_settings_keeps_yaml_openai_websocket_passthrough()
|
|||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
assert ps.general_settings["enable_openai_websocket_passthrough"] is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_vector_stores_in_db_refreshes_a_store_already_in_the_registry(monkeypatch):
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
from litellm.types.vector_stores import MILVUS_ADMIN_CONFIGURED_CONNECTION, LiteLLM_ManagedVectorStore
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
|
||||
|
||||
stale: Final = LiteLLM_ManagedVectorStore(
|
||||
vector_store_id="managed-milvus",
|
||||
custom_llm_provider="milvus",
|
||||
litellm_params={"api_base": "http://old-milvus:19530", "milvus_transport": "grpc"},
|
||||
)
|
||||
registry: Final = VectorStoreRegistry(vector_stores=[stale])
|
||||
monkeypatch.setattr(litellm, "vector_store_registry", registry)
|
||||
saved_params: Final = {
|
||||
"api_base": "http://new-milvus:19530",
|
||||
"milvus_transport": "grpc",
|
||||
MILVUS_ADMIN_CONFIGURED_CONNECTION: True,
|
||||
}
|
||||
prisma_client: Final = MagicMock()
|
||||
prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"vector_store_id": "managed-milvus",
|
||||
"custom_llm_provider": "milvus",
|
||||
"litellm_params": json.dumps(saved_params),
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
await ProxyConfig()._init_vector_stores_in_db(prisma_client=prisma_client)
|
||||
|
||||
refreshed: Final = registry.get_litellm_managed_vector_store_from_registry(vector_store_id="managed-milvus")
|
||||
assert refreshed is not None
|
||||
assert refreshed["litellm_params"] == saved_params
|
||||
|
|
|
|||
|
|
@ -1078,11 +1078,11 @@ def test_nested_provider_cannot_bypass_milvus_grpc_registration_authorization(
|
|||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
def test_non_grpc_connection_update_keeps_replacement_semantics():
|
||||
def test_non_grpc_connection_update_keeps_replacement_semantics_for_admins():
|
||||
params = prepare_vector_store_connection_for_persistence(
|
||||
custom_llm_provider="openai",
|
||||
litellm_params={"api_key": "new-key"},
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
existing_custom_llm_provider="openai",
|
||||
existing_litellm_params={"api_key": "old-key", "api_base": "https://old.example"},
|
||||
)
|
||||
|
|
@ -1094,7 +1094,7 @@ def test_non_grpc_connection_update_drops_forged_admin_marker():
|
|||
params = prepare_vector_store_connection_for_persistence(
|
||||
custom_llm_provider="openai",
|
||||
litellm_params={"api_key": "new-key", MILVUS_ADMIN_CONFIGURED_CONNECTION: True},
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
existing_custom_llm_provider="openai",
|
||||
existing_litellm_params={"api_key": "old-key"},
|
||||
)
|
||||
|
|
@ -1102,6 +1102,59 @@ def test_non_grpc_connection_update_drops_forged_admin_marker():
|
|||
assert params == {"api_key": "new-key"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("custom_llm_provider", "litellm_params", "credential"),
|
||||
(
|
||||
("openai", {"api_key": "sk-real", "api_base": "https://attacker.example"}, None),
|
||||
("openai", {"api_key": "sk-real"}, None),
|
||||
("openai", {"api_key": "sk-real", "api_base": "https://old.example", "organization": "org-other"}, None),
|
||||
("bedrock", None, None),
|
||||
("openai", None, "someone-elses-credential"),
|
||||
),
|
||||
)
|
||||
def test_non_admin_cannot_change_a_non_grpc_connection_on_update(
|
||||
custom_llm_provider: str, litellm_params: dict[str, object] | None, credential: str | None
|
||||
) -> None:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
prepare_vector_store_connection_for_persistence(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
existing_custom_llm_provider="openai",
|
||||
existing_litellm_params={"api_key": "sk-real", "api_base": "https://old.example"},
|
||||
litellm_credential_name=credential,
|
||||
litellm_credential_name_supplied=credential is not None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Only proxy admins can configure vector store connections" in exc_info.value.detail
|
||||
|
||||
|
||||
@pytest.mark.parametrize("litellm_params", (None, {"api_key": "sk-real", "api_base": "https://old.example"}))
|
||||
def test_non_admin_update_leaving_the_non_grpc_connection_alone_is_allowed(
|
||||
litellm_params: dict[str, object] | None,
|
||||
) -> None:
|
||||
params = prepare_vector_store_connection_for_persistence(
|
||||
custom_llm_provider="openai",
|
||||
litellm_params=litellm_params,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
existing_custom_llm_provider="openai",
|
||||
existing_litellm_params={"api_key": "sk-real", "api_base": "https://old.example"},
|
||||
)
|
||||
|
||||
assert params == {"api_key": "sk-real", "api_base": "https://old.example"}
|
||||
|
||||
|
||||
def test_non_admin_can_still_create_a_non_grpc_store():
|
||||
params = prepare_vector_store_connection_for_persistence(
|
||||
custom_llm_provider="openai",
|
||||
litellm_params={"api_key": "sk-team-key"},
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
|
||||
assert params == {"api_key": "sk-team-key"}
|
||||
|
||||
|
||||
class TestCheckVectorStorePermission:
|
||||
"""Test suite for check_vector_store_permission function."""
|
||||
|
||||
|
|
@ -3225,6 +3278,48 @@ class TestUpdateVectorStoreAccessControlAndRedaction:
|
|||
assert exc_info.value.status_code == 403
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.update.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_admin_cannot_move_a_non_grpc_store_by_echoing_the_redacted_key(self):
|
||||
from litellm.constants import REDACTED_BY_LITELM_STRING
|
||||
from litellm.proxy.vector_store_endpoints.management_endpoints import update_vector_store
|
||||
from litellm.types.vector_stores import VectorStoreUpdateRequest
|
||||
|
||||
existing_row = MagicMock()
|
||||
existing_row.model_dump.return_value = {
|
||||
"vector_store_id": "vs_shared",
|
||||
"custom_llm_provider": "openai",
|
||||
"team_id": None,
|
||||
"litellm_params": {"api_key": "sk-real-openai-key-123", "api_base": "https://old.example"},
|
||||
}
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=existing_row)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: isolates the connection-authorization behavior from the feature entitlement gate
|
||||
"litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch( # test-quality-ok: injects the endpoint repository boundary with an admin-created shared store
|
||||
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
|
||||
),
|
||||
pytest.raises(HTTPException) as exc_info,
|
||||
):
|
||||
await update_vector_store(
|
||||
data=VectorStoreUpdateRequest(
|
||||
vector_store_id="vs_shared",
|
||||
litellm_params={"api_key": REDACTED_BY_LITELM_STRING, "api_base": "https://attacker.example"},
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="member",
|
||||
team_id="team-A",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Only proxy admins can configure vector store connections" in exc_info.value.detail
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.update.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_denied_when_caller_cannot_access_store(self):
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue