litellm/tests/unit/rag/test_main.py
Mateo Wang d80f8c28ca
refactor(types): replace Any with proven types in 137 files (#44478)
* refactor(types): prove runtime types at harness, search, rag and client boundaries

Replace Any with adapter-validated types in the litellm.agent() harness, the
search provider transformations, RAG ingestion and query, the vector store
pre-call hook and registry, the galileo and opik logging integrations and the
proxy client CLI. Each boundary gets unit tests for well-formed and malformed
payloads.

* chore(typing): prove types at more provider boundaries and restore search transformations

Second pass over a2a, embedding, rerank, image, audio and small provider
modules. Search transformations go back to their previous form because
validating their response bodies would change the proxy status for
malformed upstream bodies from 400 to 500.

* refactor(types): prove types at logging, files, rerank, image, audio and management boundaries

Replace Any with validated or annotation-only types in 37 more files: logging
integrations, token counters, provider files/rerank/image generation/audio
transcription transformations, pass-through logging handlers and management
endpoints. No proxy HTTP status or error type changes.

* refactor(types): prove types at repository, spend, files and router boundaries

Replace Any with repository table accessors, validated mappings and
annotation-only types in 33 more files: Prisma repositories, the enterprise
batch and responses cost checkers, budget reservation, files endpoints,
management endpoints, the policy registry, the adaptive and complexity routers
and the secret managers. No proxy HTTP status or error type changes.

* test(types): run the aiohttp transformation test in-process and cover repository row conversion

The aiohttp chat transformation test no longer starts a server. It feeds the
transformation a response whose json() returns the body under test.

The proxy unit shards now exercise stored model rows whose params are JSON
strings and the object permission create and update paths.
2026-10-05 11:01:45 -07:00

637 lines
26 KiB
Python

"""
Tests for the RAG query pipeline in litellm/rag/main.py.
The RAG pipeline forwards its kwargs (including the parent litellm_logging_obj)
into @client-decorated sub-calls (vector store search, completion). Each logging
object allows exactly one async_success event, so if sub-calls are not marked as
internal, the vector store search consumes the slot first and the LLM
completion's usage/cost is never logged (spend tracking and budget enforcement
are bypassed). These tests pin the invariant that the single billing event for
aquery carries the completion response with real usage and cost.
"""
import asyncio
import json
from types import MappingProxyType
from typing import Final
from unittest.mock import patch
import httpx
import pytest
import respx
from pydantic import ValidationError
import litellm
from litellm._internal_context import is_internal_call
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.types.utils import CallTypes, ModelResponse
async def _drain_logging_worker() -> None:
"""Run every queued logging task to completion on the current event loop.
The success event is delivered through the fire-and-forget GLOBAL_LOGGING_WORKER
singleton, whose queue survives across tests. start() rebinds any tasks left over
from a previous test's event loop onto the current one, and flush() waits until
the queue is fully processed, so tests neither miss their own event nor observe
a neighbour's
"""
await asyncio.sleep(0)
GLOBAL_LOGGING_WORKER.start()
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0)
class RecordingLogger(CustomLogger):
def __init__(self):
super().__init__()
self.success_events = []
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
self.success_events.append({"kwargs": kwargs, "response_obj": response_obj})
@pytest.mark.asyncio
@pytest.mark.parametrize("use_router", [False, True])
async def test_aquery_single_billing_event_carries_completion_usage_and_cost(use_router):
"""
litellm.aquery must produce exactly one success event, and that event must
carry the LLM completion (a ModelResponse with non-zero usage and cost),
not the vector store search response. The proxy always passes a router, so
both the router and non-router completion branches are pinned.
"""
await _drain_logging_worker()
recording_logger = RecordingLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [recording_logger]
router_kwargs = {}
if use_router:
router_kwargs["router"] = litellm.Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"},
}
]
)
try:
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "What is the secret project codename?"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
mock_response="The secret project codename is AZURE-FALCON-42.",
**router_kwargs,
)
assert isinstance(response, ModelResponse)
assert is_internal_call.get() is False
await _drain_logging_worker()
finally:
litellm.callbacks = original_callbacks
assert len(recording_logger.success_events) == 1
event = recording_logger.success_events[0]
response_obj = event["response_obj"]
assert isinstance(response_obj, ModelResponse)
assert response_obj.usage.total_tokens > 0
standard_logging_object = event["kwargs"]["standard_logging_object"]
assert standard_logging_object["call_type"] == "aquery"
assert standard_logging_object["total_tokens"] > 0
assert standard_logging_object["prompt_tokens"] > 0
assert standard_logging_object["completion_tokens"] > 0
assert standard_logging_object["response_cost"] > 0
@pytest.mark.asyncio
async def test_aquery_response_hidden_params_carry_completion_cost():
"""
The aquery response must expose the completion's response_cost via hidden
params, so the proxy can return the x-litellm-response-cost header.
"""
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
mock_response="hi there",
)
await _drain_logging_worker()
assert isinstance(response, ModelResponse)
response_cost = response._hidden_params.get("response_cost")
assert response_cost is not None
assert response_cost > 0
@pytest.mark.asyncio
async def test_aquery_billed_cost_includes_priced_vector_store_search():
"""
When the vector store provider prices search calls (e.g. per-query cost),
that cost must be folded into the aquery billing instead of being dropped
with the suppressed sub-call event.
"""
await _drain_logging_worker()
recording_logger = RecordingLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [recording_logger]
try:
with patch("litellm.rag.main.vector_store_search_cost", return_value=(0.002, 0.0)):
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
mock_response="hi there",
)
await _drain_logging_worker()
finally:
litellm.callbacks = original_callbacks
assert isinstance(response, ModelResponse)
total_cost = response._hidden_params.get("response_cost")
assert total_cost is not None
assert total_cost > 0.002
assert len(recording_logger.success_events) == 1
standard_logging_object = recording_logger.success_events[0]["kwargs"]["standard_logging_object"]
assert standard_logging_object["response_cost"] == total_cost
@pytest.mark.asyncio
async def test_aquery_with_rerank_bills_once_and_folds_rerank_cost():
"""
When rerank is enabled, its sub-call must run under the internal-call
context (no standalone billing event) and its cost must be folded into
the single aquery billing event.
"""
from litellm.types.rerank import RerankResponse
await _drain_logging_worker()
recording_logger = RecordingLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [recording_logger]
rerank_seen = {}
async def fake_arerank(**kwargs):
rerank_seen["internal"] = is_internal_call.get()
rerank_result = RerankResponse(id="rr_1", results=[{"index": 0, "relevance_score": 0.9}], meta={})
rerank_result._hidden_params["response_cost"] = 0.001
return rerank_result
try:
with patch("litellm.arerank", side_effect=fake_arerank):
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
rerank={"enabled": True, "model": "cohere/rerank-english-v3.0", "top_n": 1},
mock_response="hi there",
)
await _drain_logging_worker()
finally:
litellm.callbacks = original_callbacks
assert rerank_seen["internal"] is True
assert is_internal_call.get() is False
assert isinstance(response, ModelResponse)
total_cost = response._hidden_params.get("response_cost")
assert total_cost is not None
assert total_cost > 0.001
assert len(recording_logger.success_events) == 1
standard_logging_object = recording_logger.success_events[0]["kwargs"]["standard_logging_object"]
assert standard_logging_object["call_type"] == "aquery"
assert standard_logging_object["response_cost"] == total_cost
@pytest.mark.asyncio
async def test_aquery_streaming_bills_sub_call_costs_into_final_event():
"""
On the streaming path the response cost is computed from the assembled
chunks after the pipeline returns, so there is no response object to fold
sub-call costs into. The pipeline must instead carry the accumulated
search and rerank cost through the logging object so the single streamed
billing event includes it; otherwise a caller passing stream=true incurs
priced vector search and rerank costs that never reach spend tracking.
"""
from litellm.types.rerank import RerankResponse
await _drain_logging_worker()
recording_logger = RecordingLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [recording_logger]
rerank_seen = {}
async def fake_arerank(**kwargs):
rerank_seen["internal"] = is_internal_call.get()
rerank_result = RerankResponse(id="rr_1", results=[{"index": 0, "relevance_score": 0.9}], meta={})
rerank_result._hidden_params["response_cost"] = 0.001
return rerank_result
try:
with (
patch("litellm.rag.main.vector_store_search_cost", return_value=(0.002, 0.0)),
patch("litellm.arerank", side_effect=fake_arerank),
):
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
rerank={"enabled": True, "model": "cohere/rerank-english-v3.0", "top_n": 1},
mock_response="hi there",
stream=True,
)
async for _ in response:
pass
await _drain_logging_worker()
finally:
litellm.callbacks = original_callbacks
assert rerank_seen["internal"] is True
assert is_internal_call.get() is False
assert len(recording_logger.success_events) == 1
standard_logging_object = recording_logger.success_events[0]["kwargs"]["standard_logging_object"]
assert standard_logging_object["call_type"] == "aquery"
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():
"""
Regression: provider-specific retrieval_config keys (aws_region_name,
embedding_model, vector_bucket_name, ...) and the router must be forwarded
to the vector store search call. Pre-fix they were silently dropped, so
/v1/rag/query failed with provider config errors (e.g. S3 Vectors
"aws_region_name is required") even when the caller supplied them.
"""
from unittest.mock import AsyncMock
from litellm.types.vector_stores import VectorStoreSearchResponse
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"},
}
]
)
fake_search = AsyncMock(
return_value=VectorStoreSearchResponse(
object="vector_store.search_results.page", search_query="q", data=[]
)
)
with patch("litellm.vector_stores.asearch", new=fake_search): # test-quality-ok: asearch is the boundary the forwarding contract under test targets
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={
"vector_store_id": "bkt:idx",
"custom_llm_provider": "s3_vectors",
"top_k": 5,
"aws_region_name": "eu-west-1",
"embedding_model": "my-embed",
"vector_bucket_name": "bkt",
},
router=router,
mock_response="hi",
)
assert isinstance(response, ModelResponse)
fake_search.assert_awaited_once()
search_kwargs = fake_search.await_args.kwargs
assert search_kwargs["vector_store_id"] == "bkt:idx"
assert search_kwargs["custom_llm_provider"] == "s3_vectors"
assert search_kwargs["max_num_results"] == 5
assert search_kwargs["router"] is router
# provider-specific extras forwarded
assert search_kwargs["aws_region_name"] == "eu-west-1"
assert search_kwargs["embedding_model"] == "my-embed"
assert search_kwargs["vector_bucket_name"] == "bkt"
# consumed keys are not duplicated into the spread
assert "top_k" not in search_kwargs
@pytest.mark.asyncio
async def test_aquery_minimal_retrieval_config_forwards_no_extras():
"""
A minimal retrieval_config must not leak consumed keys (or invent extras)
into the vector store search call.
"""
from unittest.mock import AsyncMock
from litellm.types.vector_stores import VectorStoreSearchResponse
fake_search = AsyncMock(
return_value=VectorStoreSearchResponse(
object="vector_store.search_results.page", search_query="q", data=[]
)
)
with patch("litellm.vector_stores.asearch", new=fake_search): # test-quality-ok: asearch is the boundary the forwarding contract under test targets
await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
mock_response="hi",
)
fake_search.assert_awaited_once()
search_kwargs = fake_search.await_args.kwargs
assert search_kwargs["vector_store_id"] == "vs_test_123"
assert search_kwargs["custom_llm_provider"] == "openai"
assert search_kwargs["router"] is None
leaked = {"top_k", "filters", "retrieval_filter", "aws_region_name", "embedding_model", "vector_bucket_name"}
assert not (leaked & set(search_kwargs.keys()))
@pytest.mark.asyncio
async def test_aquery_does_not_forward_connection_override_keys_to_search():
"""
Only allowlisted retrieval_config keys may reach the vector store search
call. Caller-controlled connection overrides (api_base, api_key, arbitrary
extras) must be dropped, otherwise a caller could redirect store
credentials to an attacker-chosen host.
"""
from unittest.mock import AsyncMock
from litellm.types.vector_stores import VectorStoreSearchResponse
fake_search = AsyncMock(
return_value=VectorStoreSearchResponse(
object="vector_store.search_results.page", search_query="q", data=[]
)
)
with patch("litellm.vector_stores.asearch", new=fake_search): # test-quality-ok: asearch is the boundary the forwarding contract under test targets
await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={
"vector_store_id": "bkt:idx",
"custom_llm_provider": "s3_vectors",
"aws_region_name": "eu-west-1",
"api_base": "https://attacker.example.com",
"api_key": "attacker-key",
"arbitrary_extra": "nope",
},
mock_response="hi",
)
fake_search.assert_awaited_once()
search_kwargs = fake_search.await_args.kwargs
assert search_kwargs["aws_region_name"] == "eu-west-1"
blocked = {"api_base", "api_key", "arbitrary_extra"}
assert not (blocked & set(search_kwargs.keys()))
@pytest.mark.asyncio
async def test_aquery_forwards_vector_store_params_to_search_but_not_completion():
"""
Regression for LIT-6773: the server-trusted vector_store_params (a managed
store's litellm_params) must reach the search call wholesale, including the
connection keys the caller allowlist blocks, while the caller's own
retrieval_config overrides stay blocked, the caller's top-level api_key and
api_base stay on the completion only, and the completion never inherits the
store's connection params.
"""
from unittest.mock import AsyncMock
from litellm.types.vector_stores import VectorStoreSearchResponse
fake_search = AsyncMock(
return_value=VectorStoreSearchResponse(
object="vector_store.search_results.page", search_query="q", data=[]
)
)
fake_completion = AsyncMock(
return_value=ModelResponse(
id="chatcmpl-test",
choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
model="gpt-4o-mini",
)
)
with (
patch("litellm.vector_stores.asearch", new=fake_search), # test-quality-ok: the search boundary under test
patch("litellm.acompletion", new=fake_completion), # test-quality-ok: the completion boundary under test
):
await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
api_key="sk-llm-key",
api_base="https://llm.example.com",
retrieval_config={
"vector_store_id": "customer_kb",
"custom_llm_provider": "milvus",
"api_base": "https://attacker.example.com",
"api_key": "attacker-key",
},
vector_store_params={
"vector_store_id": "customer_kb",
"custom_llm_provider": "milvus",
"api_base": "http://127.0.0.1:19530",
"api_key": "root:Milvus",
"milvus_text_field": "book_intro_text",
"outputFields": ["book_intro_text"],
},
)
fake_search.assert_awaited_once()
search_kwargs = fake_search.await_args.kwargs
assert search_kwargs["vector_store_id"] == "customer_kb"
assert search_kwargs["custom_llm_provider"] == "milvus"
assert search_kwargs["api_base"] == "http://127.0.0.1:19530"
assert search_kwargs["api_key"] == "root:Milvus"
assert search_kwargs["milvus_text_field"] == "book_intro_text"
assert search_kwargs["outputFields"] == ["book_intro_text"]
fake_completion.assert_awaited_once()
completion_kwargs = fake_completion.await_args.kwargs
assert completion_kwargs["api_key"] == "sk-llm-key"
assert completion_kwargs["api_base"] == "https://llm.example.com"
assert not ({"milvus_text_field", "outputFields"} & set(completion_kwargs))
_UNDECODABLE_FILE: Final = {"filename": "notes.txt", "content": "x"}
@pytest.mark.parametrize(
("ingest_options", "expected_provider"),
[
({"vector_store": {"custom_llm_provider": "bedrock"}}, "bedrock"),
({"vector_store": MappingProxyType({"custom_llm_provider": "bedrock"})}, "bedrock"),
({"vector_store": {"vector_store_id": "vs_1", 7: "ignored"}}, None),
({}, None),
],
)
def test_ingest_failure_is_attributed_to_the_vector_store_provider(
ingest_options: dict[str, object], expected_provider: str | None
) -> None:
with pytest.raises(litellm.APIConnectionError, match="Invalid base64-encoded string") as raised:
litellm.ingest(ingest_options=ingest_options, file=_UNDECODABLE_FILE)
assert raised.value.llm_provider == expected_provider
@pytest.mark.asyncio
@pytest.mark.parametrize(
("ingest_options", "expected_provider"),
[
({"vector_store": {"custom_llm_provider": "bedrock"}}, "bedrock"),
({"vector_store": MappingProxyType({"custom_llm_provider": "bedrock"})}, "bedrock"),
({"vector_store": {"vector_store_id": "vs_1", 7: "ignored"}}, None),
({}, None),
],
)
async def test_aingest_failure_is_attributed_to_the_vector_store_provider(
ingest_options: dict[str, object], expected_provider: str | None
) -> None:
with pytest.raises(litellm.APIConnectionError, match="Invalid base64-encoded string") as raised:
await litellm.aingest(ingest_options=ingest_options, file=_UNDECODABLE_FILE)
assert raised.value.llm_provider == expected_provider
@pytest.mark.parametrize("vector_store", [None, "openai", ["openai"], [{"api_key": "sk-test"}]])
def test_ingest_failure_with_a_vector_store_that_is_not_a_mapping_raises_a_validation_error(
vector_store: object,
) -> None:
with pytest.raises(ValidationError) as raised:
litellm.ingest(ingest_options={"vector_store": vector_store}, file=_UNDECODABLE_FILE)
assert "sk-test" not in str(raised.value)
@pytest.mark.asyncio
@pytest.mark.parametrize("vector_store", [None, "openai", ["openai"], [{"api_key": "sk-test"}]])
async def test_aingest_failure_with_a_vector_store_that_is_not_a_mapping_raises_a_validation_error(
vector_store: object,
) -> None:
with pytest.raises(ValidationError) as raised:
await litellm.aingest(ingest_options={"vector_store": vector_store}, file=_UNDECODABLE_FILE)
assert "sk-test" not in str(raised.value)
@pytest.mark.asyncio
@pytest.mark.parametrize("provider", [None, 5])
async def test_aquery_with_a_provider_that_is_not_a_string_bills_only_the_completion(provider: object) -> None:
messages: Final = [{"role": "user", "content": "hello"}]
default_provider_response: Final = await litellm.aquery(
model="gpt-4o-mini",
messages=messages,
retrieval_config={"vector_store_id": "vs_test_123"},
mock_response="hi there",
)
response: Final = await litellm.aquery(
model="gpt-4o-mini",
messages=messages,
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": provider},
mock_response="hi there",
)
await _drain_logging_worker()
assert response._hidden_params["response_cost"] > 0
assert response._hidden_params["response_cost"] == default_provider_response._hidden_params["response_cost"]
def test_rag_call_types_are_registered():
"""
query/aquery/ingest/aingest are @client-decorated entry points, so their
function names must resolve to CallTypes members (deployment hooks and
call-type driven logic silently no-op for unregistered call types).
"""
assert CallTypes("query") is CallTypes.query
assert CallTypes("aquery") is CallTypes.aquery
assert CallTypes("ingest") is CallTypes.ingest
assert CallTypes("aingest") is CallTypes.aingest