Merge pull request #34427 from BerriAI/litellm_bedrock_rag_retrieval_filter

fix(rag): forward retrieval_filter from retrieval_config to vector store search
This commit is contained in:
Mateo Wang 2026-09-16 11:34:44 -07:00 • committed by GitHub
commit 78848d01a3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 139 additions and 7 deletions

View file

@ -69,6 +69,11 @@ def _response_attr(source: object, name: str) -> object:
return getattr(source, name, None)
def _upstream_status_code(error: Exception) -> int:
code: Final = getattr(error, "status_code", None)
return code if isinstance(code, int) else 500
def _raise_vector_store_scan_depth_exceeded() -> None:
raise HTTPException(
status_code=400,
@ -814,6 +819,6 @@ async def rag_query(
except Exception as e:
verbose_proxy_logger.exception("RAG Query failed: %s", e)
raise HTTPException(
status_code=500,
status_code=_upstream_status_code(e),
detail={"error": str(e)},
)

View file

@ -245,6 +245,9 @@ async def _execute_query_pipeline(
raise ValueError("No query found in messages for RAG query")
# 2. Search vector store
top_level_filters: Final = kwargs.pop("filters", None)
filters: Final = retrieval_config.get("retrieval_filter") or retrieval_config.get("filters") or top_level_filters
filter_search_params: Final = MappingProxyType({"filters": filters} if filters else {})
# Forward allowlisted provider retrieval_config extras (region, embedding
# model, bucket, credential refs) to the search call; the managed store's
# params win on conflict.
@ -258,7 +261,9 @@ async def _execute_query_pipeline(
if k not in _SEARCH_ARGS_SET_BY_PIPELINE
}
)
forwarded_search_params: Final = MappingProxyType({**provider_search_params, **kwargs, **store_search_params})
forwarded_search_params: Final = MappingProxyType(
{**provider_search_params, **kwargs, **filter_search_params, **store_search_params}
)
with _suppressed_sub_call_billing():
search_response: Final = await litellm.vector_stores.asearch(
vector_store_id=retrieval_config["vector_store_id"],

View file

@ -2,10 +2,11 @@
Type definitions for RAG (Retrieval Augmented Generation) Ingest API.
"""
from collections.abc import Mapping
from typing import Any, Literal
from pydantic import BaseModel, ConfigDict
from typing_extensions import TypedDict
from typing_extensions import ReadOnly, TypedDict
from litellm.types.utils import ModelResponse
@ -237,10 +238,11 @@ class RAGIngestRequest(BaseModel):
class RAGRetrievalConfig(TypedDict, total=False):
"""Configuration for vector store retrieval."""
vector_store_id: str
custom_llm_provider: str
top_k: int # max results from vector store
filters: dict[str, Any] | None # optional - vector store filters
vector_store_id: ReadOnly[str]
custom_llm_provider: ReadOnly[str]
top_k: ReadOnly[int]
filters: ReadOnly[Mapping[str, object] | None]
retrieval_filter: ReadOnly[Mapping[str, object] | None]
class RAGRerankConfig(TypedDict, total=False):

View file

@ -11,6 +11,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi.testclient import TestClient
import litellm
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.proxy_server import app
@ -282,6 +283,41 @@ def test_rag_query_returns_response_cost_header(client_internal_user):
assert response.headers.get("x-litellm-response-cost") == "3.45e-06"
@pytest.mark.parametrize(
("upstream_error", "expected_status"),
[
(litellm.BadRequestError(message="filter andAll needs two clauses", model="kb", llm_provider="bedrock"), 400),
(litellm.NotFoundError(message="Knowledge Base does not exist", model="kb", llm_provider="bedrock"), 404),
(RuntimeError("pipeline blew up"), 500),
],
)
def test_rag_query_surfaces_upstream_status_code(client_internal_user, upstream_error, expected_status):
"""A vector store rejection must reach the caller with its own status code, never a blanket 500."""
with (
patch( # test-quality-ok: the handler calls the module-level litellm.aquery directly; no injection seam
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
new=AsyncMock(side_effect=upstream_error),
),
patch("litellm.vector_store_registry", None), # test-quality-ok: proxy module global, no injection seam
patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: proxy module global, no injection seam
):
response = client_internal_user.post(
"/v1/rag/query",
json={
"model": "bedrock/us.anthropic.claude-sonnet-5",
"messages": [{"role": "user", "content": "How was this document ingested?"}],
"retrieval_config": {
"vector_store_id": "L7INRFMVQT",
"custom_llm_provider": "bedrock",
"retrieval_filter": {"andAll": [{"equals": {"key": "department", "value": "billing"}}]},
},
},
)
assert response.status_code == expected_status, response.text
assert str(upstream_error) in response.json()["detail"]["error"]
def test_rag_query_stream_returns_event_stream(client_internal_user):
"""
A stream=true /v1/rag/query must return an SSE response. Returning the raw

View file

@ -11,9 +11,13 @@ aquery carries the completion response with real usage and cost.
"""
import asyncio
import json
from typing import Final
from unittest.mock import patch
import httpx
import pytest
import respx
import litellm
from litellm._internal_context import is_internal_call
@ -259,6 +263,86 @@ async def test_aquery_streaming_bills_sub_call_costs_into_final_event():
assert standard_logging_object["response_cost"] >= 0.003
@pytest.mark.asyncio
@pytest.mark.parametrize(
("retrieval_config_json", "top_level_filter_json", "expected_filter_json"),
(
(
'{"vector_store_id":"vs_test_123","custom_llm_provider":"openai","top_k":50,'
'"retrieval_filter":{"equals":{"key":"tenant","value":"retrieval"}}}',
None,
'{"equals":{"key":"tenant","value":"retrieval"}}',
),
(
'{"vector_store_id":"vs_test_123","custom_llm_provider":"openai","top_k":50,'
'"filters":{"equals":{"key":"tenant","value":"alias"}}}',
None,
'{"equals":{"key":"tenant","value":"alias"}}',
),
(
'{"vector_store_id":"vs_test_123","custom_llm_provider":"openai","top_k":50}',
'{"equals":{"key":"tenant","value":"top-level"}}',
'{"equals":{"key":"tenant","value":"top-level"}}',
),
(
'{"vector_store_id":"vs_test_123","custom_llm_provider":"openai","top_k":50,'
'"retrieval_filter":{"equals":{"key":"tenant","value":"retrieval"}},'
'"filters":{"equals":{"key":"tenant","value":"alias"}}}',
'{"equals":{"key":"tenant","value":"top-level"}}',
'{"equals":{"key":"tenant","value":"retrieval"}}',
),
(
'{"vector_store_id":"vs_test_123","custom_llm_provider":"openai","top_k":50}',
None,
None,
),
),
)
async def test_aquery_forwards_filters_to_vector_store_search(
retrieval_config_json: str,
top_level_filter_json: str | None,
expected_filter_json: str | None,
monkeypatch,
):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
retrieval_config: Final = json.loads(retrieval_config_json)
top_level_filter: Final = json.loads(top_level_filter_json) if top_level_filter_json is not None else None
expected_filter: Final = json.loads(expected_filter_json) if expected_filter_json is not None else None
with respx.mock(assert_all_called=True) as respx_mock:
search_route: Final = respx_mock.post("https://example.com/v1/vector_stores/vs_test_123/search").mock(
return_value=httpx.Response(
200,
content='{"object":"vector_store.search_results.page","search_query":"q","data":[]}',
)
)
respx_mock.post("https://example.com/v1/chat/completions").mock(
return_value=httpx.Response(
200,
content=(
'{"id":"chatcmpl-test","object":"chat.completion","created":1,"model":"gpt-4o-mini",'
'"choices":[{"index":0,"message":{"role":"assistant","content":"answer"},"finish_reason":"stop"}],'
'"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}'
),
)
)
response: Final = await litellm.aquery(
model="openai/gpt-4o-mini",
messages=json.loads('[{"role":"user","content":"most frequent causes of low nicotine"}]'),
retrieval_config=retrieval_config,
filters=top_level_filter,
api_key="sk-test",
api_base="https://example.com/v1",
)
request_body: Final = json.loads(search_route.calls.last.request.content)
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "answer"
assert request_body["query"] == "most frequent causes of low nicotine"
assert request_body.get("filters") == expected_filter
assert request_body["max_num_results"] == 50
@pytest.mark.asyncio
async def test_aquery_forwards_provider_retrieval_config_and_router_to_search():
"""