mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #36798 from BerriAI/litellm_azure_ai_docs_index_write_grant_rc
fix(azure_ai): recognize real Search doc endpoints so teams can read/write via passthrough
This commit is contained in:
commit
d74cb6de1b
5 changed files with 484 additions and 17 deletions
|
|
@ -37,9 +37,32 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM):
|
|||
super().__init__()
|
||||
|
||||
def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints:
|
||||
"""
|
||||
Every ``GET`` under ``/indexes/`` is a read: get details, stats, and the
|
||||
document reads (GET-form search, ``$count``, point lookup, and the
|
||||
GET forms of suggest and autocomplete).
|
||||
|
||||
``POST`` splits by endpoint. Search, suggest, autocomplete, and analyze
|
||||
are query endpoints, so they read; ``/docs/index`` is the batch endpoint
|
||||
carrying upload, merge, mergeOrUpload, and delete actions, so it writes.
|
||||
|
||||
Patterns stay literal rather than ``{placeholder}`` templates because the
|
||||
matcher falls back to the substring before a ``{``, which here is always
|
||||
``/indexes/``. The matcher is substring-based, so an index name may
|
||||
itself contain a read fragment (an index named ``analyze*`` puts
|
||||
``/analyze`` inside the batch-write path); writes are classified before
|
||||
reads, so such a path demands the write grant rather than being
|
||||
shadowed into a read.
|
||||
"""
|
||||
return {
|
||||
"read": [("GET", "/docs/search"), ("POST", "/docs/search")],
|
||||
"write": [("PUT", "/docs")],
|
||||
"read": [
|
||||
("GET", "/indexes/"),
|
||||
("POST", "/docs/search"),
|
||||
("POST", "/docs/suggest"),
|
||||
("POST", "/docs/autocomplete"),
|
||||
("POST", "/analyze"),
|
||||
],
|
||||
"write": [("POST", "/docs/index")],
|
||||
}
|
||||
|
||||
def get_auth_credentials(self, litellm_params: dict) -> BaseVectorStoreAuthCredentials:
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
|||
)
|
||||
from litellm.proxy.utils import is_known_model
|
||||
from litellm.proxy.vector_store_endpoints.utils import (
|
||||
assert_proxy_admin_for_vector_store_index_management,
|
||||
assert_user_can_access_vector_store,
|
||||
get_litellm_managed_vector_store,
|
||||
is_allowed_to_call_vector_store_endpoint,
|
||||
|
|
@ -1234,6 +1235,37 @@ async def assemblyai_proxy_route(
|
|||
return received_value
|
||||
|
||||
|
||||
def get_azure_ai_search_index_from_endpoint(endpoint: str) -> str | None:
|
||||
"""Return the index name in the ``/indexes/{name}`` position of an Azure AI
|
||||
Search passthrough path, or ``None`` when the path targets no index.
|
||||
|
||||
Only the segment immediately after ``indexes`` is the operable target. Any
|
||||
other segment (for example the trailing ``index`` in ``.../docs/index``) must
|
||||
never be treated as the index, otherwise a caller authorized on one index
|
||||
could have Azure apply the operation to a different index on the same service.
|
||||
"""
|
||||
segments: Final = endpoint.split("?", 1)[0].strip("/").split("/")
|
||||
for position, segment in enumerate(segments):
|
||||
if segment == "indexes" and position + 1 < len(segments):
|
||||
return segments[position + 1] or None
|
||||
return None
|
||||
|
||||
|
||||
def is_azure_ai_search_service_level_index_create(method: str, endpoint: str) -> bool:
|
||||
"""Return True for ``POST /indexes``, Azure AI Search's service-level index create.
|
||||
|
||||
No index name appears in that path, so ``get_azure_ai_search_index_from_endpoint``
|
||||
yields None and the managed-index branch can never claim the request. Without an
|
||||
explicit guard it reaches the generic Azure passthrough on the proxy's own
|
||||
credential, so a non-admin could create an index whenever ``AZURE_API_BASE``
|
||||
points at the Search service.
|
||||
"""
|
||||
if method != "POST":
|
||||
return False
|
||||
path: Final = endpoint.split("?", 1)[0].strip("/")
|
||||
return path == "indexes" or path.endswith("/indexes")
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/azure_ai/{endpoint:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
|
||||
|
|
@ -1259,10 +1291,15 @@ async def azure_proxy_route(
|
|||
"""
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
if is_azure_ai_search_service_level_index_create(method=request.method, endpoint=endpoint):
|
||||
assert_proxy_admin_for_vector_store_index_management(user_api_key_dict, operation="create")
|
||||
|
||||
parts: Final = endpoint.split(
|
||||
"/"
|
||||
) # azure model is in the url - e.g. https://{endpoint}/openai/deployments/{deployment-id}/completions?api-version=2024-10-21
|
||||
|
||||
search_index_name: Final = get_azure_ai_search_index_from_endpoint(endpoint)
|
||||
|
||||
if len(parts) > 1 and llm_router:
|
||||
for part in parts:
|
||||
# check if LLM MODEL
|
||||
|
|
@ -1271,9 +1308,9 @@ async def azure_proxy_route(
|
|||
)
|
||||
# check if vector store index
|
||||
is_vector_store_index = (
|
||||
(litellm.vector_store_index_registry.is_vector_store_index(vector_store_index_name=part))
|
||||
if litellm.vector_store_index_registry is not None
|
||||
else False
|
||||
part == search_index_name
|
||||
and litellm.vector_store_index_registry is not None
|
||||
and litellm.vector_store_index_registry.is_vector_store_index(vector_store_index_name=part)
|
||||
)
|
||||
|
||||
if is_router_model:
|
||||
|
|
|
|||
|
|
@ -86,7 +86,7 @@ def _is_vector_store_index_lifecycle_request(
|
|||
return True
|
||||
|
||||
# POST /indexes (create index at service level; no index name in path).
|
||||
normalized: Final = request_path.rstrip("/")
|
||||
normalized: Final = request_path.split("?", 1)[0].rstrip("/")
|
||||
if request_method == "POST" and normalized.endswith("/indexes"):
|
||||
return True
|
||||
|
||||
|
|
@ -387,17 +387,19 @@ def is_allowed_to_call_vector_store_endpoint(
|
|||
)
|
||||
return True
|
||||
|
||||
# Determine the permission type based on the request
|
||||
# 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["read"]:
|
||||
for endpoint in provider_vector_store_endpoints["write"]:
|
||||
if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route):
|
||||
permission_type = "read"
|
||||
permission_type = "write"
|
||||
break
|
||||
|
||||
if permission_type is None:
|
||||
for endpoint in provider_vector_store_endpoints["write"]:
|
||||
for endpoint in provider_vector_store_endpoints["read"]:
|
||||
if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route):
|
||||
permission_type = "write"
|
||||
permission_type = "read"
|
||||
break
|
||||
|
||||
if permission_type is None:
|
||||
|
|
@ -454,15 +456,15 @@ def is_allowed_to_call_vector_store_files_endpoint(
|
|||
request_route: Final = get_request_route(request)
|
||||
|
||||
permission_type: str | None = None
|
||||
for endpoint in provider_vector_store_endpoints.get("read", ()):
|
||||
for endpoint in provider_vector_store_endpoints.get("write", ()):
|
||||
if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route):
|
||||
permission_type = "read"
|
||||
permission_type = "write"
|
||||
break
|
||||
|
||||
if permission_type is None:
|
||||
for endpoint in provider_vector_store_endpoints.get("write", ()):
|
||||
for endpoint in provider_vector_store_endpoints.get("read", ()):
|
||||
if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route):
|
||||
permission_type = "write"
|
||||
permission_type = "read"
|
||||
break
|
||||
|
||||
if permission_type is None:
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import Request, Response
|
||||
from fastapi import HTTPException, Request, Response
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(
|
||||
|
|
@ -19,10 +19,13 @@ import litellm
|
|||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
BaseOpenAIPassThroughHandler,
|
||||
RouteChecks,
|
||||
azure_proxy_route,
|
||||
bedrock_llm_proxy_route,
|
||||
create_pass_through_route,
|
||||
cursor_proxy_route,
|
||||
get_azure_ai_search_index_from_endpoint,
|
||||
get_vertex_base_url,
|
||||
is_azure_ai_search_service_level_index_create,
|
||||
llm_passthrough_factory_proxy_route,
|
||||
milvus_proxy_route,
|
||||
mistral_proxy_route,
|
||||
|
|
@ -31,7 +34,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
|||
vertex_proxy_route,
|
||||
vllm_proxy_route,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
|
||||
|
||||
|
||||
|
|
@ -3249,3 +3252,221 @@ def test_is_passthrough_request_streaming_tolerates_non_object_bodies(request_bo
|
|||
)
|
||||
|
||||
assert is_passthrough_request_streaming(request_body) is expected
|
||||
|
||||
|
||||
class TestGetAzureAISearchIndexFromEndpoint:
|
||||
"""The operable index is only the segment right after ``indexes``.
|
||||
|
||||
A doc-write path ends in ``.../docs/index``; the trailing ``index`` must not
|
||||
be mistaken for the target, otherwise a caller could be authorized on one
|
||||
index while Azure applies the write to another.
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint, expected",
|
||||
[
|
||||
("indexes/my-index/docs/index", "my-index"),
|
||||
("indexes/my-index/docs/search", "my-index"),
|
||||
("indexes/my-index", "my-index"),
|
||||
("indexes/my-index?api-version=2024-07-01", "my-index"),
|
||||
("/indexes/my-index/docs/index", "my-index"),
|
||||
("indexes/victim/docs/index", "victim"),
|
||||
("openai/deployments/gpt-4o/chat/completions", None),
|
||||
("indexes", None),
|
||||
("indexes/", None),
|
||||
],
|
||||
)
|
||||
def test_extracts_positional_index_only(self, endpoint, expected):
|
||||
assert get_azure_ai_search_index_from_endpoint(endpoint) == expected
|
||||
|
||||
|
||||
class TestAzureProxyRouteCrossIndexAuthorization:
|
||||
"""Regression tests: the passthrough must authorize the index that the request
|
||||
actually targets (the ``/indexes/{name}`` segment), never a different segment
|
||||
that merely happens to match a managed index the caller can access.
|
||||
"""
|
||||
|
||||
def _request(self, method: str, path: str) -> MagicMock:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = method
|
||||
request.headers = {"content-type": "application/json"}
|
||||
request.url = MagicMock()
|
||||
request.url.path = path
|
||||
return request
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorizes_the_targeted_index(self):
|
||||
index_object = MagicMock()
|
||||
index_object.litellm_params.vector_store_name = "my-store"
|
||||
vector_store = {"litellm_params": {"api_base": "https://svc.search.windows.net"}}
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config"
|
||||
) as mock_get_config,
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint"
|
||||
) as mock_is_allowed,
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.assert_user_can_access_vector_store",
|
||||
new=AsyncMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler",
|
||||
new=AsyncMock(return_value=Response()),
|
||||
),
|
||||
patch.object(litellm, "vector_store_index_registry") as mock_index_registry,
|
||||
patch.object(litellm, "vector_store_registry") as mock_vector_registry,
|
||||
):
|
||||
mock_get_config.return_value.get_auth_credentials.return_value = {"headers": {"api-key": "k"}}
|
||||
mock_index_registry.is_vector_store_index.side_effect = lambda vector_store_index_name: (
|
||||
vector_store_index_name == "my-index"
|
||||
)
|
||||
mock_index_registry.get_vector_store_index_by_name.return_value = index_object
|
||||
mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = vector_store
|
||||
|
||||
await azure_proxy_route(
|
||||
endpoint="indexes/my-index/docs/index",
|
||||
request=self._request("POST", "/azure_ai/indexes/my-index/docs/index"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
|
||||
)
|
||||
|
||||
mock_is_allowed.assert_called_once()
|
||||
assert mock_is_allowed.call_args.kwargs["index_name"] == "my-index"
|
||||
mock_index_registry.get_vector_store_index_by_name.assert_called_once_with(
|
||||
vector_store_index_name="my-index"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_trailing_index_segment_does_not_authorize_a_different_index(self):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint"
|
||||
) as mock_is_allowed,
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str",
|
||||
return_value="https://azure-openai.example.com",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="azure-key",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler",
|
||||
new=AsyncMock(return_value=Response()),
|
||||
) as mock_handler,
|
||||
patch.object(litellm, "vector_store_index_registry") as mock_index_registry,
|
||||
):
|
||||
mock_index_registry.is_vector_store_index.side_effect = lambda vector_store_index_name: (
|
||||
vector_store_index_name == "index"
|
||||
)
|
||||
|
||||
await azure_proxy_route(
|
||||
endpoint="indexes/victim/docs/index",
|
||||
request=self._request("POST", "/azure_ai/indexes/victim/docs/index"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
|
||||
)
|
||||
|
||||
mock_is_allowed.assert_not_called()
|
||||
mock_handler.assert_awaited_once()
|
||||
assert mock_handler.await_args.kwargs["custom_llm_provider"] == litellm.LlmProviders.AZURE
|
||||
|
||||
|
||||
class TestAzureProxyRouteServiceLevelIndexCreate:
|
||||
"""``POST /indexes`` carries no index name, so the managed-index branch cannot
|
||||
claim it and it would otherwise reach the generic Azure passthrough on the
|
||||
proxy's own credential. The admin-only index management guard has to be
|
||||
enforced on the route itself, not just on the permission gate the route skips.
|
||||
"""
|
||||
|
||||
def _request(self, method: str, path: str) -> MagicMock:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = method
|
||||
request.headers = {"content-type": "application/json"}
|
||||
request.url = MagicMock()
|
||||
request.url.path = path
|
||||
return request
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"method, endpoint, expected",
|
||||
[
|
||||
("POST", "indexes", True),
|
||||
("POST", "indexes?api-version=2024-07-01", True),
|
||||
("POST", "/indexes/", True),
|
||||
("POST", "indexes/my-index", False),
|
||||
("POST", "indexes/my-index/docs/index", False),
|
||||
("GET", "indexes", False),
|
||||
("POST", "openai/deployments/gpt-4o/chat/completions", False),
|
||||
],
|
||||
)
|
||||
def test_recognizes_service_level_create(self, method, endpoint, expected):
|
||||
assert is_azure_ai_search_service_level_index_create(method=method, endpoint=endpoint) is expected
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_admin_cannot_create_an_index(self):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str",
|
||||
return_value="https://svc.search.windows.net",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler",
|
||||
new=AsyncMock(return_value=Response()),
|
||||
) as mock_handler,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await azure_proxy_route(
|
||||
endpoint="indexes?api-version=2024-07-01",
|
||||
request=self._request("POST", "/azure_ai/indexes"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
token="sk-team-token",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Only proxy admins can create" in exc_info.value.detail
|
||||
mock_handler.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_can_still_create_an_index(self):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str",
|
||||
return_value="https://svc.search.windows.net",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="azure-key",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler",
|
||||
new=AsyncMock(return_value=Response()),
|
||||
) as mock_handler,
|
||||
):
|
||||
await azure_proxy_route(
|
||||
endpoint="indexes?api-version=2024-07-01",
|
||||
request=self._request("POST", "/azure_ai/indexes"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
token="sk-admin-token",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
),
|
||||
)
|
||||
|
||||
mock_handler.assert_awaited_once()
|
||||
|
|
|
|||
|
|
@ -2928,3 +2928,187 @@ class TestUpdateVectorStoreAccessControlAndRedaction:
|
|||
params = response["vector_store"]["litellm_params"]
|
||||
assert params["api_key"] == REDACTED_BY_LITELM_STRING
|
||||
assert params["api_base"] == "https://api.openai.com/v1"
|
||||
|
||||
|
||||
class TestAzureAIDocumentWritePassthroughPermission:
|
||||
"""Regression tests for the Azure AI Search passthrough write mapping.
|
||||
|
||||
Azure's batch document write/merge/delete endpoint is
|
||||
``POST /indexes/{name}/docs/index``. A non-admin team holding a ``write``
|
||||
grant on the index must be allowed to call it, while index lifecycle
|
||||
(create / update / delete the index itself) stays proxy-admin only.
|
||||
|
||||
These exercise the real ``AzureAIVectorStoreConfig`` endpoint map on
|
||||
purpose (no mocked provider config), so reverting the map to the old
|
||||
``("PUT", "/docs")`` entry makes ``test_team_with_write_grant_can_upload``
|
||||
fail.
|
||||
"""
|
||||
|
||||
INDEX = "my-index"
|
||||
|
||||
READ_ROUTES = [
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/stats"),
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/docs"),
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/docs/$count"),
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/docs/seed-doc-1"),
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/docs/suggest"),
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/docs/autocomplete"),
|
||||
("POST", f"/azure_ai/indexes/{INDEX}/docs/suggest"),
|
||||
("POST", f"/azure_ai/indexes/{INDEX}/docs/autocomplete"),
|
||||
("POST", f"/azure_ai/indexes/{INDEX}/analyze"),
|
||||
]
|
||||
|
||||
def _request(self, method: str, path: str) -> MagicMock:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = method
|
||||
request.url.path = path
|
||||
return request
|
||||
|
||||
def _team_member(self, permissions: list) -> MagicMock:
|
||||
user = MagicMock(spec=UserAPIKeyAuth)
|
||||
user.user_role = None
|
||||
user.metadata = {"allowed_vector_store_indexes": [{"index_name": self.INDEX, "index_permissions": permissions}]}
|
||||
user.team_metadata = None
|
||||
return user
|
||||
|
||||
def test_team_with_write_grant_can_upload(self):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request("POST", f"/azure_ai/indexes/{self.INDEX}/docs/index"),
|
||||
user_api_key_dict=self._team_member(["read", "write"]),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_team_without_write_grant_cannot_upload(self):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request("POST", f"/azure_ai/indexes/{self.INDEX}/docs/index"),
|
||||
user_api_key_dict=self._team_member(["read"]),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
def test_team_with_read_grant_can_search(self):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request("POST", f"/azure_ai/indexes/{self.INDEX}/docs/search"),
|
||||
user_api_key_dict=self._team_member(["read"]),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_team_with_read_grant_can_get_index_details(self):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request("GET", f"/azure_ai/indexes/{self.INDEX}"),
|
||||
user_api_key_dict=self._team_member(["read"]),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_team_without_read_grant_cannot_get_index_details(self):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request("GET", f"/azure_ai/indexes/{self.INDEX}"),
|
||||
user_api_key_dict=self._team_member(["write"]),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.parametrize("method, path", READ_ROUTES)
|
||||
def test_team_with_read_grant_can_call_every_read_route(self, method, path):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request(method, path),
|
||||
user_api_key_dict=self._team_member(["read"]),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
@pytest.mark.parametrize("method, path", READ_ROUTES)
|
||||
def test_team_without_read_grant_cannot_call_read_routes(self, method, path):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request(method, path),
|
||||
user_api_key_dict=self._team_member(["write"]),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"method, operation, path",
|
||||
[
|
||||
("PUT", "update", f"/azure_ai/indexes/{INDEX}?api-version=2024-07-01"),
|
||||
("DELETE", "delete", f"/azure_ai/indexes/{INDEX}?api-version=2024-07-01"),
|
||||
("POST", "create", "/azure_ai/indexes?api-version=2024-07-01"),
|
||||
],
|
||||
)
|
||||
def test_team_cannot_manage_index_lifecycle_even_with_write_grant(self, method, operation, path):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request(method, path),
|
||||
user_api_key_dict=self._team_member(["read", "write"]),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
assert f"Only proxy admins can {operation}" in exc_info.value.detail
|
||||
|
||||
|
||||
class TestAzureAIAnalyzeNamedIndexClassification:
|
||||
"""Regression tests for write-before-read endpoint classification.
|
||||
|
||||
The endpoint matcher is substring-based, so the batch-write path of an
|
||||
index named ``analyze*`` contains the ``("POST", "/analyze")`` read
|
||||
fragment. Reads-first classification labeled that write a read, letting a
|
||||
read-only grant upload, merge, and delete documents (and refusing
|
||||
legitimate write-only grants). Writes are classified first now, so an
|
||||
ambiguous path demands the stronger grant.
|
||||
"""
|
||||
|
||||
def _request(self, method: str, path: str) -> MagicMock:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = method
|
||||
request.url.path = path
|
||||
return request
|
||||
|
||||
def _team_member(self, index: str, permissions: list) -> MagicMock:
|
||||
user = MagicMock(spec=UserAPIKeyAuth)
|
||||
user.user_role = None
|
||||
user.metadata = {"allowed_vector_store_indexes": [{"index_name": index, "index_permissions": permissions}]}
|
||||
user.team_metadata = None
|
||||
return user
|
||||
|
||||
@pytest.mark.parametrize("index", ["analyze", "analyzer-reports"])
|
||||
def test_read_only_grant_cannot_upload_to_analyze_named_index(self, index):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=index,
|
||||
request=self._request("POST", f"/azure_ai/indexes/{index}/docs/index"),
|
||||
user_api_key_dict=self._team_member(index, ["read"]),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.parametrize("index", ["analyze", "analyzer-reports"])
|
||||
def test_write_grant_can_upload_to_analyze_named_index(self, index):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=index,
|
||||
request=self._request("POST", f"/azure_ai/indexes/{index}/docs/index"),
|
||||
user_api_key_dict=self._team_member(index, ["write"]),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_read_only_grant_can_still_analyze_on_analyze_named_index(self):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name="analyze",
|
||||
request=self._request("POST", "/azure_ai/indexes/analyze/analyze"),
|
||||
user_api_key_dict=self._team_member("analyze", ["read"]),
|
||||
)
|
||||
assert result is True
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue