mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
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:
parent
9e39c751ed
commit
78a751c049
27 changed files with 290 additions and 403 deletions
0
tests/unit/messages/__init__.py
Normal file
0
tests/unit/messages/__init__.py
Normal 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
|
||||
|
||||
|
|
@ -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",
|
||||
0
tests/unit/ocr/__init__.py
Normal file
0
tests/unit/ocr/__init__.py
Normal 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"
|
||||
|
||||
|
||||
0
tests/unit/passthrough/__init__.py
Normal file
0
tests/unit/passthrough/__init__.py
Normal 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
|
||||
|
||||
|
||||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
0
tests/unit/rag/ingestion/__init__.py
Normal file
0
tests/unit/rag/ingestion/__init__.py
Normal 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)
|
||||
|
||||
|
||||
|
|
@ -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] = {}
|
||||
|
||||
|
|
@ -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:
|
||||
|
|
@ -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,
|
||||
):
|
||||
|
|
@ -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=[
|
||||
|
|
@ -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 = {
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
0
tests/unit/rust_bridge/__init__.py
Normal file
0
tests/unit/rust_bridge/__init__.py
Normal file
0
tests/unit/rust_bridge/chat_completions/__init__.py
Normal file
0
tests/unit/rust_bridge/chat_completions/__init__.py
Normal file
0
tests/unit/rust_bridge/messages/__init__.py
Normal file
0
tests/unit/rust_bridge/messages/__init__.py
Normal file
Loading…
Add table
Reference in a new issue