mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
chore(vector stores): address tenant guard followups
This commit is contained in:
parent
d3ab59e059
commit
363c0de6f7
5 changed files with 128 additions and 60 deletions
|
|
@ -3572,7 +3572,7 @@
|
|||
"/anthropic/{endpoint}": {
|
||||
"delete": {
|
||||
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
|
||||
"operationId": "anthropic_proxy_route_anthropic__endpoint__put",
|
||||
"operationId": "anthropic_proxy_route_anthropic__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -3616,7 +3616,7 @@
|
|||
},
|
||||
"get": {
|
||||
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
|
||||
"operationId": "anthropic_proxy_route_anthropic__endpoint__put",
|
||||
"operationId": "anthropic_proxy_route_anthropic__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -3660,7 +3660,7 @@
|
|||
},
|
||||
"patch": {
|
||||
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
|
||||
"operationId": "anthropic_proxy_route_anthropic__endpoint__put",
|
||||
"operationId": "anthropic_proxy_route_anthropic__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -3704,7 +3704,7 @@
|
|||
},
|
||||
"post": {
|
||||
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
|
||||
"operationId": "anthropic_proxy_route_anthropic__endpoint__put",
|
||||
"operationId": "anthropic_proxy_route_anthropic__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -3748,7 +3748,7 @@
|
|||
},
|
||||
"put": {
|
||||
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
|
||||
"operationId": "anthropic_proxy_route_anthropic__endpoint__put",
|
||||
"operationId": "anthropic_proxy_route_anthropic__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -13260,7 +13260,7 @@
|
|||
"/langfuse/{endpoint}": {
|
||||
"delete": {
|
||||
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
|
||||
"operationId": "langfuse_proxy_route_langfuse__endpoint__put",
|
||||
"operationId": "langfuse_proxy_route_langfuse__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -13299,7 +13299,7 @@
|
|||
},
|
||||
"get": {
|
||||
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
|
||||
"operationId": "langfuse_proxy_route_langfuse__endpoint__put",
|
||||
"operationId": "langfuse_proxy_route_langfuse__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -13338,7 +13338,7 @@
|
|||
},
|
||||
"patch": {
|
||||
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
|
||||
"operationId": "langfuse_proxy_route_langfuse__endpoint__put",
|
||||
"operationId": "langfuse_proxy_route_langfuse__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -13377,7 +13377,7 @@
|
|||
},
|
||||
"post": {
|
||||
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
|
||||
"operationId": "langfuse_proxy_route_langfuse__endpoint__put",
|
||||
"operationId": "langfuse_proxy_route_langfuse__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -13416,7 +13416,7 @@
|
|||
},
|
||||
"put": {
|
||||
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
|
||||
"operationId": "langfuse_proxy_route_langfuse__endpoint__put",
|
||||
"operationId": "langfuse_proxy_route_langfuse__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -26883,7 +26883,7 @@
|
|||
"/toolset/{toolset_name}/mcp": {
|
||||
"delete": {
|
||||
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -26922,7 +26922,7 @@
|
|||
},
|
||||
"get": {
|
||||
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -26961,7 +26961,7 @@
|
|||
},
|
||||
"head": {
|
||||
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -27000,7 +27000,7 @@
|
|||
},
|
||||
"options": {
|
||||
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -27039,7 +27039,7 @@
|
|||
},
|
||||
"patch": {
|
||||
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -27078,7 +27078,7 @@
|
|||
},
|
||||
"post": {
|
||||
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
@ -27117,7 +27117,7 @@
|
|||
},
|
||||
"put": {
|
||||
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put",
|
||||
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
|||
from litellm.proxy.utils import is_known_model
|
||||
from litellm.proxy.vector_store_endpoints.utils import (
|
||||
assert_user_can_access_vector_store,
|
||||
get_litellm_managed_vector_store,
|
||||
is_allowed_to_call_vector_store_endpoint,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
|
@ -1927,11 +1928,11 @@ async def vertex_discovery_proxy_route(
|
|||
"Extracted vector store ID from endpoint: %s", vector_store_id
|
||||
)
|
||||
|
||||
# Retrieve vector store credentials from the registry
|
||||
vector_store_credentials = (
|
||||
passthrough_endpoint_router.get_vector_store_credentials(
|
||||
vector_store_id=vector_store_id
|
||||
)
|
||||
# Retrieve LiteLLM-managed vector store credentials if the datastore id
|
||||
# is registered with LiteLLM. Unknown datastore ids keep the existing
|
||||
# direct Vertex pass-through behavior.
|
||||
vector_store_credentials = await get_litellm_managed_vector_store(
|
||||
vector_store_id=vector_store_id
|
||||
)
|
||||
|
||||
if vector_store_credentials:
|
||||
|
|
@ -1939,14 +1940,10 @@ async def vertex_discovery_proxy_route(
|
|||
"Found vector store credentials for ID: %s", vector_store_id
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.warning(
|
||||
verbose_proxy_logger.debug(
|
||||
"Vector store ID %s found in endpoint but no credentials found in registry",
|
||||
vector_store_id,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Access denied: You do not have permission to access this vector store",
|
||||
)
|
||||
|
||||
discovery_handler = get_vertex_pass_through_handler(call_type="discovery")
|
||||
return await _base_vertex_proxy_route(
|
||||
|
|
|
|||
|
|
@ -31,21 +31,27 @@ router = APIRouter()
|
|||
|
||||
def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]:
|
||||
vector_store_ids: set[str] = set()
|
||||
payload_stack = [payload]
|
||||
|
||||
if isinstance(payload, dict):
|
||||
for key, value in payload.items():
|
||||
if key == "vector_store_id":
|
||||
if not isinstance(value, str) or not value:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "vector_store_id must be a non-empty string"},
|
||||
)
|
||||
vector_store_ids.add(value)
|
||||
continue
|
||||
vector_store_ids.update(_collect_vector_store_ids_from_payload(value))
|
||||
elif isinstance(payload, list):
|
||||
for item in payload:
|
||||
vector_store_ids.update(_collect_vector_store_ids_from_payload(item))
|
||||
while payload_stack:
|
||||
current_payload = payload_stack.pop()
|
||||
|
||||
if isinstance(current_payload, dict):
|
||||
for key, value in current_payload.items():
|
||||
if key == "vector_store_id":
|
||||
if not isinstance(value, str) or not value:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "vector_store_id must be a non-empty string"
|
||||
},
|
||||
)
|
||||
vector_store_ids.add(value)
|
||||
continue
|
||||
if isinstance(value, (dict, list)):
|
||||
payload_stack.append(value)
|
||||
elif isinstance(current_payload, list):
|
||||
payload_stack.extend(current_payload)
|
||||
|
||||
return vector_store_ids
|
||||
|
||||
|
|
|
|||
|
|
@ -141,7 +141,7 @@ async def get_litellm_managed_vector_store(
|
|||
vector_store_id: str,
|
||||
) -> Optional[LiteLLM_ManagedVectorStore]:
|
||||
"""
|
||||
Resolve a LiteLLM-managed vector store from the registry or database.
|
||||
Resolve a LiteLLM-managed vector store from the registry or shared cache.
|
||||
|
||||
Provider-native vector store IDs will not be present in either location and
|
||||
return None, preserving direct provider behavior while still protecting
|
||||
|
|
@ -165,19 +165,31 @@ async def get_litellm_managed_vector_store(
|
|||
)
|
||||
|
||||
try:
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
get_managed_vector_store_rows_by_uuids,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
return None
|
||||
row = await prisma_client.db.litellm_managedvectorstorestable.find_unique(
|
||||
where={"vector_store_id": vector_store_id}
|
||||
rows = await get_managed_vector_store_rows_by_uuids(
|
||||
uuids=[vector_store_id],
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if row is None:
|
||||
if not rows:
|
||||
return None
|
||||
return _normalize_litellm_params(LiteLLM_ManagedVectorStore(**row.model_dump()))
|
||||
return _normalize_litellm_params(
|
||||
LiteLLM_ManagedVectorStore(**rows[0].model_dump())
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to resolve vector store id=%s from database: %s",
|
||||
"Failed to resolve vector store id=%s from shared cache: %s",
|
||||
vector_store_id,
|
||||
e,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import pytest
|
|||
from fastapi import HTTPException, Request, Response
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable, UserAPIKeyAuth
|
||||
|
||||
|
||||
def _mock_request() -> MagicMock:
|
||||
|
|
@ -288,7 +288,53 @@ async def test_vertex_discovery_denies_other_team_vector_store_credentials():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_discovery_denies_unregistered_vector_store_id():
|
||||
async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback():
|
||||
from litellm.proxy.vector_store_endpoints.utils import (
|
||||
get_litellm_managed_vector_store,
|
||||
)
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = None
|
||||
cache_helper = AsyncMock(
|
||||
return_value=[
|
||||
LiteLLM_ManagedVectorStoresTable(
|
||||
vector_store_id="vs_cached",
|
||||
custom_llm_provider="openai",
|
||||
vector_store_name=None,
|
||||
vector_store_description=None,
|
||||
vector_store_metadata=None,
|
||||
created_at=None,
|
||||
updated_at=None,
|
||||
litellm_credential_name=None,
|
||||
litellm_params={"api_base": "https://example.com"},
|
||||
team_id="team-a",
|
||||
user_id=None,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(litellm, "vector_store_registry", mock_registry),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids",
|
||||
new=cache_helper,
|
||||
),
|
||||
):
|
||||
vector_store = await get_litellm_managed_vector_store(
|
||||
vector_store_id="vs_cached"
|
||||
)
|
||||
|
||||
assert vector_store is not None
|
||||
assert vector_store["vector_store_id"] == "vs_cached"
|
||||
assert vector_store["team_id"] == "team-a"
|
||||
cache_helper.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_discovery_allows_unregistered_provider_native_datastore_id():
|
||||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
vertex_discovery_proxy_route,
|
||||
)
|
||||
|
|
@ -296,18 +342,25 @@ async def test_vertex_discovery_denies_unregistered_vector_store_id():
|
|||
request = _mock_request()
|
||||
request.method = "GET"
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_vector_store_credentials",
|
||||
return_value=None,
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_litellm_managed_vector_store",
|
||||
new=AsyncMock(return_value=None),
|
||||
) as mock_lookup,
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._base_vertex_proxy_route",
|
||||
new=AsyncMock(return_value={"ok": True}),
|
||||
) as mock_base_route,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await vertex_discovery_proxy_route(
|
||||
endpoint="projects/p/locations/us-central1/dataStores/vs_unknown",
|
||||
request=request,
|
||||
fastapi_response=Response(),
|
||||
)
|
||||
response = await vertex_discovery_proxy_route(
|
||||
endpoint="projects/p/locations/us-central1/dataStores/vs_unknown",
|
||||
request=request,
|
||||
fastapi_response=Response(),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert response == {"ok": True}
|
||||
mock_lookup.assert_awaited_once_with(vector_store_id="vs_unknown")
|
||||
assert mock_base_route.call_args.kwargs["router_credentials"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue