mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-05 08:07:05 +00:00
Merge 7774c308ca into 67c7b97fd2
This commit is contained in:
commit
10523ee818
4 changed files with 201 additions and 0 deletions
|
|
@ -1568,6 +1568,7 @@ _MODEL_ROUTING_ROUTE_MARKERS: Final = (
|
|||
_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS: Final = (
|
||||
"/files",
|
||||
"/batches",
|
||||
"/vector_stores",
|
||||
"/skills",
|
||||
"/evals",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -94,6 +94,31 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
|
|||
return data
|
||||
|
||||
|
||||
def _update_request_data_with_model_from_request(
|
||||
data: Dict,
|
||||
request: Request,
|
||||
) -> Dict:
|
||||
"""
|
||||
Add model routing from request query/header when no body or registry model exists.
|
||||
|
||||
The vector store retrieve/delete endpoints do not have request bodies, but
|
||||
clients can still pass the proxy model through query params or
|
||||
x-litellm-model. Preserve registry-derived routing first.
|
||||
"""
|
||||
if (
|
||||
data.get("model")
|
||||
or data.get("custom_llm_provider")
|
||||
or data.get("litellm_credential_name")
|
||||
):
|
||||
return data
|
||||
|
||||
model = request.query_params.get("model") or request.headers.get("x-litellm-model")
|
||||
if model:
|
||||
data["model"] = model
|
||||
|
||||
return data
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/vector_stores/{vector_store_id:path}/search",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
|
|
@ -312,6 +337,7 @@ async def vector_store_retrieve(
|
|||
data = await _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
data=data, vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
data = _update_request_data_with_model_from_request(data=data, request=request)
|
||||
|
||||
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
|
|
@ -511,6 +537,7 @@ async def vector_store_delete(
|
|||
data = await _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
data=data, vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
data = _update_request_data_with_model_from_request(data=data, request=request)
|
||||
|
||||
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -439,6 +439,28 @@ def test_get_model_from_request_includes_file_endpoint_header_model():
|
|||
)
|
||||
|
||||
|
||||
def test_get_model_from_request_includes_vector_store_query_model():
|
||||
assert (
|
||||
get_model_from_request(
|
||||
request_data={"vector_store_id": "vs_plain_openai"},
|
||||
route="/v1/vector_stores/{vector_store_id}",
|
||||
request_query_params={"model": "restricted-model"},
|
||||
)
|
||||
== "restricted-model"
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_from_request_includes_vector_store_header_model():
|
||||
assert (
|
||||
get_model_from_request(
|
||||
request_data={"vector_store_id": "vs_plain_openai"},
|
||||
route="/v1/vector_stores/{vector_store_id}",
|
||||
request_headers={"x-litellm-model": "restricted-model"},
|
||||
)
|
||||
== "restricted-model"
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_from_request_ignores_routing_header_on_standard_llm_routes():
|
||||
assert (
|
||||
get_model_from_request(
|
||||
|
|
|
|||
|
|
@ -55,6 +55,157 @@ async def test_vector_store_search_forces_path_id_over_body_id():
|
|||
assert captured_data["vector_store_id"] == "vs_path_allowed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_retrieve_uses_query_model_for_routing():
|
||||
from litellm.proxy.vector_store_endpoints.endpoints import vector_store_retrieve
|
||||
|
||||
captured_data = {}
|
||||
|
||||
async def fake_base_process(self, **kwargs):
|
||||
captured_data.update(self.data)
|
||||
return {"ok": True}
|
||||
|
||||
request = _mock_request()
|
||||
request.method = "GET"
|
||||
request.query_params = {"model": "gpt-5.2"}
|
||||
request.headers = {"x-litellm-model": "header-model"}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.endpoints.get_litellm_managed_vector_store",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
|
||||
new=fake_base_process,
|
||||
),
|
||||
):
|
||||
response = await vector_store_retrieve(
|
||||
request=request,
|
||||
vector_store_id="vs_plain_openai",
|
||||
fastapi_response=Response(),
|
||||
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
|
||||
)
|
||||
|
||||
assert response == {"ok": True}
|
||||
assert captured_data["vector_store_id"] == "vs_plain_openai"
|
||||
assert captured_data["model"] == "gpt-5.2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_delete_uses_header_model_for_routing():
|
||||
from litellm.proxy.vector_store_endpoints.endpoints import vector_store_delete
|
||||
|
||||
captured_data = {}
|
||||
|
||||
async def fake_base_process(self, **kwargs):
|
||||
captured_data.update(self.data)
|
||||
return {"ok": True}
|
||||
|
||||
request = _mock_request()
|
||||
request.method = "DELETE"
|
||||
request.headers = {"x-litellm-model": "gpt-5.2"}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.endpoints.get_litellm_managed_vector_store",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
|
||||
new=fake_base_process,
|
||||
),
|
||||
):
|
||||
response = await vector_store_delete(
|
||||
request=request,
|
||||
vector_store_id="vs_plain_openai",
|
||||
fastapi_response=Response(),
|
||||
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
|
||||
)
|
||||
|
||||
assert response == {"ok": True}
|
||||
assert captured_data["vector_store_id"] == "vs_plain_openai"
|
||||
assert captured_data["model"] == "gpt-5.2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_delete_uses_query_model_for_routing():
|
||||
from litellm.proxy.vector_store_endpoints.endpoints import vector_store_delete
|
||||
|
||||
captured_data = {}
|
||||
|
||||
async def fake_base_process(self, **kwargs):
|
||||
captured_data.update(self.data)
|
||||
return {"ok": True}
|
||||
|
||||
request = _mock_request()
|
||||
request.method = "DELETE"
|
||||
request.query_params = {"model": "gpt-5.2"}
|
||||
request.headers = {"x-litellm-model": "header-model"}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.endpoints.get_litellm_managed_vector_store",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
|
||||
new=fake_base_process,
|
||||
),
|
||||
):
|
||||
response = await vector_store_delete(
|
||||
request=request,
|
||||
vector_store_id="vs_plain_openai",
|
||||
fastapi_response=Response(),
|
||||
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
|
||||
)
|
||||
|
||||
assert response == {"ok": True}
|
||||
assert captured_data["vector_store_id"] == "vs_plain_openai"
|
||||
assert captured_data["model"] == "gpt-5.2"
|
||||
|
||||
|
||||
def test_vector_store_model_from_request_preserves_existing_model():
|
||||
from litellm.proxy.vector_store_endpoints.endpoints import (
|
||||
_update_request_data_with_model_from_request,
|
||||
)
|
||||
|
||||
request = _mock_request()
|
||||
request.query_params = {"model": "query-model"}
|
||||
request.headers = {"x-litellm-model": "header-model"}
|
||||
data = {"model": "registry-model"}
|
||||
|
||||
updated_data = _update_request_data_with_model_from_request(
|
||||
data=data, request=request
|
||||
)
|
||||
|
||||
assert updated_data is data
|
||||
assert updated_data["model"] == "registry-model"
|
||||
|
||||
|
||||
def test_vector_store_model_from_request_preserves_registry_routing():
|
||||
from litellm.proxy.vector_store_endpoints.endpoints import (
|
||||
_update_request_data_with_model_from_request,
|
||||
)
|
||||
|
||||
request = _mock_request()
|
||||
request.query_params = {"model": "query-model"}
|
||||
request.headers = {"x-litellm-model": "header-model"}
|
||||
data = {
|
||||
"custom_llm_provider": "openai",
|
||||
"litellm_credential_name": "vector-store-credential",
|
||||
}
|
||||
|
||||
updated_data = _update_request_data_with_model_from_request(
|
||||
data=data, request=request
|
||||
)
|
||||
|
||||
assert updated_data is data
|
||||
assert "model" not in updated_data
|
||||
assert updated_data["custom_llm_provider"] == "openai"
|
||||
assert updated_data["litellm_credential_name"] == "vector-store-credential"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_file_create_forces_path_id_over_body_id():
|
||||
from litellm.proxy.vector_store_files_endpoints.endpoints import (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue