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:
Mateo Wang 2026-08-14 17:24:06 -07:00 • committed by GitHub
commit d74cb6de1b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 484 additions and 17 deletions

View file

@ -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:

View file

@ -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:

View file

@ -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:

View file

@ -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()

View file

@ -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