mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-22 00:31:44 +00:00
Merge remote-tracking branch 'origin/main' into litellm-ocr-new-mapping
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> # Conflicts: # litellm-rust/crates/core/src/ocr/adapters/azure/cohere.rs # litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/mod.rs # litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/polling.rs # litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs # litellm-rust/crates/core/src/ocr/adapters/cohere.rs # litellm-rust/crates/core/src/ocr/adapters/mistral.rs # litellm-rust/crates/core/src/ocr/adapters/reducto/mod.rs # litellm-rust/crates/core/src/ocr/adapters/vertex/deepseek.rs # litellm-rust/crates/core/src/ocr/adapters/vertex/mistral.rs # litellm-rust/crates/core/src/ocr/adapters/vertex/mod.rs # litellm-rust/crates/core/src/ocr/registry.rs # litellm-rust/crates/core/src/ocr/wire.rs # litellm-rust/crates/core/tests/host_lifecycle.rs # litellm-rust/crates/core/tests/ocr.rs # litellm-rust/crates/core/tests/vertex_ai_deepseek_ocr.rs # litellm-rust/crates/core/tests/vertex_ai_ocr.rs # litellm-rust/crates/python-bridge/src/routes/definition.rs # litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs # litellm-rust/crates/python-bridge/src/routes/ocr/value.rs
This commit is contained in:
commit
e09da4d485
7 changed files with 339 additions and 8 deletions
|
|
@ -108,6 +108,37 @@ def _trace_id_from_traceparent(traceparent: str) -> str | None:
|
|||
return trace_id if trace_id != "0" * 32 else None
|
||||
|
||||
|
||||
def _trace_id_from_otel_span(span: "OtelSpan | None") -> str | None:
|
||||
if span is None:
|
||||
return None
|
||||
try:
|
||||
span_context: Final = span.get_span_context()
|
||||
is_valid: Final = span_context.is_valid
|
||||
trace_id: Final = span_context.trace_id
|
||||
except AttributeError:
|
||||
return None
|
||||
if not is_valid or not isinstance(trace_id, int):
|
||||
return None
|
||||
return format(trace_id, "032x")
|
||||
|
||||
|
||||
def add_otel_trace_id_to_request(
|
||||
data: dict[str, object], _metadata_variable_name: str, parent_otel_span: "OtelSpan | None"
|
||||
) -> None:
|
||||
if data.get("litellm_trace_id"):
|
||||
return
|
||||
metadata: Final = data.get(_metadata_variable_name)
|
||||
requester_metadata: Final = data.get("metadata")
|
||||
if any(isinstance(m, dict) and m.get("trace_id") for m in (metadata, requester_metadata)):
|
||||
return
|
||||
trace_id: Final = _trace_id_from_otel_span(parent_otel_span)
|
||||
if trace_id is None:
|
||||
return
|
||||
data["litellm_trace_id"] = trace_id # rebind-ok: data is an out-param
|
||||
if isinstance(metadata, dict):
|
||||
metadata["trace_id"] = trace_id # rebind-ok: metadata is the request's own out-param dict
|
||||
|
||||
|
||||
def _session_id_from_baggage(baggage: str) -> str | None:
|
||||
"""Extract a session.id entry from a W3C Baggage header
|
||||
(https://www.w3.org/TR/baggage/), e.g. "session.id=abc-123,user.id=42"."""
|
||||
|
|
@ -173,6 +204,8 @@ _ENABLE_TEAM_STALE_ALIAS_BYPASS: bool | None = None
|
|||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as OtelSpan
|
||||
|
||||
from litellm.integrations.otel.model.destination import OtelDestination
|
||||
from litellm.proxy.policy_engine.attachment_registry import AttachmentRegistry
|
||||
from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig
|
||||
|
|
@ -2042,6 +2075,13 @@ async def add_litellm_data_to_request(
|
|||
data=data,
|
||||
_metadata_variable_name=_metadata_variable_name,
|
||||
)
|
||||
add_otel_trace_id_to_request(
|
||||
data=data,
|
||||
_metadata_variable_name=_metadata_variable_name,
|
||||
parent_otel_span=user_api_key_dict.parent_otel_span
|
||||
if user_api_key_dict.parent_otel_span is not None
|
||||
else getattr(request.state, "parent_otel_span", None),
|
||||
)
|
||||
apply_missing_session_id_policy(
|
||||
data=data,
|
||||
_metadata_variable_name=_metadata_variable_name,
|
||||
|
|
|
|||
|
|
@ -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)},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import pytest
|
||||
from botocore.credentials import Credentials
|
||||
from fastapi import Request
|
||||
from opentelemetry.trace import INVALID_SPAN, NonRecordingSpan, SpanContext
|
||||
from pydantic import ValidationError as PydanticValidationError
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
|
|
@ -559,6 +560,7 @@ def _batches_request_mock() -> MagicMock:
|
|||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
request_mock.state.parent_otel_span = None
|
||||
return request_mock
|
||||
|
||||
|
||||
|
|
@ -2813,7 +2815,7 @@ def test_add_headers_to_llm_call_by_model_group_existing_headers_in_data():
|
|||
litellm.model_group_settings = original_model_group_settings
|
||||
|
||||
|
||||
from typing import Optional
|
||||
from typing import Final, Optional
|
||||
|
||||
from fastapi.responses import Response
|
||||
|
||||
|
|
@ -3536,6 +3538,163 @@ def test_add_litellm_metadata_from_request_headers_explicit_trace_id_beats_trace
|
|||
assert data["litellm_session_id"] == "explicit-trace-id-value"
|
||||
|
||||
|
||||
def _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan:
|
||||
return NonRecordingSpan(SpanContext(trace_id=trace_id, span_id=0x00F067AA0BA902B7, is_remote=False))
|
||||
|
||||
|
||||
def _request_mock_without_trace_headers() -> MagicMock:
|
||||
request_mock: Final = MagicMock(spec=Request)
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
return request_mock
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_defaults_trace_id_to_otel_server_span():
|
||||
"""With OTel on and a client that sends no trace headers, the request's
|
||||
litellm_trace_id (and so the spend log session_id) must be the W3C trace-id
|
||||
of the proxy's server span, so a trace in the OTel backend can be looked up
|
||||
in the Logs UI and vice versa."""
|
||||
otel_trace_id: Final = 0x4BF92F3577B34DA6A3CE929D0E0E4736
|
||||
user_api_key_dict: Final = UserAPIKeyAuth(
|
||||
api_key="hashed-key", parent_otel_span=_otel_span_with_trace_id(otel_trace_id)
|
||||
)
|
||||
|
||||
data: Final = await add_litellm_data_to_request(
|
||||
data={"model": "gpt-5.6", "messages": [{"role": "user", "content": "hi"}]},
|
||||
request=_request_mock_without_trace_headers(),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
)
|
||||
|
||||
assert data["litellm_trace_id"] == format(otel_trace_id, "032x")
|
||||
assert data["metadata"]["trace_id"] == format(otel_trace_id, "032x")
|
||||
assert "litellm_session_id" not in data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_falls_back_to_request_state_otel_span():
|
||||
"""Custom auth hooks return a UserAPIKeyAuth without parent_otel_span even
|
||||
though user_api_key_auth already opened the server span on request.state,
|
||||
so the fallback must read the span from there or custom-auth requests would
|
||||
keep getting an unrelated session id."""
|
||||
otel_trace_id: Final = 0x4BF92F3577B34DA6A3CE929D0E0E4736
|
||||
request_mock: Final = _request_mock_without_trace_headers()
|
||||
request_mock.state.parent_otel_span = _otel_span_with_trace_id(otel_trace_id)
|
||||
|
||||
data: Final = await add_litellm_data_to_request(
|
||||
data={"model": "gpt-5.6"},
|
||||
request=request_mock,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", parent_otel_span=None),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
)
|
||||
|
||||
assert data["litellm_trace_id"] == format(otel_trace_id, "032x")
|
||||
assert data["metadata"]["trace_id"] == format(otel_trace_id, "032x")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_otel_span_does_not_override_caller_trace_id():
|
||||
"""A caller's own trace identity (x-litellm-trace-id header or body
|
||||
metadata.trace_id) keeps priority over the OTel server span's trace-id."""
|
||||
span: Final = _otel_span_with_trace_id(0x4BF92F3577B34DA6A3CE929D0E0E4736)
|
||||
|
||||
header_request: Final = _request_mock_without_trace_headers()
|
||||
header_request.headers = {"Content-Type": "application/json", "x-litellm-trace-id": "caller-trace"}
|
||||
from_header: Final = await add_litellm_data_to_request(
|
||||
data={"model": "gpt-5.6"},
|
||||
request=header_request,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", parent_otel_span=span),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
)
|
||||
assert from_header["litellm_trace_id"] == "caller-trace"
|
||||
assert from_header["metadata"]["trace_id"] == "caller-trace"
|
||||
|
||||
from_body: Final = await add_litellm_data_to_request(
|
||||
data={"model": "gpt-5.6", "metadata": {"trace_id": "body-trace"}},
|
||||
request=_request_mock_without_trace_headers(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", parent_otel_span=span),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
)
|
||||
assert "litellm_trace_id" not in from_body
|
||||
assert from_body["metadata"]["trace_id"] == "body-trace"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"])
|
||||
async def test_add_litellm_data_to_request_otel_span_does_not_override_body_trace_id_on_litellm_metadata_routes(path):
|
||||
"""On routes that keep LiteLLM state in litellm_metadata, the caller's body
|
||||
metadata.trace_id is only promoted into litellm_metadata later in the
|
||||
pipeline, so the OTel fallback must look at the requester metadata too or
|
||||
it would claim the slot first and the caller's id would be lost."""
|
||||
request_mock: Final = _request_mock_without_trace_headers()
|
||||
request_mock.url.path = path
|
||||
request_mock.url.__str__.return_value = f"http://localhost{path}"
|
||||
data: Final = await add_litellm_data_to_request(
|
||||
data={"model": "gpt-5.6", "metadata": {"trace_id": "body-trace"}},
|
||||
request=request_mock,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="hashed-key", parent_otel_span=_otel_span_with_trace_id(0x4BF92F3577B34DA6A3CE929D0E0E4736)
|
||||
),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
)
|
||||
assert "litellm_trace_id" not in data
|
||||
assert data["litellm_metadata"]["trace_id"] == "body-trace"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("empty_trace_id", [None, ""])
|
||||
async def test_add_litellm_data_to_request_otel_span_fills_empty_body_trace_id(empty_trace_id):
|
||||
"""A serialized-but-empty litellm_trace_id in the body (null or "") carries
|
||||
no identity, so it must not block the OTel server span fallback."""
|
||||
otel_trace_id: Final = 0x4BF92F3577B34DA6A3CE929D0E0E4736
|
||||
data: Final = await add_litellm_data_to_request(
|
||||
data={"model": "gpt-5.6", "litellm_trace_id": empty_trace_id},
|
||||
request=_request_mock_without_trace_headers(),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="hashed-key", parent_otel_span=_otel_span_with_trace_id(otel_trace_id)
|
||||
),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
)
|
||||
assert data["litellm_trace_id"] == format(otel_trace_id, "032x")
|
||||
assert data["metadata"]["trace_id"] == format(otel_trace_id, "032x")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("parent_otel_span", [None, "invalid_span", "not_a_span", "plain_string"])
|
||||
async def test_add_litellm_data_to_request_no_trace_id_without_valid_otel_span(parent_otel_span):
|
||||
"""No OTel span (OTel off), a span with an invalid context, an object that
|
||||
only quacks like a span, or a value that is not a span at all (custom auth
|
||||
is typed loosely and can hand back anything) must leave litellm_trace_id
|
||||
unset, and never fail the request, so downstream keeps generating its own id."""
|
||||
span: Final = {
|
||||
"invalid_span": INVALID_SPAN,
|
||||
"not_a_span": MagicMock(),
|
||||
"plain_string": "not-a-span",
|
||||
}.get(parent_otel_span)
|
||||
data: Final = await add_litellm_data_to_request(
|
||||
data={"model": "gpt-5.6"},
|
||||
request=_request_mock_without_trace_headers(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", parent_otel_span=span),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
)
|
||||
assert "litellm_trace_id" not in data
|
||||
assert "trace_id" not in data["metadata"]
|
||||
|
||||
|
||||
def test_add_litellm_metadata_from_request_headers_anthropic_metadata_beats_baggage():
|
||||
"""The existing Anthropic metadata.user_id session_id path must win over a
|
||||
baggage session.id fallback."""
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue