test: migrate phase 15 legacy tests to tests/unit

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yuneng 2026-09-20 11:55:43 +00:00
parent 9e39c751ed
commit 78a751c049
27 changed files with 290 additions and 403 deletions

View file

View file

@ -29,9 +29,7 @@ RUST_RULES: Final[Rules] = (Rule(Route.MESSAGES, Rollout.RUST_REQUIRED),)
def messages_binding(native: NativeMessages | None) -> NativeBinding[NativeMessages]:
binding: Final[NativeBinding[NativeMessages]] = NativeBinding(
"anthropic_messages_handler", validate=lambda _: None
)
binding: Final[NativeBinding[NativeMessages]] = NativeBinding("anthropic_messages_handler", validate=lambda _: None)
binding.override(native)
return binding
@ -99,7 +97,8 @@ async def test_async_python_route_forwards_original_call_shape() -> None:
expected: Final = response()
async def python(
*call_args: object, **call_kwargs: object # kwargs-ok: records call shape
*call_args: object,
**call_kwargs: object, # kwargs-ok: records call shape
) -> AnthropicMessagesResponse:
captured.append((call_args, call_kwargs))
return expected
@ -217,7 +216,9 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
def python(*call_args: object, **call_kwargs: object) -> AnthropicMessagesResponse: # kwargs-ok: records invalid call
def python(
*call_args: object, **call_kwargs: object
) -> AnthropicMessagesResponse: # kwargs-ok: records invalid call
captured.append((call_args, call_kwargs))
return expected

View file

@ -5,7 +5,7 @@ Tests for backend domain models.
from datetime import datetime, timezone
import pytest
from pydantic import BaseModel, TypeAdapter
from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm.models.access_group import LiteLLM_AccessGroupTable
from litellm.models.autorouter_session import LiteLLM_AutoRouterSession
@ -19,7 +19,6 @@ from litellm.models.credentials import CreateCredentialItem, CredentialItem
from litellm.models.end_user import LiteLLM_EndUserTable
from litellm.models.managed_files import (
LiteLLM_ManagedFileTable,
LiteLLM_ManagedObjectTable,
LiteLLM_ManagedVectorStoresTable,
)
from litellm.models.mcp_server import LiteLLM_MCPServerTable
@ -41,7 +40,6 @@ from litellm.models.verification_token import (
LiteLLM_DeletedVerificationToken,
LiteLLM_VerificationToken,
)
from pydantic import ValidationError
class TestBudget:
@ -121,9 +119,7 @@ class TestCredentials:
assert item.credential_values is None
def test_create_credential_item_requires_values_or_model_id(self):
with pytest.raises(
ValueError, match="Either credential_values or model_id must be set"
):
with pytest.raises(ValueError, match="Either credential_values or model_id must be set"):
CreateCredentialItem(credential_name="bad", credential_info={})
@ -141,12 +137,8 @@ class TestModel:
assert model.team_public_model_name == "my-gpt4"
def test_is_blocked(self):
model_blocked = LiteLLM_ProxyModelTable(
model_id="m1", model_name="test", litellm_params={}, blocked=True
)
model_unblocked = LiteLLM_ProxyModelTable(
model_id="m2", model_name="test", litellm_params={}, blocked=False
)
model_blocked = LiteLLM_ProxyModelTable(model_id="m1", model_name="test", litellm_params={}, blocked=True)
model_unblocked = LiteLLM_ProxyModelTable(model_id="m2", model_name="test", litellm_params={}, blocked=False)
assert model_blocked.is_blocked
assert not model_unblocked.is_blocked
@ -188,9 +180,7 @@ class TestModel:
assert model.blocked is True
def test_team_helpers_none_when_no_model_info(self):
model = LiteLLM_ProxyModelTable(
model_id="m1", model_name="gpt-4", litellm_params={}, model_info=None
)
model = LiteLLM_ProxyModelTable(model_id="m1", model_name="gpt-4", litellm_params={}, model_info=None)
assert model.team_id is None
assert model.team_public_model_name is None
@ -292,9 +282,7 @@ class TestTeam:
assert team.model_max_budget == {"gpt-4": 5.0}
def test_cached_team(self):
cached = LiteLLM_TeamTableCachedObj(
team_id="t1", last_refreshed_at=1234567890.0
)
cached = LiteLLM_TeamTableCachedObj(team_id="t1", last_refreshed_at=1234567890.0)
assert cached.last_refreshed_at == 1234567890.0
def test_deleted_team(self):
@ -345,9 +333,7 @@ class TestUser:
assert "password" not in user.model_dump()
assert "password" not in user.model_dump_json()
with_keys = LiteLLM_UserTableWithKeyCount(
user_id="u1", user_email="a@b.c", password=secret, key_count=2
)
with_keys = LiteLLM_UserTableWithKeyCount(user_id="u1", user_email="a@b.c", password=secret, key_count=2)
assert with_keys.password == secret
assert "password" not in with_keys.model_dump()
assert "password" not in with_keys.model_dump_json()
@ -479,9 +465,7 @@ class TestEndUserTable:
class TestBudgetTableFull:
def test_full_adds_server_managed_fields(self):
now = datetime.now()
budget = LiteLLM_BudgetTableFull(
budget_id="b1", max_budget=10.0, created_at=now, budget_reset_at=now
)
budget = LiteLLM_BudgetTableFull(budget_id="b1", max_budget=10.0, created_at=now, budget_reset_at=now)
assert budget.created_at == now
assert budget.budget_reset_at == now
assert budget.max_budget == 10.0
@ -493,9 +477,7 @@ class TestBudgetTableFull:
class TestTeamMemberTable:
def test_tracks_user_within_team(self):
member = LiteLLM_TeamMemberTable(
user_id="u1", team_id="t1", spend=3.0, budget_id="b1", max_budget=5.0
)
member = LiteLLM_TeamMemberTable(user_id="u1", team_id="t1", spend=3.0, budget_id="b1", max_budget=5.0)
assert member.user_id == "u1"
assert member.team_id == "t1"
assert member.spend == 3.0
@ -585,9 +567,7 @@ class TestSpendLogs:
assert log.updated_at == updated_at
def test_error_logs_creation(self):
log = LiteLLM_ErrorLogs(
request_id="r1", startTime=None, endTime=None, status_code="500"
)
log = LiteLLM_ErrorLogs(request_id="r1", startTime=None, endTime=None, status_code="500")
assert log.request_id == "r1"
assert log.status_code == "500"
@ -603,12 +583,6 @@ class TestManagedTables:
assert table.model_mappings == {"gpt-4": "file-abc"}
assert table.flat_model_file_ids == ["file-abc"]
def test_managed_object_table_requires_purpose(self):
with pytest.raises(ValidationError):
LiteLLM_ManagedObjectTable(
unified_object_id="o1", model_object_id="m1", file_object={}
)
def test_managed_vector_stores_table(self):
table = LiteLLM_ManagedVectorStoresTable(
vector_store_id="vs1",

View file

View file

@ -73,9 +73,7 @@ class TestConvertFileDocumentToUrlDocument:
tmp_path = Path(f.name)
try:
result = convert_file_document_to_url_document(
{"type": "file", "file": tmp_path}
)
result = convert_file_document_to_url_document({"type": "file", "file": tmp_path})
assert result["type"] == "document_url"
assert result["document_url"].startswith("data:application/pdf;base64,")
@ -95,9 +93,7 @@ class TestConvertFileDocumentToUrlDocument:
tmp_path = Path(f.name)
try:
result = convert_file_document_to_url_document(
{"type": "file", "file": tmp_path}
)
result = convert_file_document_to_url_document({"type": "file", "file": tmp_path})
assert result["type"] == "image_url"
assert result["image_url"].startswith("data:image/png;base64,")
@ -112,9 +108,7 @@ class TestConvertFileDocumentToUrlDocument:
request handler the value is attacker-controlled, and opening it as
a path is an arbitrary local file read on the proxy host."""
with pytest.raises(ValueError, match="does not accept bare str values"):
convert_file_document_to_url_document(
{"type": "file", "file": "/etc/passwd"}
)
convert_file_document_to_url_document({"type": "file", "file": "/etc/passwd"})
def test_should_convert_pathlib_path(self):
"""pathlib.Path objects should work the same as string paths."""
@ -126,9 +120,7 @@ class TestConvertFileDocumentToUrlDocument:
tmp_path = Path(f.name)
try:
result = convert_file_document_to_url_document(
{"type": "file", "file": tmp_path}
)
result = convert_file_document_to_url_document({"type": "file", "file": tmp_path})
assert result["type"] == "document_url"
assert result["document_url"].startswith("data:application/pdf;base64,")
@ -139,9 +131,7 @@ class TestConvertFileDocumentToUrlDocument:
"""Raw bytes should be converted using a fallback MIME type."""
content = b"raw bytes content"
result = convert_file_document_to_url_document(
{"type": "file", "file": content}
)
result = convert_file_document_to_url_document({"type": "file", "file": content})
assert result["type"] == "document_url"
assert "base64," in result["document_url"]
@ -164,9 +154,7 @@ class TestConvertFileDocumentToUrlDocument:
"""Raw bytes with an image MIME type should produce type=image_url."""
content = b"raw image content"
result = convert_file_document_to_url_document(
{"type": "file", "file": content, "mime_type": "image/jpeg"}
)
result = convert_file_document_to_url_document({"type": "file", "file": content, "mime_type": "image/jpeg"})
assert result["type"] == "image_url"
assert result["image_url"].startswith("data:image/jpeg;base64,")
@ -176,9 +164,7 @@ class TestConvertFileDocumentToUrlDocument:
content = b"file-like content"
file_obj = BytesIO(content)
result = convert_file_document_to_url_document(
{"type": "file", "file": file_obj}
)
result = convert_file_document_to_url_document({"type": "file", "file": file_obj})
assert result["type"] == "document_url"
assert "base64," in result["document_url"]
@ -189,9 +175,7 @@ class TestConvertFileDocumentToUrlDocument:
file_obj = BytesIO(content)
file_obj.name = "test_image.png"
result = convert_file_document_to_url_document(
{"type": "file", "file": file_obj}
)
result = convert_file_document_to_url_document({"type": "file", "file": file_obj})
assert result["type"] == "image_url"
assert result["image_url"].startswith("data:image/png;base64,")
@ -204,9 +188,7 @@ class TestConvertFileDocumentToUrlDocument:
def test_should_raise_error_for_nonexistent_pathlib_path(self):
"""Non-existent pathlib.Path should raise FileNotFoundError."""
with pytest.raises(FileNotFoundError, match="File not found"):
convert_file_document_to_url_document(
{"type": "file", "file": Path("/nonexistent/path/to/file.pdf")}
)
convert_file_document_to_url_document({"type": "file", "file": Path("/nonexistent/path/to/file.pdf")})
def test_should_raise_error_for_empty_file(self):
"""Empty file should raise ValueError."""
@ -215,9 +197,7 @@ class TestConvertFileDocumentToUrlDocument:
try:
with pytest.raises(ValueError, match="File is empty"):
convert_file_document_to_url_document(
{"type": "file", "file": tmp_path}
)
convert_file_document_to_url_document({"type": "file", "file": tmp_path})
finally:
os.unlink(str(tmp_path))
@ -248,9 +228,7 @@ class TestConvertFileDocumentToUrlDocument:
tmp_path = Path(f.name)
try:
result = convert_file_document_to_url_document(
{"type": "file", "file": tmp_path, "mime_type": "image/png"}
)
result = convert_file_document_to_url_document({"type": "file", "file": tmp_path, "mime_type": "image/png"})
assert result["type"] == "image_url"
assert result["image_url"].startswith("data:image/png;base64,")
@ -477,9 +455,7 @@ class TestProxySecurityGuard:
result = await self._parse_multipart(mock_request)
assert result["document"]["type"] == "document_url"
assert result["document"]["document_url"].startswith(
"data:application/pdf;base64,"
)
assert result["document"]["document_url"].startswith("data:application/pdf;base64,")
assert result["model"] == "mistral/mistral-ocr-latest"

View file

View file

@ -21,9 +21,7 @@ def _make_mock_response(status_code: int, body: bytes, headers: dict = None): #
def _raise_for_status():
if status_code >= 400:
request = httpx.Request(
"POST", "https://azure.example.com/openai/responses"
)
request = httpx.Request("POST", "https://azure.example.com/openai/responses")
real_response = httpx.Response(
status_code=status_code,
content=body,
@ -55,16 +53,15 @@ def _make_mock_logging_obj():
async def test_async_streaming_429_raises():
"""429 from upstream should raise HTTPStatusError, not yield error bytes."""
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
error_body = json.dumps(
{"error": {"code": "429", "message": "Rate limit exceeded."}}
).encode()
error_body = json.dumps({"error": {"code": "429", "message": "Rate limit exceeded."}}).encode()
mock_response = _make_mock_response(429, error_body)
async def response_coro():
return mock_response
chunks = []
async def _drain():
async for chunk in AsyncPassthroughStreamingResponse(
response=response_coro(),
@ -84,15 +81,13 @@ async def test_async_streaming_429_raises():
async def test_async_streaming_500_raises():
"""500 from upstream should also raise, not yield error bytes."""
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
error_body = json.dumps(
{"error": {"code": "500", "message": "Internal server error"}}
).encode()
error_body = json.dumps({"error": {"code": "500", "message": "Internal server error"}}).encode()
mock_response = _make_mock_response(500, error_body)
async def response_coro():
return mock_response
with pytest.raises(httpx.HTTPStatusError) as exc_info:
async for _ in AsyncPassthroughStreamingResponse(
response=response_coro(),
@ -100,7 +95,7 @@ async def test_async_streaming_500_raises():
provider_config=MagicMock(),
):
pass
assert exc_info.value.response.status_code == 500

View file

@ -3,14 +3,9 @@ from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi.testclient import TestClient
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.passthrough.main import allm_passthrough_route, llm_passthrough_route
@ -37,10 +32,7 @@ def test_llm_passthrough_route():
client=client,
)
assert (
mock_post.call_args.kwargs["request"].url
== "http://localhost:8090/v1/chat/completions"
)
assert mock_post.call_args.kwargs["request"].url == "http://localhost:8090/v1/chat/completions"
assert response.status_code == 200
assert response.json == {"message": "Hello, world!"}
@ -74,12 +66,9 @@ def test_bedrock_application_inference_profile_url_encoding():
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("test-model", "bedrock", "test-key", "test-base"),
),
patch.object(
client.client, "send", return_value=MagicMock(status_code=200)
) as mock_send,
patch.object(client.client, "send", return_value=MagicMock(status_code=200)),
patch.object(client.client, "build_request") as mock_build_request,
):
# Mock logging object
mock_logging_obj = MagicMock()
mock_logging_obj.update_environment_variables = MagicMock()
@ -132,12 +121,9 @@ def test_bedrock_non_application_inference_profile_no_encoding():
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("test-model", "bedrock", "test-key", "test-base"),
),
patch.object(
client.client, "send", return_value=MagicMock(status_code=200)
) as mock_send,
patch.object(client.client, "send", return_value=MagicMock(status_code=200)),
patch.object(client.client, "build_request") as mock_build_request,
):
# Mock logging object
mock_logging_obj = MagicMock()
mock_logging_obj.update_environment_variables = MagicMock()
@ -202,7 +188,6 @@ def test_update_stream_param_based_on_request_body():
@pytest.fixture
def mock_request():
"""Create a mock request with headers"""
from typing import Optional
class QueryParams:
def __init__(self):
@ -215,9 +200,7 @@ def mock_request():
return self._dict.items()
class MockRequest:
def __init__(
self, headers=None, method="POST", request_body: Optional[dict] = None
):
def __init__(self, headers=None, method="POST", request_body: dict | None = None):
self.headers = headers or {}
self.query_params = QueryParams()
self.method = method
@ -245,9 +228,7 @@ def mock_user_api_key_dict():
@pytest.mark.asyncio
async def test_pass_through_request_stream_param_override(
mock_request, mock_user_api_key_dict
):
async def test_pass_through_request_stream_param_override(mock_request, mock_user_api_key_dict):
"""
Test that when stream=None is passed as parameter but stream=True
is in request body, the request body value takes precedence and
@ -346,9 +327,7 @@ async def test_pass_through_request_stream_param_override(
@pytest.mark.asyncio
async def test_pass_through_request_stream_param_no_override(
mock_request, mock_user_api_key_dict
):
async def test_pass_through_request_stream_param_no_override(mock_request, mock_user_api_key_dict):
"""
Test that when stream=False is passed as parameter and no stream
is in request body, the function parameter is used and
@ -448,15 +427,11 @@ def test_azure_with_custom_api_base_and_key():
# Mock the provider config and its methods
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
httpx.URL(
"https://my-custom-base/openai/deployments/gpt-4.1/chat/completions?api-version=2024-02-01"
),
httpx.URL("https://my-custom-base/openai/deployments/gpt-4.1/chat/completions?api-version=2024-02-01"),
"https://my-custom-base",
)
mock_provider_config.get_api_key.return_value = "my-custom-key"
mock_provider_config.validate_environment.return_value = {
"api-key": "my-custom-key"
}
mock_provider_config.validate_environment.return_value = {"api-key": "my-custom-key"}
mock_provider_config.sign_request.return_value = (
{"api-key": "my-custom-key"},
None,
@ -484,13 +459,10 @@ def test_azure_with_custom_api_base_and_key():
patch.object(
client.client,
"send",
return_value=MagicMock(
status_code=200, json=lambda: {"id": "chatcmpl-123", "choices": []}
),
) as mock_send,
return_value=MagicMock(status_code=200, json=lambda: {"id": "chatcmpl-123", "choices": []}),
),
patch.object(client.client, "build_request") as mock_build_request,
):
# Mock logging object
mock_logging_obj = MagicMock()
mock_logging_obj.update_environment_variables = MagicMock()
@ -541,9 +513,7 @@ def test_content_param_forwarded_to_build_request():
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
httpx.URL(
"https://my-azure.openai.azure.com/openai/deployments/gpt-4/chat/completions"
),
httpx.URL("https://my-azure.openai.azure.com/openai/deployments/gpt-4/chat/completions"),
"https://my-azure.openai.azure.com",
)
mock_provider_config.get_api_key.return_value = "test-key"
@ -575,7 +545,6 @@ def test_content_param_forwarded_to_build_request():
patch.object(client.client, "send", return_value=MagicMock(status_code=200)),
patch.object(client.client, "build_request") as mock_build_request,
):
mock_logging_obj = MagicMock()
mock_logging_obj.update_environment_variables = MagicMock()
@ -656,15 +625,11 @@ async def test_allm_passthrough_route_429_streaming_raises():
"""
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
httpx.URL(
"https://my-azure.openai.azure.com/openai/deployments/gpt-4/responses"
),
httpx.URL("https://my-azure.openai.azure.com/openai/deployments/gpt-4/responses"),
"https://my-azure.openai.azure.com",
)
mock_provider_config.get_api_key.return_value = "fake-azure-key"
mock_provider_config.validate_environment.return_value = {
"api-key": "fake-azure-key"
}
mock_provider_config.validate_environment.return_value = {"api-key": "fake-azure-key"}
mock_provider_config.sign_request.return_value = (
{"api-key": "fake-azure-key"},
None,
@ -752,9 +717,7 @@ def test_llm_passthrough_route_sync_streaming_error_maps_upstream_status():
headers={"content-type": "application/json"},
)
sync_client = HTTPHandler(
client=httpx.Client(transport=httpx.MockTransport(_handler))
)
sync_client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_handler)))
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
@ -762,18 +725,14 @@ def test_llm_passthrough_route_sync_streaming_error_maps_upstream_status():
"https://gigachat.devices.sberbank.ru/api/v1",
)
mock_provider_config.get_api_key.return_value = "fake-key"
mock_provider_config.validate_environment.return_value = {
"Authorization": "Bearer fake-key"
}
mock_provider_config.validate_environment.return_value = {"Authorization": "Bearer fake-key"}
mock_provider_config.sign_request.return_value = (
{"Authorization": "Bearer fake-key"},
None,
)
mock_provider_config.is_streaming_request.return_value = True
mock_provider_config.get_error_class.side_effect = (
lambda error_message, status_code, headers: BaseLLMException(
status_code=status_code, message=error_message, headers=headers
)
mock_provider_config.get_error_class.side_effect = lambda error_message, status_code, headers: BaseLLMException(
status_code=status_code, message=error_message, headers=headers
)
mock_logging_obj = MagicMock()

View file

@ -68,9 +68,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion():
chunks = [b"chunk-1", b"chunk-2", b"chunk-3"]
mock_response = _make_streaming_response(chunks)
mock_response.headers = httpx.Headers(
{"content-type": "application/octet-stream", "x-request-id": "req-123"}
)
mock_response.headers = httpx.Headers({"content-type": "application/octet-stream", "x-request-id": "req-123"})
async def response_coro():
return mock_response
@ -88,7 +86,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion():
received.append(chunk)
assert received == chunks
assert received_response.headers["content-type"] == "application/octet-stream"
assert received_response.headers["x-request-id"] == "req-123"
@ -107,9 +105,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_client_disconnect():
b'{"chunk": 3, "outputTokens": 8}',
]
mock_response = _make_streaming_response(chunks)
mock_response.headers = httpx.Headers(
{"content-type": "application/octet-stream", "x-request-id": "req-123"}
)
mock_response.headers = httpx.Headers({"content-type": "application/octet-stream", "x-request-id": "req-123"})
async def response_coro():
return mock_response
@ -138,17 +134,13 @@ async def test_asyncpassthroughstreamingresponse_does_not_flush_on_4xx():
err_response = MagicMock(spec=httpx.Response)
err_response.status_code = 429
err_response.headers = httpx.Headers(
{"content-type": "application/octet-stream"}
)
err_response.headers = httpx.Headers({"content-type": "application/octet-stream"})
def _raise():
raise httpx.HTTPStatusError(
"429",
request=httpx.Request("POST", "https://example.com"),
response=httpx.Response(
429, request=httpx.Request("POST", "https://example.com")
),
response=httpx.Response(429, request=httpx.Request("POST", "https://example.com")),
)
err_response.raise_for_status = _raise
@ -180,9 +172,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_upstream_exception_w
mock_response.status_code = 200
mock_response.raise_for_status = MagicMock(return_value=None)
mock_response.aclose = AsyncMock()
mock_response.headers = httpx.Headers(
{"content-type": "application/octet-stream", "x-request-id": "req-123"}
)
mock_response.headers = httpx.Headers({"content-type": "application/octet-stream", "x-request-id": "req-123"})
async def _aiter_bytes_then_raise():
for c in partial_chunks:
@ -197,6 +187,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_upstream_exception_w
mock_logging_obj = _make_logging_obj()
received = []
async def _drain():
async for chunk in AsyncPassthroughStreamingResponse(
response=response_coro(),
@ -222,9 +213,7 @@ def test_passthroughstreamingresponse_flushes_on_normal_completion():
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = httpx.Headers(
{"content-type": "application/octet-stream", "x-request-id": "req-123"}
)
mock_response.headers = httpx.Headers({"content-type": "application/octet-stream", "x-request-id": "req-123"})
def _iter_bytes():
yield from chunks
@ -258,9 +247,7 @@ def test_passthroughstreamingresponse_flushes_on_early_close():
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = httpx.Headers(
{"content-type": "application/octet-stream", "x-request-id": "req-123"}
)
mock_response.headers = httpx.Headers({"content-type": "application/octet-stream", "x-request-id": "req-123"})
def _iter_bytes():
yield from chunks

View file

View file

@ -21,10 +21,14 @@ class _RecordingRouter:
def _ingestion(embedding=REQUEST_EMBEDDING, router=None, **vector_store):
vector_store_options = {"custom_llm_provider": "s3_vectors", "aws_region_name": "us-west-2", **vector_store}
ingest_options = {"vector_store": vector_store_options} if embedding is None else {
"embedding": embedding,
"vector_store": vector_store_options,
}
ingest_options = (
{"vector_store": vector_store_options}
if embedding is None
else {
"embedding": embedding,
"vector_store": vector_store_options,
}
)
return S3VectorsRAGIngestion(ingest_options=ingest_options, router=router)

View file

@ -12,6 +12,15 @@ from litellm.realtime_api import main as realtime_main
from litellm.realtime_api.main import _with_resolved_session_model
@pytest.fixture
def local_model_cost_map(monkeypatch):
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
litellm.get_model_info.cache_clear()
yield
litellm.get_model_info.cache_clear()
class FakeLogging:
def update_from_kwargs(self, **kwargs):
pass
@ -502,8 +511,8 @@ async def test_arealtime_azure_env_beta_protocol_wins_over_a_ga_client(monkeypat
async def _vertex_provider_config_for(monkeypatch, model: str, vertex_location: str | None):
from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig
from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import VertexChirpRealtimeConfig
from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig
captured: dict[str, object] = {}

View file

@ -78,17 +78,11 @@ class MockTable:
record_data = dict(data)
if self._pk_field and self._pk_field not in record_data:
record_data[self._pk_field] = f"{self._pk_field}-{len(self._records)}"
key = (
record_data.get(self._pk_field)
if self._pk_field
else record_data.get("id", str(len(self._records)))
)
key = record_data.get(self._pk_field) if self._pk_field else record_data.get("id", str(len(self._records)))
self._records[key] = record_data
return MockRecord(record_data)
async def update(
self, where: Dict[str, Any], data: Dict[str, Any]
) -> Optional[MockRecord]:
async def update(self, where: Dict[str, Any], data: Dict[str, Any]) -> Optional[MockRecord]:
key_field = list(where.keys())[0]
key_value = where[key_field]
if key_value in self._records:
@ -140,9 +134,7 @@ class MockPrismaClient:
self.db.litellm_config = MockTable()
self.db.litellm_organizationtable = MockTable()
self.db.litellm_projecttable = MockTable(pk_field="project_id")
self.db.litellm_objectpermissiontable = MockTable(
pk_field="object_permission_id"
)
self.db.litellm_objectpermissiontable = MockTable(pk_field="object_permission_id")
self.db.litellm_credentialstable = MockTable()
@ -200,9 +192,7 @@ class TestBaseRepository:
prisma_client.db.litellm_budgettable._records = {
"b1": {"budget_id": "b1", "max_budget": 100.0},
}
budgets = await repo.find_many(
where={"budget_id": "b1"}, skip=0, take=10, order={"budget_id": "asc"}
)
budgets = await repo.find_many(where={"budget_id": "b1"}, skip=0, take=10, order={"budget_id": "asc"})
assert len(budgets) == 1
def test_record_to_dict_branches(self):
@ -1518,9 +1508,7 @@ class TestVerificationTokenRepositoryExtended:
class MockTx:
def __init__(self, client):
self.litellm_deletedverificationtoken = (
client.db.litellm_deletedverificationtoken
)
self.litellm_deletedverificationtoken = client.db.litellm_deletedverificationtoken
self.litellm_verificationtoken = client.db.litellm_verificationtoken
async def __aenter__(self):
@ -1563,9 +1551,7 @@ class TestVerificationTokenRepositoryExtended:
class MockTx:
def __init__(self, client):
self.litellm_deletedverificationtoken = (
client.db.litellm_deletedverificationtoken
)
self.litellm_deletedverificationtoken = client.db.litellm_deletedverificationtoken
self.litellm_verificationtoken = client.db.litellm_verificationtoken
async def __aenter__(self):
@ -1578,9 +1564,7 @@ class TestVerificationTokenRepositoryExtended:
await repo.delete_token("sk-arch", deleted_by="admin")
archived = list(
repo._prisma_client.db.litellm_deletedverificationtoken._records.values()
)[0]
archived = list(repo._prisma_client.db.litellm_deletedverificationtoken._records.values())[0]
assert isinstance(archived["aliases"], str)
assert json.loads(archived["aliases"]) == {"a": "b"}
@ -1599,9 +1583,7 @@ class TestVerificationTokenRepositoryExtended:
):
assert relation_field not in archived
assert (
"sk-arch" not in repo._prisma_client.db.litellm_verificationtoken._records
)
assert "sk-arch" not in repo._prisma_client.db.litellm_verificationtoken._records
@pytest.mark.asyncio
async def test_find_by_id_maps_org_and_budget_columns(self, repo):
@ -1977,9 +1959,7 @@ class TestDomainModelExtended:
DomainModel.from_db_record(None)
def test_from_db_record_dict(self):
model = _SampleDomainModel.from_db_record(
{"budget_id": "b1", "max_budget": 100.0}
)
model = _SampleDomainModel.from_db_record({"budget_id": "b1", "max_budget": 100.0})
assert model.budget_id == "b1"
def test_from_db_record_model_dump(self):
@ -2174,9 +2154,7 @@ class TestPrismaTableRepository:
assert self.CONFIG_SYNCED_TABLE_NAMES <= seen
def _json_path_equals(
metadata: Optional[Dict[str, Any]], path: List[str], expected: Any
) -> bool:
def _json_path_equals(metadata: Optional[Dict[str, Any]], path: List[str], expected: Any) -> bool:
"""Reproduce Postgres jsonb path-equals semantics: a missing path yields
SQL NULL, which never matches `equals`."""
value: Any = metadata
@ -2201,11 +2179,7 @@ class _ScimAwareUserTable:
json_filter = where["metadata"]
path = json_filter["path"]
expected = getattr(json_filter["equals"], "data", json_filter["equals"])
return sum(
1
for metadata in self._metadatas
if _json_path_equals(metadata, path, expected)
)
return sum(1 for metadata in self._metadatas if _json_path_equals(metadata, path, expected))
class TestCountBillableUsers:

View file

@ -14,6 +14,30 @@ from litellm.router_utils.pre_call_checks.deployment_affinity_check import (
)
@pytest.fixture(autouse=True)
def isolate_litellm_router_state():
saved = {
name: getattr(litellm, name).copy()
if isinstance(getattr(litellm, name, None), list)
else getattr(litellm, name, None)
for name in (
"callbacks",
"input_callback",
"success_callback",
"failure_callback",
"_async_input_callback",
"_async_success_callback",
"_async_failure_callback",
"model_fallbacks",
"cache",
)
if hasattr(litellm, name)
}
yield
for name, value in saved.items():
setattr(litellm, name, value)
class MockResponse:
def __init__(self, json_data, status_code):
self._json_data = json_data
@ -43,9 +67,7 @@ async def test_async_user_key_affinity_routes_to_same_deployment():
"id": "msg_123",
"status": "completed",
"role": "assistant",
"content": [
{"type": "output_text", "text": "Hello there!", "annotations": []}
],
"content": [{"type": "output_text", "text": "Hello there!", "annotations": []}],
}
],
"parallel_tool_calls": True,
@ -348,9 +370,7 @@ async def test_async_previous_response_id_priority_over_user_key_affinity():
model_group=model_group,
user_key=user_api_key_hash,
)
await router.cache.async_set_cache(
affinity_cache_key, {"model_id": other_model_id}, ttl=3600
)
await router.cache.async_set_cache(affinity_cache_key, {"model_id": other_model_id}, ttl=3600)
# Even though user-key affinity points elsewhere, previous_response_id should pin
# to the deployment that created the original response.
@ -519,9 +539,7 @@ async def test_async_filter_deployments_uses_stable_model_map_key_for_affinity_s
},
{
"model_name": stable_model_map_key,
"litellm_params": {
"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"
},
"litellm_params": {"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"},
"model_info": {"id": "deployment-2"},
},
]
@ -542,9 +560,7 @@ async def test_async_filter_deployments_uses_stable_model_map_key_for_affinity_s
model="some-router-model-group",
healthy_deployments=healthy_deployments,
messages=None,
request_kwargs={
"metadata": {"user_api_key_hash": user_key, "model_group": "alias-group"}
},
request_kwargs={"metadata": {"user_api_key_hash": user_key, "model_group": "alias-group"}},
parent_otel_span=None,
)
@ -580,9 +596,7 @@ async def test_async_filter_deployments_falls_back_when_cached_deployment_is_unh
},
{
"model_name": stable_model_map_key,
"litellm_params": {
"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"
},
"litellm_params": {"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"},
"model_info": {"id": "deployment-2"},
},
]
@ -618,9 +632,7 @@ async def test_async_filter_deployments_does_not_pin_when_target_order_is_set():
},
{
"model_name": stable_model_map_key,
"litellm_params": {
"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"
},
"litellm_params": {"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"},
"model_info": {"id": "deployment-2"},
},
]
@ -660,9 +672,7 @@ async def test_async_user_key_affinity_ttl_expiry_allows_reroute():
},
{
"model_name": stable_model_map_key,
"litellm_params": {
"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"
},
"litellm_params": {"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"},
"model_info": {"id": "deployment-2"},
},
]
@ -706,9 +716,7 @@ def test_cache_key_does_not_double_hash_user_api_key_hash():
The affinity cache key should not hash it again.
"""
user_api_key_hash = (
"b95b015b66dd02a1c14e1e0a8729211f8ee53ec962658764f4cf58546c2c68e1"
)
user_api_key_hash = "b95b015b66dd02a1c14e1e0a8729211f8ee53ec962658764f4cf58546c2c68e1"
key = DeploymentAffinityCheck.get_affinity_cache_key(
model_group="any-model-group",
user_key=user_api_key_hash,
@ -746,9 +754,7 @@ def test_get_effective_flags_returns_per_group_config():
assert session_id is True
# unconfigured-model: falls back to global flags
user_key, responses_api, session_id = callback._get_effective_flags(
"unconfigured-model"
)
user_key, responses_api, session_id = callback._get_effective_flags("unconfigured-model")
assert user_key is True
assert responses_api is True
assert session_id is False
@ -980,12 +986,8 @@ async def test_model_group_affinity_config_overrides_global():
]
# Set up user-key affinity cache for claude-3
cache_key = DeploymentAffinityCheck.get_affinity_cache_key(
model_group=stable_model_map_key, user_key=user_key
)
await callback.cache.async_set_cache(
cache_key, {"model_id": "deployment-1"}, ttl=60
)
cache_key = DeploymentAffinityCheck.get_affinity_cache_key(model_group=stable_model_map_key, user_key=user_key)
await callback.cache.async_set_cache(cache_key, {"model_id": "deployment-1"}, ttl=60)
# claude-3 has per-group config (session_affinity only), so user-key affinity
# should NOT apply even though it's globally enabled
@ -1050,7 +1052,7 @@ async def test_async_jwt_user_affinity_routes_to_same_deployment():
return seq[0]
return seq[1] if len(seq) > 1 else seq[0]
with patch( # test-quality-ok: simple-shuffle has no injectable RNG; forcing the other pick is what proves the pin overrides the strategy
with patch(
"litellm.router_strategy.simple_shuffle.random.choice",
side_effect=deterministic_choice,
):

View file

@ -25,6 +25,31 @@ from litellm.models.credentials import CredentialItem
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import ResponsesAPIResponse
@pytest.fixture(autouse=True)
def isolate_litellm_router_state():
saved = {
name: getattr(litellm, name).copy()
if isinstance(getattr(litellm, name, None), list)
else getattr(litellm, name, None)
for name in (
"callbacks",
"input_callback",
"success_callback",
"failure_callback",
"_async_input_callback",
"_async_success_callback",
"_async_failure_callback",
"model_fallbacks",
"cache",
)
if hasattr(litellm, name)
}
yield
for name, value in saved.items():
setattr(litellm, name, value)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@ -1088,21 +1113,19 @@ def test_boundary_key_resolves_missing_values_from_named_credential():
EncryptedContentAffinityCheck,
)
with (
patch.object( # test-quality-ok: credential registry is the direct dependency under test
litellm,
"credential_list",
[
CredentialItem(
credential_name="account-a",
credential_values={
"api_base": "https://account-a.example.com",
"api_key": "credential-key-a",
},
credential_info={},
)
],
)
with patch.object(
litellm,
"credential_list",
[
CredentialItem(
credential_name="account-a",
credential_values={
"api_base": "https://account-a.example.com",
"api_key": "credential-key-a",
},
credential_info={},
)
],
):
boundary = EncryptedContentAffinityCheck._encryption_boundary_key({"litellm_credential_name": "account-a"})
@ -1114,21 +1137,19 @@ def test_boundary_key_matches_named_credential_precedence():
EncryptedContentAffinityCheck,
)
with (
patch.object( # test-quality-ok: credential registry is the direct dependency under test
litellm,
"credential_list",
[
CredentialItem(
credential_name="account-a",
credential_values={
"api_base": "https://credential.example.com",
"api_key": "credential-key-a",
},
credential_info={},
)
],
)
with patch.object(
litellm,
"credential_list",
[
CredentialItem(
credential_name="account-a",
credential_values={
"api_base": "https://credential.example.com",
"api_key": "credential-key-a",
},
credential_info={},
)
],
):
boundary = EncryptedContentAffinityCheck._encryption_boundary_key(
{
@ -1146,21 +1167,19 @@ def test_boundary_key_resolves_credential_when_explicit_values_are_empty():
EncryptedContentAffinityCheck,
)
with (
patch.object( # test-quality-ok: credential registry is the direct dependency under test
litellm,
"credential_list",
[
CredentialItem(
credential_name="account-a",
credential_values={
"api_base": "https://credential.example.com",
"api_key": "credential-key-a",
},
credential_info={},
)
],
)
with patch.object(
litellm,
"credential_list",
[
CredentialItem(
credential_name="account-a",
credential_values={
"api_base": "https://credential.example.com",
"api_key": "credential-key-a",
},
credential_info={},
)
],
):
boundary = EncryptedContentAffinityCheck._encryption_boundary_key(
{
@ -1178,37 +1197,35 @@ def test_boundary_fallback_matches_deployments_with_same_named_credential_values
EncryptedContentAffinityCheck,
)
with (
patch.object( # test-quality-ok: credential registry is the direct dependency under test
litellm,
"credential_list",
[
CredentialItem(
credential_name="account-a",
credential_values={
"api_base": "https://account-a.example.com",
"api_key": "credential-key-a",
},
credential_info={},
),
CredentialItem(
credential_name="account-a-peer",
credential_values={
"api_base": "https://account-a.example.com",
"api_key": "credential-key-a",
},
credential_info={},
),
CredentialItem(
credential_name="account-b",
credential_values={
"api_base": "https://account-b.example.com",
"api_key": "credential-key-b",
},
credential_info={},
),
],
)
with patch.object(
litellm,
"credential_list",
[
CredentialItem(
credential_name="account-a",
credential_values={
"api_base": "https://account-a.example.com",
"api_key": "credential-key-a",
},
credential_info={},
),
CredentialItem(
credential_name="account-a-peer",
credential_values={
"api_base": "https://account-a.example.com",
"api_key": "credential-key-a",
},
credential_info={},
),
CredentialItem(
credential_name="account-b",
credential_values={
"api_base": "https://account-b.example.com",
"api_key": "credential-key-b",
},
credential_info={},
),
],
):
router = litellm.Router(
model_list=[

View file

@ -1,10 +1,9 @@
import asyncio
import copy
from typing import List, cast
from typing import cast
import pytest
import litellm
from litellm.caching.dual_cache import DualCache
from litellm.constants import DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT
@ -22,6 +21,15 @@ MODEL_GROUP_ALIAS = "my-claude-group"
OPUS_4_6_MIN_TOKENS = 4096
@pytest.fixture
def local_model_cost_map(monkeypatch):
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
litellm.get_model_info.cache_clear()
yield
litellm.get_model_info.cache_clear()
@pytest.fixture(autouse=True)
def _local_model_cost_map_autouse(local_model_cost_map):
"""Every test here reads `prompt_cache_min_tokens`, which only the in-repo map
@ -30,8 +38,7 @@ def _local_model_cost_map_autouse(local_model_cost_map):
yield
def _deployments(*models: str) -> List[dict]:
def _deployments(*models: str) -> list[dict]:
return [
{
"model_name": MODEL_GROUP_ALIAS,
@ -42,9 +49,9 @@ def _deployments(*models: str) -> List[dict]:
]
def _messages(word_count: int) -> List[AllMessageValues]:
def _messages(word_count: int) -> list[AllMessageValues]:
return cast(
List[AllMessageValues],
list[AllMessageValues],
[
{
"role": "user",
@ -84,7 +91,9 @@ def test_write_gate_is_what_prevents_a_pin_below_the_model_minimum():
"""
messages = _messages(word_count=1400)
token_count = token_counter(messages=messages, model="anthropic/claude-opus-4-5", use_default_image_token_count=True)
token_count = token_counter(
messages=messages, model="anthropic/claude-opus-4-5", use_default_image_token_count=True
)
assert 1024 < token_count < 4096
assert is_prompt_caching_valid_prompt(model="anthropic/claude-opus-4-5", messages=messages) is False
@ -110,7 +119,9 @@ async def test_async_filter_deployments_does_not_narrow_prompt_below_model_minim
deployments = _deployments("anthropic/claude-opus-4-6", "anthropic/claude-opus-4-6")
messages = _messages(word_count=1400)
token_count = token_counter(messages=messages, model="anthropic/claude-opus-4-6", use_default_image_token_count=True)
token_count = token_counter(
messages=messages, model="anthropic/claude-opus-4-6", use_default_image_token_count=True
)
assert DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT < token_count < OPUS_4_6_MIN_TOKENS
await PromptCachingCache(cache=cache).async_add_model_id(model_id="dep-2", messages=messages, tools=None)
@ -136,7 +147,9 @@ async def test_async_filter_deployments_narrows_prompt_above_model_minimum():
deployments = _deployments("anthropic/claude-opus-4-6", "anthropic/claude-opus-4-6")
messages = _messages(word_count=5000)
token_count = token_counter(messages=messages, model="anthropic/claude-opus-4-6", use_default_image_token_count=True)
token_count = token_counter(
messages=messages, model="anthropic/claude-opus-4-6", use_default_image_token_count=True
)
assert token_count > OPUS_4_6_MIN_TOKENS
await PromptCachingCache(cache=cache).async_add_model_id(model_id="dep-2", messages=messages, tools=None)
@ -197,10 +210,10 @@ async def test_async_filter_deployments_narrows_for_group_whose_model_minimum_is
AUTO_CACHING_MODEL = "anthropic/claude-sonnet-4-5"
def _auto_caching_messages() -> List[AllMessageValues]:
def _auto_caching_messages() -> list[AllMessageValues]:
"""A prompt over the model minimum that carries no client cache_control."""
return cast(
List[AllMessageValues],
list[AllMessageValues],
[
{"role": "system", "content": "word " * 3000},
{"role": "user", "content": "hello"},
@ -208,7 +221,7 @@ def _auto_caching_messages() -> List[AllMessageValues]:
)
def _affinity_messages(messages: List[AllMessageValues]) -> List[AllMessageValues]:
def _affinity_messages(messages: list[AllMessageValues]) -> list[AllMessageValues]:
"""The messages the check keys deployment affinity on, for a group of `AUTO_CACHING_MODEL`."""
return AnthropicCacheControlHook.messages_with_default_injections(
messages=messages,
@ -218,7 +231,7 @@ def _affinity_messages(messages: List[AllMessageValues]) -> List[AllMessageValue
class _SentMessagesCapture(CustomLogger):
def __init__(self):
self.messages: List[AllMessageValues] | None = None
self.messages: list[AllMessageValues] | None = None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
standard_logging_object = kwargs.get("standard_logging_object")
@ -338,7 +351,7 @@ async def test_claude_code_one_shot_subagent_does_not_reuse_an_auto_injected_aff
cache = DualCache()
check = PromptCachingDeploymentCheck(cache=cache)
deployments = _deployments(AUTO_CACHING_MODEL, AUTO_CACHING_MODEL)
messages = cast(List[AllMessageValues], [{"role": "user", "content": "unique " * 3000}])
messages = cast(list[AllMessageValues], [{"role": "user", "content": "unique " * 3000}])
request_kwargs = {
"system": [
{
@ -441,7 +454,7 @@ def test_client_supplied_cache_control_keeps_its_own_prefix_boundary(monkeypatch
"""
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
messages = cast(
List[AllMessageValues],
list[AllMessageValues],
[
{
"role": "system",
@ -491,7 +504,7 @@ async def test_async_filter_deployments_counts_the_prompt_off_the_event_loop():
warm_tokenizer("anthropic/claude-fable-5")
check = PromptCachingDeploymentCheck(cache=DualCache())
deployments = _deployments("anthropic/claude-fable-5")
messages = cast(List[AllMessageValues], [{"role": "user", "content": text * 100}])
messages = cast(list[AllMessageValues], [{"role": "user", "content": text * 100}])
result, took, lags = await timed_with_loop_lags(
lambda: check.async_filter_deployments(
@ -516,7 +529,7 @@ async def test_async_log_success_event_counts_the_prompt_off_the_event_loop():
cache = DualCache()
check = PromptCachingDeploymentCheck(cache=cache)
messages = cast(
List[AllMessageValues],
list[AllMessageValues],
[{"role": "user", "content": [{"type": "text", "text": text * 100, "cache_control": {"type": "ephemeral"}}]}],
)
standard_logging_object = {

View file

@ -1,21 +1,10 @@
import asyncio
from typing import Optional
import json
from unittest.mock import AsyncMock, patch
import pytest
import json
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.types.llms.openai import (
IncompleteDetails,
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
)
from litellm.types.utils import StandardLoggingPayload
@pytest.mark.asyncio
@ -119,14 +108,11 @@ async def test_async_responses_api_routing_with_previous_response_id():
input="Hello, how are you?",
truncation="auto",
)
print("RESPONSE", response)
# Store the model_id from the response
expected_model_id = response._hidden_params["model_id"]
response_id = response.id
print("Response ID=", response_id, "came from model_id=", expected_model_id)
# Make 10 other requests with previous_response_id, assert that they are sent to the same model_id
for i in range(10):
# Reset the mock for the next call
@ -137,7 +123,7 @@ async def test_async_responses_api_routing_with_previous_response_id():
response = await router.aresponses(
model=MODEL,
input=f"Follow-up question {i+1}",
input=f"Follow-up question {i + 1}",
truncation="auto",
previous_response_id=response_id,
)
@ -163,9 +149,7 @@ async def test_async_routing_without_previous_response_id():
"id": "msg_123",
"status": "completed",
"role": "assistant",
"content": [
{"type": "output_text", "text": "Hello there!", "annotations": []}
],
"content": [{"type": "output_text", "text": "Hello there!", "annotations": []}],
}
],
"parallel_tool_calls": True,
@ -266,9 +250,7 @@ async def test_async_routing_without_previous_response_id():
used_model_ids.add(response._hidden_params["model_id"])
# We should have used more than one model_id if load balancing is working
assert (
len(used_model_ids) > 1
), "Load balancing isn't working, only one deployment was used"
assert len(used_model_ids) > 1, "Load balancing isn't working, only one deployment was used"
@pytest.mark.asyncio

View file

@ -1,12 +1,10 @@
import asyncio
import json
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import json
import litellm
from litellm.caching.affinity_cache import claim_affinity_pin
from litellm.caching.dual_cache import DualCache
@ -46,9 +44,7 @@ async def test_async_session_id_affinity_routes_to_same_deployment():
"id": "msg_123",
"status": "completed",
"role": "assistant",
"content": [
{"type": "output_text", "text": "Hello there!", "annotations": []}
],
"content": [{"type": "output_text", "text": "Hello there!", "annotations": []}],
}
],
"parallel_tool_calls": True,
@ -164,9 +160,7 @@ async def test_async_session_id_affinity_priority_over_user_key():
)
await callback.cache.async_set_cache(
DeploymentAffinityCheck.get_session_affinity_cache_key(
"model_group", "session1", user_key="user1"
),
DeploymentAffinityCheck.get_session_affinity_cache_key("model_group", "session1", user_key="user1"),
{"model_id": "deployment-2"},
)
@ -175,9 +169,7 @@ async def test_async_session_id_affinity_priority_over_user_key():
model="model_group",
healthy_deployments=healthy_deployments,
messages=[],
request_kwargs={
"metadata": {"user_api_key_hash": "user1", "session_id": "session1"}
},
request_kwargs={"metadata": {"user_api_key_hash": "user1", "session_id": "session1"}},
)
assert len(filtered) == 1
@ -575,16 +567,17 @@ async def test_claim_pin_falls_back_to_pod_local_when_redis_is_down():
(None, {"model": "second"}),
],
)
async def test_eligible_affinity_claim_replaces_stale_pins_and_slides_ttl(
stored: object, expected: object
) -> None:
async def test_eligible_affinity_claim_replaces_stale_pins_and_slides_ttl(stored: object, expected: object) -> None:
clock: Final = MagicMock(return_value=100.0)
cache: Final = DualCache(in_memory_cache=InMemoryCache(clock=clock))
cache.in_memory_cache.set_cache("tier-pin", stored, ttl=10)
clock.return_value = 105.0
winner: Final = await claim_affinity_pin(
cache, "tier-pin", {"model": "second"}, 30,
cache,
"tier-pin",
{"model": "second"},
30,
eligible_values=({"model": "first"}, {"model": "second"}),
)
@ -600,13 +593,18 @@ async def test_eligible_affinity_claim_replaces_stale_pins_and_slides_ttl(
async def test_concurrent_eligible_claims_return_one_winner() -> None:
cache: Final = DualCache()
candidates: Final = ({"model": "first"}, {"model": "second"})
winners: Final = await asyncio.gather(*(
claim_affinity_pin(
cache, "tier-pin", candidates[index % 2], 30,
eligible_values=candidates,
winners: Final = await asyncio.gather(
*(
claim_affinity_pin(
cache,
"tier-pin",
candidates[index % 2],
30,
eligible_values=candidates,
)
for index in range(20)
)
for index in range(20)
))
)
assert winners == [{"model": "first"}] * 20
assert cache.in_memory_cache.get_cache("tier-pin") == {"model": "first"}
@ -628,23 +626,19 @@ async def test_legacy_deployment_claim_retains_decoder_and_keepalive(
clock: Final = MagicMock(return_value=100.0)
cache: Final = DualCache(in_memory_cache=InMemoryCache(clock=clock))
callback: Final = DeploymentAffinityCheck(
cache=cache, ttl_seconds=30,
enable_user_key_affinity=False, enable_responses_api_affinity=False,
cache=cache,
ttl_seconds=30,
enable_user_key_affinity=False,
enable_responses_api_affinity=False,
)
cache.in_memory_cache.set_cache("deployment-pin", stored, ttl=10)
clock.return_value = 105.0
winner: Final = await callback._claim_pin(
"deployment-pin", {"model_id": "7"}, 30
)
winner: Final = await callback._claim_pin("deployment-pin", {"model_id": "7"}, 30)
assert winner == expected
assert cache.in_memory_cache.ttl_dict["deployment-pin"] == (
135.0 if refresh else 110.0
)
assert cache.in_memory_cache.get_cache("deployment-pin") == (
{"model_id": "7"} if refresh else stored
)
assert cache.in_memory_cache.ttl_dict["deployment-pin"] == (135.0 if refresh else 110.0)
assert cache.in_memory_cache.get_cache("deployment-pin") == ({"model_id": "7"} if refresh else stored)
@pytest.mark.asyncio
@ -668,13 +662,13 @@ async def test_redis_deployment_claim_preserves_legacy_result_decoding(
redis.async_register_script.return_value = AsyncMock(return_value=raw)
cache: Final = DualCache(redis_cache=redis)
callback: Final = DeploymentAffinityCheck(
cache=cache, ttl_seconds=30,
enable_user_key_affinity=False, enable_responses_api_affinity=False,
cache=cache,
ttl_seconds=30,
enable_user_key_affinity=False,
enable_responses_api_affinity=False,
)
winner: Final = await callback._claim_pin(
"deployment-pin", {"model_id": "candidate"}, 30
)
winner: Final = await callback._claim_pin("deployment-pin", {"model_id": "candidate"}, 30)
assert winner == expected
assert cache.in_memory_cache.get_cache("deployment-pin") == stored

View file