mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
test: make 77 legacy live tests offline in litellm_utils, router_unit and responses dirs (#45298)
* test: make 77 legacy live tests offline in litellm_utils, router_unit and responses dirs * test: restore lost coverage in litellm_utils offline replacements Bound live health-check tasks to the concurrency limit so eager task creation behind a semaphore fails, route the Langfuse trace-id check through completion -> Logging.get_trace_id with all four metadata combinations, run the callback dedup checks through acompletion, pass the dynamic key via litellm_params, cover the env-inferred default model list, and add the per-provider audio transcription config lookup * test: restore router coverage lost in the offline move Add the missing non-stream LIT-3058 header/count-once test, pin the UTC minute on every usage-counter test, assert router-level client reuse for transcription, and replace the private selector-attribute checks with routing outcomes per strategy. Assert outbound bodies for speech, rerank, image, assistants and moderation, add timeouts to event waits, and drop the router_unit_tests husk files that no longer hold tests * test: restore sync stream, router sync stream, error event and field-type coverage for migrated responses tests Add the sync streaming logging and sync Router.responses streaming cases the offline replacements dropped, port the legacy per-event and response field-type validation, assert raw headers, the float created_at conversion, MCP tool headers, the search_context_size input, and add an offline replacement for the in-stream context-window error event. Remove helpers left dead in the legacy file. * test: run numpydoc-backed unit tests in GHA and drop a test that pinned a crash * test: restore sequence number, item id and content part checks in the responses stream validator * test: match router logging events to the test's own deployment so late events from other tests are ignored * test: restore the s3 cold storage history test to its original assertions --------- Co-authored-by: yuneng <yuneng@berri.ai>
This commit is contained in:
parent
7c7b0ea85b
commit
04d97abffb
26 changed files with 2752 additions and 3051 deletions
4
.github/workflows/_test-unit-base.yml
vendored
4
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -143,9 +143,9 @@ jobs:
|
|||
run: |
|
||||
diff -u model_prices_and_context_window.json litellm/model_prices_and_context_window_backup.json
|
||||
if [ -z "$RUST_BRIDGE_ARTIFACT" ]; then
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --extra utils
|
||||
else
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --no-install-project
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --extra utils --no-install-project
|
||||
uv pip install --no-deps --python .venv/bin/python rust-bridge-dist/*.whl
|
||||
cp rust-bridge-dist/litellm/rust_bridge/_native.abi3.so litellm/rust_bridge/_native.abi3.so
|
||||
uv run --no-sync python -c "import importlib.metadata; import litellm.rust_bridge._native; print(importlib.metadata.version('litellm'))"
|
||||
|
|
|
|||
|
|
@ -6,17 +6,20 @@ import uuid
|
|||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from collections.abc import Set as AbstractSet
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
import yaml
|
||||
from _pytest.mark.structures import ParameterSet
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from integration._support.client import Gateway, Scenario, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from openai.types.responses import ResponseCompletedEvent
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
|
@ -463,6 +466,26 @@ def _spend_row(*request_ids: str) -> Mapping[str, JsonValue]:
|
|||
return rows[0]
|
||||
|
||||
|
||||
def _spend_proxy_server_request(request_id: str) -> dict[str, JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT proxy_server_request FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
|
||||
(request_id,),
|
||||
),
|
||||
lambda found: len(found) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
return _JSON_OBJECT.validate_python(rows[0]["proxy_server_request"])
|
||||
|
||||
|
||||
def _config_storing_prompts(directory: Path) -> Path:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["general_settings"]["store_prompts_in_spend_logs"] = True
|
||||
path: Final = directory / "store-prompts.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
def _assert_caller_id(endpoint: Endpoint, stream: bool, caller_id: str, response_id: str) -> None:
|
||||
if endpoint == "messages" and stream:
|
||||
assert caller_id.startswith("msg_"), caller_id
|
||||
|
|
@ -654,6 +677,74 @@ def test_gemini_own_tool_call_signature_is_still_replayed_on_the_function_call_p
|
|||
assert _spend_row(rounds[1])["status"] == "success"
|
||||
|
||||
|
||||
def test_responses_previous_response_id_replays_gemini_session_history(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
first_prompt: Final = f"gemini-session-first-{uuid.uuid4().hex}"
|
||||
second_prompt: Final = f"gemini-session-second-{uuid.uuid4().hex}"
|
||||
first_answer: Final = f"gemini-answer-first-{uuid.uuid4().hex}"
|
||||
second_answer: Final = f"gemini-answer-second-{uuid.uuid4().hex}"
|
||||
first_provider_id: Final = f"gemini-response-first-{uuid.uuid4().hex}"
|
||||
second_provider_id: Final = f"gemini-response-second-{uuid.uuid4().hex}"
|
||||
calls: Final = itertools.count()
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
if next(calls) == 0:
|
||||
return _reply(first_provider_id, stream=False, parts=({"text": first_answer},))
|
||||
return _reply(second_provider_id, stream=False, parts=({"text": second_answer},))
|
||||
|
||||
with (
|
||||
wire_server(respond) as wire,
|
||||
owned_proxy(gateway, tmp_path, {}, config=_config_storing_prompts(tmp_path), workers=1) as prompt_gateway,
|
||||
prompt_gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _register(prompt_gateway, scenario, "gemini", wire)
|
||||
first_response: Final = prompt_gateway.request(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
{"model": model, "input": first_prompt, **_CACHE_BUST},
|
||||
)
|
||||
assert first_response.status_code == 200, first_response.text
|
||||
first_body: Final = _JSON_OBJECT.validate_json(first_response.content)
|
||||
first_response_id: Final = str(first_body["id"])
|
||||
assert _output_text(first_body) == first_answer
|
||||
first_spend_id: Final = _upstream_response_id("responses", first_response_id)
|
||||
assert _spend_row(first_spend_id)["status"] == "success"
|
||||
first_proxy_request: Final = _spend_proxy_server_request(first_spend_id)
|
||||
assert first_proxy_request["input"] == first_prompt, first_proxy_request
|
||||
|
||||
second_response: Final = prompt_gateway.request(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
{
|
||||
"model": model,
|
||||
"input": second_prompt,
|
||||
"previous_response_id": first_response_id,
|
||||
**_CACHE_BUST,
|
||||
},
|
||||
)
|
||||
assert second_response.status_code == 200, second_response.text
|
||||
second_body: Final = _JSON_OBJECT.validate_json(second_response.content)
|
||||
second_response_id: Final = str(second_body["id"])
|
||||
assert _output_text(second_body) == second_answer
|
||||
assert _spend_row(_upstream_response_id("responses", second_response_id))["status"] == "success"
|
||||
|
||||
requests: Final = wire.drain()
|
||||
assert [(request.method, request.target) for request in requests] == [
|
||||
("POST", _target("gemini", False)),
|
||||
("POST", _target("gemini", False)),
|
||||
]
|
||||
first_provider_body: Final = _JSON_OBJECT.validate_json(requests[0].body)
|
||||
second_provider_body: Final = _JSON_OBJECT.validate_json(requests[1].body)
|
||||
assert first_provider_body == _expected_body("responses", _user(first_prompt))
|
||||
assert second_provider_body == _expected_body(
|
||||
"responses",
|
||||
_user(first_prompt),
|
||||
_model_turn({"text": first_answer}),
|
||||
_user(second_prompt),
|
||||
)
|
||||
|
||||
|
||||
_FIVE_KB: Final = "s" * 5000
|
||||
_ASSISTANT_SHAPES: Final = (
|
||||
pytest.param({"reasoning_content": _REASONING, "thinking_blocks": 5}, _REPLAYED_TURN, id="blocks-int"),
|
||||
|
|
|
|||
|
|
@ -17,7 +17,6 @@ import pytest
|
|||
|
||||
|
||||
from litellm.llms.base_llm.base_utils import BaseTokenCounter
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
|
||||
class BaseTokenCounterTest(ABC):
|
||||
|
|
@ -71,69 +70,3 @@ class BaseTokenCounterTest(ABC):
|
|||
):
|
||||
pytest.skip(f"Missing or invalid credentials: {e}")
|
||||
raise
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_basic(self):
|
||||
"""
|
||||
Test basic token counting functionality.
|
||||
|
||||
Verifies that:
|
||||
- Token counter returns a TokenCountResponse
|
||||
- total_tokens is greater than 0
|
||||
- tokenizer_type is set
|
||||
- No error occurred
|
||||
"""
|
||||
token_counter = self.get_token_counter()
|
||||
model = self.get_test_model()
|
||||
messages = self.get_test_messages()
|
||||
deployment = self.get_deployment_config()
|
||||
|
||||
result = await token_counter.count_tokens(
|
||||
model_to_use=model,
|
||||
messages=messages,
|
||||
contents=None,
|
||||
deployment=deployment,
|
||||
request_model=model,
|
||||
)
|
||||
|
||||
print(f"Token count result: {result}")
|
||||
|
||||
assert result is not None, "Token counter should return a result"
|
||||
assert isinstance(
|
||||
result, TokenCountResponse
|
||||
), "Result should be TokenCountResponse"
|
||||
assert (
|
||||
result.total_tokens > 0
|
||||
), f"Token count should be > 0, got {result.total_tokens}"
|
||||
assert result.tokenizer_type is not None, "tokenizer_type should be set"
|
||||
assert (
|
||||
result.error is not True
|
||||
), f"Token counting should not error: {result.error_message}"
|
||||
|
||||
def test_should_use_token_counting_api(self):
|
||||
"""
|
||||
Test that should_use_token_counting_api returns True for the correct provider.
|
||||
|
||||
Verifies that the token counter correctly identifies when it should be used
|
||||
based on the custom_llm_provider.
|
||||
"""
|
||||
token_counter = self.get_token_counter()
|
||||
provider = self.get_custom_llm_provider()
|
||||
|
||||
result = token_counter.should_use_token_counting_api(
|
||||
custom_llm_provider=provider
|
||||
)
|
||||
|
||||
assert (
|
||||
result is True
|
||||
), f"should_use_token_counting_api should return True for {provider}"
|
||||
|
||||
# Also verify it returns False for other providers
|
||||
other_provider = "some_other_provider_that_doesnt_exist"
|
||||
result_other = token_counter.should_use_token_counting_api(
|
||||
custom_llm_provider=other_provider
|
||||
)
|
||||
|
||||
assert (
|
||||
result_other is False
|
||||
), f"should_use_token_counting_api should return False for {other_provider}"
|
||||
|
|
|
|||
|
|
@ -1,326 +0,0 @@
|
|||
#### What this tests ####
|
||||
# This tests if ahealth_check() actually works
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_health_check():
|
||||
response = await litellm.ahealth_check(
|
||||
model_params={
|
||||
"model": "azure/gpt-4.1-mini",
|
||||
"messages": [{"role": "user", "content": "Hey, how's it going?"}],
|
||||
"api_key": os.getenv("AZURE_AI_API_KEY"),
|
||||
"api_base": os.getenv("AZURE_AI_API_BASE"),
|
||||
"api_version": os.getenv("AZURE_AI_API_VERSION"),
|
||||
}
|
||||
)
|
||||
print(f"response: {response}")
|
||||
|
||||
assert "x-ratelimit-remaining-tokens" in response
|
||||
return response
|
||||
|
||||
|
||||
# asyncio.run(test_azure_health_check())
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_embedding_health_check():
|
||||
response = await litellm.ahealth_check(
|
||||
model_params={
|
||||
"model": "azure/text-embedding-ada-002",
|
||||
"api_key": os.getenv("AZURE_AI_API_KEY"),
|
||||
"api_base": os.getenv("AZURE_AI_API_BASE"),
|
||||
"api_version": os.getenv("AZURE_AI_API_VERSION"),
|
||||
},
|
||||
input=["test for litellm"],
|
||||
mode="embedding",
|
||||
)
|
||||
print(f"response: {response}")
|
||||
|
||||
assert "x-ratelimit-remaining-tokens" in response
|
||||
return response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_img_gen_health_check():
|
||||
response = await litellm.ahealth_check(
|
||||
model_params={
|
||||
"model": "gpt-image-1",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
mode="image_generation",
|
||||
prompt="cute baby sea otter",
|
||||
)
|
||||
print(f"response: {response}")
|
||||
|
||||
assert isinstance(response, dict) and "error" not in response
|
||||
return response
|
||||
|
||||
|
||||
# asyncio.run(test_openai_img_gen_health_check())
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# asyncio.run(test_sagemaker_embedding_health_check())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_groq_health_check():
|
||||
"""
|
||||
This should not fail
|
||||
|
||||
ensure that provider wildcard model passes health check
|
||||
"""
|
||||
litellm.set_verbose = True
|
||||
response = await litellm.ahealth_check(
|
||||
model_params={
|
||||
"api_key": os.environ.get("GROQ_API_KEY"),
|
||||
"model": "groq/*",
|
||||
"messages": [{"role": "user", "content": "What's 1 + 1?"}],
|
||||
},
|
||||
mode=None,
|
||||
prompt="What's 1 + 1?",
|
||||
input=["test from litellm"],
|
||||
)
|
||||
print(f"response: {response}")
|
||||
assert response == {}
|
||||
|
||||
return response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cohere_rerank_health_check():
|
||||
response = await litellm.ahealth_check(
|
||||
model_params={
|
||||
"model": "cohere/rerank-english-v3.0",
|
||||
"api_key": os.getenv("COHERE_API_KEY"),
|
||||
},
|
||||
mode="rerank",
|
||||
prompt="Hey, how's it going",
|
||||
)
|
||||
|
||||
assert "error" not in response
|
||||
|
||||
print(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_speech_health_check():
|
||||
response = await litellm.ahealth_check(
|
||||
model_params={
|
||||
"model": "openai/tts-1",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
mode="audio_speech",
|
||||
prompt="Hey",
|
||||
)
|
||||
|
||||
assert "error" not in response
|
||||
|
||||
print(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_speech_health_check_with_another_voice():
|
||||
response = await litellm.ahealth_check(
|
||||
model_params={
|
||||
"model": "openai/tts-1",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
"health_check_voice": "en-US-JennyNeural",
|
||||
},
|
||||
mode="audio_speech",
|
||||
prompt="Hey",
|
||||
)
|
||||
|
||||
assert "error" not in response
|
||||
|
||||
print(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_transcription_health_check():
|
||||
litellm.set_verbose = True
|
||||
response = await litellm.ahealth_check(
|
||||
model_params={
|
||||
"model": "openai/whisper-1",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
mode="audio_transcription",
|
||||
)
|
||||
|
||||
print(f"response: {response}")
|
||||
|
||||
assert "error" not in response
|
||||
|
||||
print(response)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_bad_model():
|
||||
import time
|
||||
|
||||
from litellm.proxy.health_check import _perform_health_check
|
||||
|
||||
model_list = [
|
||||
{
|
||||
"model_name": "openai-gpt-4o",
|
||||
"litellm_params": {
|
||||
"api_key": "sk-9876",
|
||||
"api_base": "https://exampleopenaiendpoint-production.up.railway.app",
|
||||
"model": "openai/my-fake-openai-endpoint",
|
||||
"mock_timeout": True,
|
||||
"timeout": 60,
|
||||
},
|
||||
"model_info": {
|
||||
"id": "ca27ca2eeea2f9e38bb274ead831948a26621a3738d06f1797253f0e6c4278c0",
|
||||
"db_model": False,
|
||||
"health_check_timeout": 1,
|
||||
},
|
||||
},
|
||||
]
|
||||
details = None
|
||||
healthy_endpoints, unhealthy_endpoints, _ = await _perform_health_check(
|
||||
model_list, details
|
||||
)
|
||||
print(f"healthy_endpoints: {healthy_endpoints}")
|
||||
print(f"unhealthy_endpoints: {unhealthy_endpoints}")
|
||||
|
||||
# Track which model is actually used in the health check
|
||||
health_check_calls = []
|
||||
|
||||
async def mock_health_check(litellm_params, **kwargs):
|
||||
health_check_calls.append(litellm_params["model"])
|
||||
await asyncio.sleep(10)
|
||||
return {"status": "healthy"}
|
||||
|
||||
with patch(
|
||||
"litellm.ahealth_check", side_effect=mock_health_check
|
||||
) as mock_health_check:
|
||||
start_time = time.time()
|
||||
healthy_endpoints, unhealthy_endpoints, _ = await _perform_health_check(
|
||||
model_list
|
||||
)
|
||||
end_time = time.time()
|
||||
print("health check calls: ", health_check_calls)
|
||||
assert len(healthy_endpoints) == 0
|
||||
assert len(unhealthy_endpoints) == 1
|
||||
assert (
|
||||
end_time - start_time < 2
|
||||
), "Health check took longer than health_check_timeout"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_respects_concurrency_limit():
|
||||
from litellm.proxy.health_check import _perform_health_check
|
||||
|
||||
model_list = [
|
||||
{"litellm_params": {"model": f"openai/gpt-4o-mini-{i}", "api_key": "fake-key"}}
|
||||
for i in range(6)
|
||||
]
|
||||
|
||||
active = 0
|
||||
max_active = 0
|
||||
|
||||
async def mock_health_check(litellm_params, **kwargs):
|
||||
nonlocal active, max_active
|
||||
active += 1
|
||||
max_active = max(max_active, active)
|
||||
await asyncio.sleep(0.05)
|
||||
active -= 1
|
||||
return {"status": "healthy"}
|
||||
|
||||
with patch("litellm.ahealth_check", side_effect=mock_health_check):
|
||||
await _perform_health_check(model_list, max_concurrency=2)
|
||||
|
||||
assert max_active <= 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_creates_only_bounded_initial_tasks():
|
||||
from litellm.proxy.health_check import _perform_health_check
|
||||
|
||||
model_list = [
|
||||
{"litellm_params": {"model": f"openai/gpt-4o-mini-{i}", "api_key": "fake-key"}}
|
||||
for i in range(10)
|
||||
]
|
||||
release_event = asyncio.Event()
|
||||
create_task_call_count = 0
|
||||
real_create_task = asyncio.create_task
|
||||
|
||||
async def mock_health_check(litellm_params, **kwargs):
|
||||
await release_event.wait()
|
||||
return {"status": "healthy"}
|
||||
|
||||
def tracked_create_task(coro):
|
||||
nonlocal create_task_call_count
|
||||
create_task_call_count += 1
|
||||
return real_create_task(coro)
|
||||
|
||||
with (
|
||||
patch("litellm.ahealth_check", side_effect=mock_health_check),
|
||||
patch(
|
||||
"litellm.proxy.health_check.asyncio.create_task",
|
||||
side_effect=tracked_create_task,
|
||||
),
|
||||
):
|
||||
perform_task = real_create_task(
|
||||
_perform_health_check(model_list, max_concurrency=2)
|
||||
)
|
||||
await asyncio.sleep(0.05)
|
||||
assert create_task_call_count == 2
|
||||
release_event.set()
|
||||
await perform_task
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_timeout_does_not_cancel_other_health_checks():
|
||||
from litellm.proxy.health_check import _perform_health_check
|
||||
|
||||
model_list = [
|
||||
{
|
||||
"litellm_params": {"model": "openai/slow-model", "api_key": "fake-key"},
|
||||
"model_info": {"health_check_timeout": 0.05},
|
||||
},
|
||||
{
|
||||
"litellm_params": {"model": "openai/fast-model", "api_key": "fake-key"},
|
||||
"model_info": {"health_check_timeout": 1},
|
||||
},
|
||||
]
|
||||
|
||||
async def mock_health_check(litellm_params, **kwargs):
|
||||
if litellm_params["model"] == "openai/slow-model":
|
||||
await asyncio.sleep(0.2)
|
||||
return {"status": "healthy"}
|
||||
await asyncio.sleep(0.01)
|
||||
return {"status": "healthy"}
|
||||
|
||||
with patch("litellm.ahealth_check", side_effect=mock_health_check):
|
||||
healthy_endpoints, unhealthy_endpoints, _ = await _perform_health_check(
|
||||
model_list, max_concurrency=1
|
||||
)
|
||||
|
||||
healthy_models = {endpoint["model"] for endpoint in healthy_endpoints}
|
||||
unhealthy_models = {endpoint["model"] for endpoint in unhealthy_endpoints}
|
||||
|
||||
assert "openai/fast-model" in healthy_models
|
||||
assert "openai/slow-model" in unhealthy_models
|
||||
|
|
@ -1,4 +1,3 @@
|
|||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
|
|
@ -7,7 +6,6 @@ from dotenv import load_dotenv
|
|||
|
||||
load_dotenv()
|
||||
import tempfile
|
||||
from unittest.mock import MagicMock, patch
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
|
@ -93,8 +91,6 @@ def test_oidc_google():
|
|||
)
|
||||
|
||||
print(f"secret_val: {redact_oidc_signature(secret_val)}")
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
os.environ.get("ACTIONS_ID_TOKEN_REQUEST_TOKEN") is None,
|
||||
reason="Cannot run without being in GitHub Actions",
|
||||
|
|
@ -105,112 +101,3 @@ def test_oidc_github():
|
|||
)
|
||||
|
||||
print(f"secret_val: {redact_oidc_signature(secret_val)}")
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
os.environ.get("CIRCLE_OIDC_TOKEN") is None,
|
||||
reason="Cannot run without being in CircleCI Runner",
|
||||
)
|
||||
def test_oidc_circleci():
|
||||
secret_val = get_secret("oidc/circleci/")
|
||||
|
||||
print(f"secret_val: {redact_oidc_signature(secret_val)}")
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
os.environ.get("CIRCLE_OIDC_TOKEN_V2") is None,
|
||||
reason="Cannot run without being in CircleCI Runner",
|
||||
)
|
||||
def test_oidc_circleci_v2():
|
||||
secret_val = get_secret(
|
||||
"oidc/circleci_v2/https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.titan-text-express-v1/invoke"
|
||||
)
|
||||
|
||||
print(f"secret_val: {redact_oidc_signature(secret_val)}")
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_google_secret_manager():
|
||||
"""
|
||||
Test that we can get a secret from Google Secret Manager
|
||||
"""
|
||||
os.environ["GOOGLE_SECRET_MANAGER_PROJECT_ID"] = "litellm-ci-cd"
|
||||
|
||||
from litellm.secret_managers.google_secret_manager import GoogleSecretManager
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"payload": {
|
||||
"data": base64.b64encode(b"anything").decode("utf-8"),
|
||||
}
|
||||
}
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch.object(
|
||||
GoogleSecretManager,
|
||||
"sync_construct_request_headers",
|
||||
return_value={"Authorization": "Bearer mock_token"},
|
||||
),
|
||||
):
|
||||
secret_manager = GoogleSecretManager()
|
||||
secret_manager.sync_httpx_client = MagicMock()
|
||||
secret_manager.sync_httpx_client.get.return_value = mock_response
|
||||
|
||||
secret_val = secret_manager.get_secret_from_google_secret_manager(
|
||||
secret_name="OPENAI_API_KEY"
|
||||
)
|
||||
print("secret_val: {}".format(secret_val))
|
||||
|
||||
assert (
|
||||
secret_val == "anything"
|
||||
), "did not get expected secret value. expect 'anything', got '{}'".format(
|
||||
secret_val
|
||||
)
|
||||
|
||||
secret_manager.sync_httpx_client.get.assert_called_once()
|
||||
call_url = secret_manager.sync_httpx_client.get.call_args[1]["url"]
|
||||
assert "projects/litellm-ci-cd/secrets/OPENAI_API_KEY" in call_url
|
||||
|
||||
|
||||
def test_google_secret_manager_read_in_memory():
|
||||
"""
|
||||
Test that Google Secret manager returns in memory value when it exists
|
||||
"""
|
||||
from litellm.secret_managers.google_secret_manager import GoogleSecretManager
|
||||
|
||||
os.environ["GOOGLE_SECRET_MANAGER_PROJECT_ID"] = "litellm-ci-cd"
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch.object(
|
||||
GoogleSecretManager,
|
||||
"sync_construct_request_headers",
|
||||
return_value={"Authorization": "Bearer mock_token"},
|
||||
),
|
||||
):
|
||||
secret_manager = GoogleSecretManager()
|
||||
secret_manager.cache.cache_dict["UNIQUE_KEY"] = None
|
||||
secret_manager.cache.cache_dict["UNIQUE_KEY_2"] = "lite-llm"
|
||||
|
||||
secret_val = secret_manager.get_secret_from_google_secret_manager(
|
||||
secret_name="UNIQUE_KEY"
|
||||
)
|
||||
print("secret_val: {}".format(secret_val))
|
||||
assert secret_val is None
|
||||
|
||||
secret_val = secret_manager.get_secret_from_google_secret_manager(
|
||||
secret_name="UNIQUE_KEY_2"
|
||||
)
|
||||
print("secret_val: {}".format(secret_val))
|
||||
assert secret_val == "lite-llm"
|
||||
|
|
|
|||
|
|
@ -1,7 +1,5 @@
|
|||
import copy
|
||||
import logging
|
||||
import time
|
||||
from datetime import datetime
|
||||
from unittest import mock
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
|
@ -15,22 +13,14 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, headers
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.litellm_core_utils.duration_parser import (
|
||||
get_last_day_of_month,
|
||||
_extract_from_regex,
|
||||
)
|
||||
from litellm.utils import (
|
||||
check_valid_key,
|
||||
get_llm_provider,
|
||||
get_supported_openai_params,
|
||||
get_token_count,
|
||||
get_valid_models,
|
||||
trim_messages,
|
||||
validate_environment,
|
||||
)
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
||||
# Assuming your trim_messages, shorten_message_to_fit_limit, and get_token_count functions are all in a module named 'message_utils'
|
||||
|
|
@ -47,676 +37,29 @@ def reset_mock_cache():
|
|||
# test_basic_trimming()
|
||||
|
||||
|
||||
|
||||
|
||||
# test_basic_trimming_no_max_tokens_specified()
|
||||
|
||||
|
||||
|
||||
|
||||
# test_multiple_messages_trimming()
|
||||
|
||||
|
||||
|
||||
|
||||
# test_multiple_messages_no_trimming()
|
||||
|
||||
|
||||
|
||||
|
||||
# test_large_trimming()
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["anthropic", "xai"])
|
||||
def test_get_valid_models_with_custom_llm_provider(custom_llm_provider):
|
||||
from litellm.utils import ProviderConfigManager
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
provider_config = ProviderConfigManager.get_provider_model_info(
|
||||
model=None,
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
assert provider_config is not None
|
||||
valid_models = get_valid_models(
|
||||
check_provider_endpoint=True, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
print(valid_models)
|
||||
assert len(valid_models) > 0
|
||||
assert set(provider_config.get_models()) == set(valid_models)
|
||||
|
||||
|
||||
# test_get_valid_models()
|
||||
|
||||
|
||||
def test_bad_key():
|
||||
key = "bad-key"
|
||||
response = check_valid_key(model="gpt-5-mini", api_key=key)
|
||||
print(response, key)
|
||||
assert response == False
|
||||
|
||||
|
||||
def test_good_key():
|
||||
key = os.environ["OPENAI_API_KEY"]
|
||||
response = check_valid_key(model="gpt-5-mini", api_key=key)
|
||||
assert response == True
|
||||
|
||||
|
||||
# test validate environment
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_function_to_dict():
|
||||
print("testing function to dict for get current weather")
|
||||
|
||||
def get_current_weather(location: str, unit: str):
|
||||
"""Get the current weather in a given location
|
||||
|
||||
Parameters
|
||||
----------
|
||||
location : str
|
||||
The city and state, e.g. San Francisco, CA
|
||||
unit : {'celsius', 'fahrenheit'}
|
||||
Temperature unit
|
||||
|
||||
Returns
|
||||
-------
|
||||
str
|
||||
a sentence indicating the weather
|
||||
"""
|
||||
if location == "Boston, MA":
|
||||
return "The weather is 12F"
|
||||
|
||||
function_json = litellm.utils.function_to_dict(get_current_weather)
|
||||
print(function_json)
|
||||
|
||||
expected_output = {
|
||||
"name": "get_current_weather",
|
||||
"description": "Get the current weather in a given location",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {
|
||||
"type": "string",
|
||||
"description": "The city and state, e.g. San Francisco, CA",
|
||||
},
|
||||
"unit": {
|
||||
"type": "string",
|
||||
"description": "Temperature unit",
|
||||
"enum": "['fahrenheit', 'celsius']",
|
||||
},
|
||||
},
|
||||
"required": ["location", "unit"],
|
||||
},
|
||||
}
|
||||
print(expected_output)
|
||||
|
||||
assert function_json["name"] == expected_output["name"]
|
||||
assert function_json["description"] == expected_output["description"]
|
||||
assert function_json["parameters"]["type"] == expected_output["parameters"]["type"]
|
||||
assert (
|
||||
function_json["parameters"]["properties"]["location"]
|
||||
== expected_output["parameters"]["properties"]["location"]
|
||||
)
|
||||
|
||||
# the enum can change it can be - which is why we don't assert on unit
|
||||
# {'type': 'string', 'description': 'Temperature unit', 'enum': "['fahrenheit', 'celsius']"}
|
||||
# {'type': 'string', 'description': 'Temperature unit', 'enum': "['celsius', 'fahrenheit']"}
|
||||
|
||||
assert (
|
||||
function_json["parameters"]["required"]
|
||||
== expected_output["parameters"]["required"]
|
||||
)
|
||||
|
||||
print("passed")
|
||||
|
||||
|
||||
# test_function_to_dict()
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_duration_in_seconds():
|
||||
"""
|
||||
Test if duration int is correctly calculated for different str
|
||||
"""
|
||||
import time
|
||||
|
||||
now = time.time()
|
||||
current_time = datetime.fromtimestamp(now)
|
||||
|
||||
if current_time.month == 12:
|
||||
target_year = current_time.year + 1
|
||||
target_month = 1
|
||||
else:
|
||||
target_year = current_time.year
|
||||
target_month = current_time.month + 1
|
||||
|
||||
# Determine the day to set for next month
|
||||
target_day = current_time.day
|
||||
last_day_of_target_month = get_last_day_of_month(target_year, target_month)
|
||||
|
||||
if target_day > last_day_of_target_month:
|
||||
target_day = last_day_of_target_month
|
||||
|
||||
next_month = datetime(
|
||||
year=target_year,
|
||||
month=target_month,
|
||||
day=target_day,
|
||||
hour=current_time.hour,
|
||||
minute=current_time.minute,
|
||||
second=current_time.second,
|
||||
microsecond=current_time.microsecond,
|
||||
)
|
||||
|
||||
# Calculate the duration until the first day of the next month
|
||||
duration_until_next_month = next_month - current_time
|
||||
expected_duration = int(duration_until_next_month.total_seconds())
|
||||
|
||||
value = duration_in_seconds(duration="1mo")
|
||||
|
||||
assert value - expected_duration < 2
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.parametrize("langfuse_trace_id", [None, "my-unique-trace-id"])
|
||||
@pytest.mark.parametrize(
|
||||
"langfuse_existing_trace_id", [None, "my-unique-existing-trace-id"]
|
||||
)
|
||||
def test_logging_trace_id(langfuse_trace_id, langfuse_existing_trace_id):
|
||||
"""
|
||||
- Unit test for `_get_trace_id` function in Logging obj
|
||||
"""
|
||||
from litellm.integrations.langfuse.langfuse_sdk import resolve_trace_id
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
litellm.success_callback = ["langfuse"]
|
||||
litellm_call_id = "my-unique-call-id"
|
||||
litellm_logging_obj = Logging(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="acompletion",
|
||||
litellm_call_id=litellm_call_id,
|
||||
start_time=datetime.now(),
|
||||
function_id="1234",
|
||||
)
|
||||
|
||||
metadata = {}
|
||||
|
||||
if langfuse_trace_id is not None:
|
||||
metadata["trace_id"] = langfuse_trace_id
|
||||
if langfuse_existing_trace_id is not None:
|
||||
metadata["existing_trace_id"] = langfuse_existing_trace_id
|
||||
|
||||
litellm.completion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "Hey how's it going?"}],
|
||||
mock_response="Hey!",
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
time.sleep(3)
|
||||
assert litellm_logging_obj.get_trace_id(service_name="langfuse") is not None
|
||||
|
||||
# langfuse addresses a trace by a 32-hex id, so the id litellm reports back is the
|
||||
# resolved form of whichever source won; that is what the alerting deep link needs
|
||||
if langfuse_existing_trace_id is not None:
|
||||
expected_source = langfuse_existing_trace_id
|
||||
elif langfuse_trace_id is not None:
|
||||
expected_source = langfuse_trace_id
|
||||
else:
|
||||
expected_source = litellm_logging_obj.litellm_trace_id
|
||||
|
||||
assert litellm_logging_obj.get_trace_id(service_name="langfuse") == resolve_trace_id(
|
||||
expected_source
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_is_base64_encoded():
|
||||
import base64
|
||||
|
||||
import requests
|
||||
|
||||
litellm.set_verbose = True
|
||||
url = "https://dummyimage.com/100/100/fff&text=Test+image"
|
||||
response = requests.get(url)
|
||||
file_data = response.content
|
||||
|
||||
encoded_file = base64.b64encode(file_data).decode("utf-8")
|
||||
base64_image = f"data:image/png;base64,{encoded_file}"
|
||||
|
||||
from litellm.utils import is_base64_encoded
|
||||
|
||||
assert is_base64_encoded(s=base64_image) is True
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_is_prompt_caching_enabled_return_default_image_dimensions():
|
||||
"""
|
||||
Assert that `is_prompt_caching_valid_prompt` counts tokens with use_default_image_token_count=True
|
||||
when processing messages containing images
|
||||
|
||||
IMPORTANT: Ensures Get token counter does not make a GET request to the image url
|
||||
"""
|
||||
mock_token_counter = MagicMock(return_value=False)
|
||||
with patch(
|
||||
"litellm.utils.get_messages_reach_token_count",
|
||||
return_value=mock_token_counter,
|
||||
):
|
||||
litellm.utils.is_prompt_caching_valid_prompt(
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "What is in this image?"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "https://www.gstatic.com/webp/gallery/1.webp",
|
||||
"detail": "high",
|
||||
},
|
||||
},
|
||||
],
|
||||
}
|
||||
],
|
||||
tools=None,
|
||||
custom_llm_provider="openai",
|
||||
model="gpt-4o-mini",
|
||||
)
|
||||
|
||||
# Assert token_counter was called with use_default_image_token_count=True
|
||||
args_to_mock_token_counter = mock_token_counter.call_args[1]
|
||||
print("args_to_mock", args_to_mock_token_counter)
|
||||
assert args_to_mock_token_counter["use_default_image_token_count"] is True
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_get_valid_models_fireworks_ai(monkeypatch):
|
||||
from litellm.utils import get_valid_models
|
||||
import litellm
|
||||
|
||||
litellm.turn_on_debug()
|
||||
|
||||
monkeypatch.setenv("FIREWORKS_API_KEY", "sk-9876")
|
||||
monkeypatch.setenv("FIREWORKS_ACCOUNT_ID", "1234")
|
||||
monkeypatch.setattr(litellm, "provider_list", ["fireworks_ai"])
|
||||
|
||||
mock_response_data = {
|
||||
"models": [
|
||||
{
|
||||
"name": "accounts/fireworks/models/llama-3.1-8b-instruct",
|
||||
"displayName": "<string>",
|
||||
"description": "<string>",
|
||||
"createTime": "2023-11-07T05:31:56Z",
|
||||
"createdBy": "<string>",
|
||||
"state": "STATE_UNSPECIFIED",
|
||||
"status": {"code": "OK", "message": "<string>"},
|
||||
"kind": "KIND_UNSPECIFIED",
|
||||
"githubUrl": "<string>",
|
||||
"huggingFaceUrl": "<string>",
|
||||
"baseModelDetails": {
|
||||
"worldSize": 123,
|
||||
"checkpointFormat": "CHECKPOINT_FORMAT_UNSPECIFIED",
|
||||
"parameterCount": "<string>",
|
||||
"moe": True,
|
||||
"tunable": True,
|
||||
},
|
||||
"peftDetails": {
|
||||
"baseModel": "<string>",
|
||||
"r": 123,
|
||||
"targetModules": ["<string>"],
|
||||
},
|
||||
"teftDetails": {},
|
||||
"public": True,
|
||||
"conversationConfig": {
|
||||
"style": "<string>",
|
||||
"system": "<string>",
|
||||
"template": "<string>",
|
||||
},
|
||||
"contextLength": 123,
|
||||
"supportsImageInput": True,
|
||||
"supportsTools": True,
|
||||
"importedFrom": "<string>",
|
||||
"fineTuningJob": "<string>",
|
||||
"defaultDraftModel": "<string>",
|
||||
"defaultDraftTokenCount": 123,
|
||||
"precisions": ["PRECISION_UNSPECIFIED"],
|
||||
"deployedModelRefs": [
|
||||
{
|
||||
"name": "<string>",
|
||||
"deployment": "<string>",
|
||||
"state": "STATE_UNSPECIFIED",
|
||||
"default": True,
|
||||
"public": True,
|
||||
}
|
||||
],
|
||||
"cluster": "<string>",
|
||||
"deprecationDate": {"year": 123, "month": 123, "day": 123},
|
||||
}
|
||||
],
|
||||
"nextPageToken": "<string>",
|
||||
"totalSize": 123,
|
||||
}
|
||||
|
||||
# Create a mock response object
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = mock_response_data
|
||||
|
||||
with patch.object(
|
||||
litellm.module_level_client, "get", return_value=mock_response
|
||||
) as mock_post:
|
||||
valid_models = get_valid_models(check_provider_endpoint=True)
|
||||
print("valid_models", valid_models)
|
||||
mock_post.assert_called_once()
|
||||
assert (
|
||||
"fireworks_ai/accounts/fireworks/models/llama-3.1-8b-instruct"
|
||||
in valid_models
|
||||
)
|
||||
|
||||
|
||||
def test_get_valid_models_default(monkeypatch):
|
||||
"""
|
||||
Ensure that the default models is used when error retrieving from model api.
|
||||
|
||||
Prevent regression for existing usage.
|
||||
"""
|
||||
from litellm.utils import get_valid_models
|
||||
|
||||
monkeypatch.setenv("FIREWORKS_API_KEY", "sk-9876")
|
||||
valid_models = get_valid_models()
|
||||
assert len(valid_models) > 0
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_add_custom_logger_callback_to_specific_event(monkeypatch):
|
||||
from litellm.utils import add_custom_logger_callback_to_specific_event
|
||||
|
||||
monkeypatch.setattr(litellm, "success_callback", [])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
|
||||
add_custom_logger_callback_to_specific_event("langfuse", "success")
|
||||
|
||||
assert len(litellm.success_callback) == 1
|
||||
assert len(litellm.failure_callback) == 0
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_custom_logger_callback_to_specific_event_with_duplicates(
|
||||
monkeypatch,
|
||||
):
|
||||
"""
|
||||
Test that when a callback exists in both success_callback and _async_success_callback,
|
||||
it's not added again
|
||||
"""
|
||||
from litellm.integrations.langfuse.langfuse_prompt_management import (
|
||||
LangfusePromptManagement,
|
||||
)
|
||||
|
||||
# Reset all callback lists
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "success_callback", [])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
|
||||
# Add logger to both success_callback and _async_success_callback
|
||||
langfuse_logger = LangfusePromptManagement()
|
||||
litellm.success_callback.append(langfuse_logger)
|
||||
litellm._async_success_callback.append(langfuse_logger)
|
||||
|
||||
# Get initial lengths
|
||||
initial_success_callback_len = len(litellm.success_callback)
|
||||
initial_async_success_callback_len = len(litellm._async_success_callback)
|
||||
|
||||
# Make a completion call
|
||||
await litellm.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
mock_response="Testing duplicate callbacks",
|
||||
)
|
||||
|
||||
# Assert no new callbacks were added
|
||||
assert len(litellm.success_callback) == initial_success_callback_len
|
||||
assert len(litellm._async_success_callback) == initial_async_success_callback_len
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_custom_logger_callback_to_specific_event_with_duplicates_success_callback(
|
||||
monkeypatch,
|
||||
):
|
||||
"""
|
||||
Test that when a callback exists in both success_callback and _async_success_callback,
|
||||
it's not added again
|
||||
"""
|
||||
from litellm.integrations.langfuse.langfuse_prompt_management import (
|
||||
LangfusePromptManagement,
|
||||
)
|
||||
|
||||
# Reset all callback lists
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "success_callback", [])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
|
||||
# Add logger to both success_callback and _async_success_callback
|
||||
langfuse_logger = LangfusePromptManagement()
|
||||
litellm.success_callback.append(langfuse_logger)
|
||||
|
||||
# Get initial lengths
|
||||
initial_success_callback_len = len(litellm.success_callback)
|
||||
initial_async_success_callback_len = len(litellm._async_success_callback)
|
||||
|
||||
# Make a completion call
|
||||
await litellm.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
mock_response="Testing duplicate callbacks",
|
||||
)
|
||||
|
||||
# Assert no new callbacks were added
|
||||
assert len(litellm.success_callback) == initial_success_callback_len
|
||||
assert len(litellm._async_success_callback) == initial_async_success_callback_len
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_custom_logger_callback_to_specific_event_with_duplicates_callbacks(
|
||||
monkeypatch,
|
||||
):
|
||||
"""
|
||||
Test that when a callback exists in both success_callback and _async_success_callback,
|
||||
it's not added again
|
||||
"""
|
||||
from litellm.integrations.langfuse.langfuse_prompt_management import (
|
||||
LangfusePromptManagement,
|
||||
)
|
||||
|
||||
# Reset all callback lists
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "success_callback", [])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
|
||||
# Add logger to both success_callback and _async_success_callback
|
||||
langfuse_logger = LangfusePromptManagement()
|
||||
litellm.callbacks.append(langfuse_logger)
|
||||
|
||||
# Make a completion call
|
||||
await litellm.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
mock_response="Testing duplicate callbacks",
|
||||
)
|
||||
|
||||
# Assert no new callbacks were added
|
||||
initial_callbacks_len = len(litellm.callbacks)
|
||||
initial_async_success_callback_len = len(litellm._async_success_callback)
|
||||
initial_success_callback_len = len(litellm.success_callback)
|
||||
print(
|
||||
f"Num callbacks before: litellm.callbacks: {len(litellm.callbacks)}, litellm._async_success_callback: {len(litellm._async_success_callback)}, litellm.success_callback: {len(litellm.success_callback)}"
|
||||
)
|
||||
|
||||
for _ in range(10):
|
||||
await litellm.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
mock_response="Testing duplicate callbacks",
|
||||
)
|
||||
|
||||
assert len(litellm.callbacks) == initial_callbacks_len
|
||||
assert len(litellm._async_success_callback) == initial_async_success_callback_len
|
||||
assert len(litellm.success_callback) == initial_success_callback_len
|
||||
|
||||
print(
|
||||
f"Num callbacks after 10 mock calls: litellm.callbacks: {len(litellm.callbacks)}, litellm._async_success_callback: {len(litellm._async_success_callback)}, litellm.success_callback: {len(litellm.success_callback)}"
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.utils import get_applied_guardrails
|
||||
from unittest.mock import Mock
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_get_provider_audio_transcription_config():
|
||||
from litellm.utils import ProviderConfigManager
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
for provider in LlmProviders:
|
||||
config = ProviderConfigManager.get_provider_audio_transcription_config(
|
||||
model="whisper-1", provider=provider
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_get_valid_models_from_dynamic_api_key():
|
||||
"""
|
||||
Test that get_valid_models returns the correct models for a given provider
|
||||
"""
|
||||
from litellm.utils import get_valid_models
|
||||
from litellm.types.router import CredentialLiteLLMParams
|
||||
|
||||
creds = CredentialLiteLLMParams(api_key="123")
|
||||
|
||||
valid_models = get_valid_models(
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=creds,
|
||||
check_provider_endpoint=True,
|
||||
)
|
||||
assert len(valid_models) == 0
|
||||
|
||||
creds = CredentialLiteLLMParams(api_key=os.getenv("ANTHROPIC_API_KEY"))
|
||||
valid_models = get_valid_models(
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=creds,
|
||||
check_provider_endpoint=True,
|
||||
)
|
||||
assert len(valid_models) > 0
|
||||
assert "anthropic/claude-sonnet-4-6" in valid_models
|
||||
|
||||
|
||||
def test_get_whitelisted_models():
|
||||
"""
|
||||
Snapshot of all bedrock models as of 12/24/2024.
|
||||
|
|
|
|||
|
|
@ -347,31 +347,6 @@ class BaseResponsesAPITest(ABC):
|
|||
else:
|
||||
raise ValueError("response is not a ResponsesAPIResponse")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiturn_responses_api(self):
|
||||
litellm.turn_on_debug()
|
||||
litellm.set_verbose = True
|
||||
try:
|
||||
base_completion_call_args = self.get_base_completion_call_args()
|
||||
response_1 = await litellm.aresponses(
|
||||
input="Basic ping", max_output_tokens=20, **base_completion_call_args
|
||||
)
|
||||
|
||||
# follow up with a second request
|
||||
response_1_id = response_1.id
|
||||
response_2 = await litellm.aresponses(
|
||||
input="Basic ping",
|
||||
max_output_tokens=20,
|
||||
previous_response_id=response_1_id,
|
||||
**base_completion_call_args,
|
||||
)
|
||||
|
||||
# assert the response is not None
|
||||
assert response_1 is not None
|
||||
assert response_2 is not None
|
||||
except litellm.InternalServerError:
|
||||
pytest.skip("Skipping test due to litellm.InternalServerError")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_with_tool_calls(self):
|
||||
"""Test that calls the Responses API with tool calls including function call and output"""
|
||||
|
|
@ -447,72 +422,6 @@ class BaseResponsesAPITest(ABC):
|
|||
else:
|
||||
assert len(response["output"]) > 0
|
||||
|
||||
def test_openai_responses_api_dict_input_filtering(self):
|
||||
"""
|
||||
Test that regular dict inputs with status fields are properly filtered
|
||||
to replicate exclude_unset=True behavior for non-Pydantic objects.
|
||||
"""
|
||||
from litellm.llms.openai.responses.transformation import (
|
||||
OpenAIResponsesAPIConfig,
|
||||
)
|
||||
|
||||
# Test input with regular dict objects (like from JSON)
|
||||
test_input = [
|
||||
{"role": "user", "content": "test"},
|
||||
{
|
||||
"id": "rs_123",
|
||||
"summary": [{"text": "test", "type": "summary_text"}],
|
||||
"type": "reasoning",
|
||||
"content": None, # Should be filtered out
|
||||
"encrypted_content": None, # Should be filtered out
|
||||
"status": None, # Should be filtered out
|
||||
},
|
||||
{
|
||||
"arguments": "{}",
|
||||
"call_id": "call_123",
|
||||
"name": "get_today",
|
||||
"type": "function_call",
|
||||
"id": "fc_123",
|
||||
"status": "completed", # Should be preserved (not a default field)
|
||||
},
|
||||
]
|
||||
|
||||
config = OpenAIResponsesAPIConfig()
|
||||
validated_input = config._validate_input_param(test_input)
|
||||
|
||||
# Verify the results
|
||||
assert len(validated_input) == 3
|
||||
|
||||
# Check reasoning item (index 1)
|
||||
reasoning_item = validated_input[1]
|
||||
assert reasoning_item["type"] == "reasoning"
|
||||
assert (
|
||||
"status" not in reasoning_item
|
||||
), "status field should be filtered out from reasoning item"
|
||||
assert (
|
||||
"content" not in reasoning_item
|
||||
), "content field should be filtered out from reasoning item"
|
||||
assert (
|
||||
"encrypted_content" not in reasoning_item
|
||||
), "encrypted_content field should be filtered out from reasoning item"
|
||||
# Note: ID auto-generation was disabled, so reasoning items may not have IDs
|
||||
# Only check for ID if it was present in the original input
|
||||
if "id" in reasoning_item:
|
||||
assert reasoning_item["id"] == "rs_123", "ID should be preserved if present"
|
||||
assert "summary" in reasoning_item, "summary field should be preserved"
|
||||
|
||||
# Check function call item (index 2)
|
||||
function_call_item = validated_input[2]
|
||||
assert function_call_item["type"] == "function_call"
|
||||
assert (
|
||||
"status" in function_call_item
|
||||
), "status field should be preserved in function call item"
|
||||
assert (
|
||||
function_call_item["status"] == "completed"
|
||||
), "status value should be preserved"
|
||||
|
||||
print("✅ OpenAI Responses API dict input filtering test passed")
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [False, True])
|
||||
@pytest.mark.flaky(retries=3, delay=2)
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -565,23 +474,6 @@ class BaseResponsesAPITest(ABC):
|
|||
else:
|
||||
raise e
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel_responses_invalid_response_id(self, sync_mode):
|
||||
"""Test cancel_responses with invalid response ID should raise appropriate error"""
|
||||
base_completion_call_args = self.get_base_completion_call_args()
|
||||
|
||||
if sync_mode:
|
||||
with pytest.raises(openai.APIError):
|
||||
litellm.cancel_responses(
|
||||
response_id="invalid_response_id_12345", **base_completion_call_args
|
||||
)
|
||||
else:
|
||||
with pytest.raises(openai.APIError):
|
||||
await litellm.acancel_responses(
|
||||
response_id="invalid_response_id_12345", **base_completion_call_args
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_context_management_server_side_compaction(self):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,35 +1,20 @@
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from base_responses_api import BaseResponsesAPITest
|
||||
from openai.types.responses.function_tool import FunctionTool
|
||||
|
||||
|
||||
class TestAnthropicResponsesAPITest(BaseResponsesAPITest):
|
||||
test_basic_openai_responses_delete_endpoint = None
|
||||
test_basic_openai_responses_streaming_delete_endpoint = None
|
||||
test_basic_openai_responses_get_endpoint = None
|
||||
test_basic_openai_responses_cancel_endpoint = None
|
||||
|
||||
def get_base_completion_call_args(self):
|
||||
# litellm.turn_on_debug()
|
||||
return {
|
||||
"model": "anthropic/claude-sonnet-4-5",
|
||||
}
|
||||
|
||||
async def test_basic_openai_responses_delete_endpoint(self, sync_mode=False):
|
||||
pytest.skip("DELETE responses is not supported for anthropic")
|
||||
|
||||
async def test_basic_openai_responses_streaming_delete_endpoint(
|
||||
self, sync_mode=False
|
||||
):
|
||||
pytest.skip("DELETE responses is not supported for anthropic")
|
||||
|
||||
async def test_basic_openai_responses_get_endpoint(self, sync_mode=False):
|
||||
pytest.skip("GET responses is not supported for anthropic")
|
||||
|
||||
async def test_basic_openai_responses_cancel_endpoint(self, sync_mode=False):
|
||||
pytest.skip("CANCEL responses is not supported for anthropic")
|
||||
|
||||
async def test_cancel_responses_invalid_response_id(self, sync_mode=False):
|
||||
pytest.skip("CANCEL responses is not supported for anthropic")
|
||||
|
||||
|
||||
def test_multiturn_tool_calls():
|
||||
# Test streaming response with tools for Anthropic
|
||||
litellm.turn_on_debug()
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ from base_responses_api import BaseResponsesAPITest
|
|||
|
||||
|
||||
class TestAzureResponsesAPITest(BaseResponsesAPITest):
|
||||
test_multiturn_responses_api = None
|
||||
test_responses_api_with_tool_calls = None
|
||||
|
||||
def get_base_completion_call_args(self):
|
||||
|
|
|
|||
|
|
@ -1,233 +1,12 @@
|
|||
import os
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
import json
|
||||
from base_responses_api import BaseResponsesAPITest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_basic_google_ai_studio_responses_api_with_tools():
|
||||
litellm.turn_on_debug()
|
||||
litellm.set_verbose = True
|
||||
request_model = "gemini/gemini-2.5-flash"
|
||||
response = await litellm.aresponses(
|
||||
model=request_model,
|
||||
input="what is the latest version of supabase python package and when was it released?",
|
||||
tools=[{"type": "web_search_preview", "search_context_size": "low"}],
|
||||
)
|
||||
print("litellm response=", json.dumps(response, indent=4, default=str))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_3_responses_api_with_thought_signatures():
|
||||
"""
|
||||
Test that Gemini 3 Responses API preserves thought signatures in function calls.
|
||||
This test verifies that provider_specific_fields with thought_signature are correctly
|
||||
preserved when using the Responses API with Gemini 3.
|
||||
"""
|
||||
if not os.getenv("GEMINI_API_KEY"):
|
||||
pytest.skip("GEMINI_API_KEY not set")
|
||||
|
||||
litellm.set_verbose = False
|
||||
request_model = "gemini/gemini-3.1-pro-preview"
|
||||
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"name": "get_weather",
|
||||
"description": "Get the current weather for a location",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {
|
||||
"type": "string",
|
||||
"description": "City and country e.g. Mumbai, India",
|
||||
},
|
||||
"units": {
|
||||
"type": "string",
|
||||
"enum": ["celsius", "fahrenheit"],
|
||||
"description": "Units the temperature will be returned in.",
|
||||
},
|
||||
},
|
||||
"required": ["location", "units"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
"strict": True,
|
||||
}
|
||||
]
|
||||
|
||||
# Step 1: Initial request with tools
|
||||
response = await litellm.aresponses(
|
||||
model=request_model,
|
||||
input="What is the weather in Mumbai?",
|
||||
tools=tools,
|
||||
reasoning_effort="low",
|
||||
)
|
||||
|
||||
# Validate response structure
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
assert isinstance(
|
||||
response, ResponsesAPIResponse
|
||||
), "Response should be a ResponsesAPIResponse"
|
||||
assert (
|
||||
hasattr(response, "output") or "output" in response
|
||||
), "Response should have 'output' field"
|
||||
assert isinstance(response.output, list), "Output should be a list"
|
||||
|
||||
# Find function call in output
|
||||
function_call_item = None
|
||||
for item in response.output:
|
||||
# Convert to dict if it's a Pydantic model for easier access
|
||||
if hasattr(item, "model_dump"):
|
||||
item_dict = item.model_dump()
|
||||
elif hasattr(item, "__dict__"):
|
||||
item_dict = dict(item) if not isinstance(item, dict) else item
|
||||
else:
|
||||
item_dict = item if isinstance(item, dict) else {}
|
||||
|
||||
if isinstance(item_dict, dict) and item_dict.get("type") == "function_call":
|
||||
function_call_item = item_dict
|
||||
break
|
||||
|
||||
# Verify function call exists
|
||||
assert (
|
||||
function_call_item is not None
|
||||
), "Response should contain a function_call item"
|
||||
assert (
|
||||
function_call_item.get("name") == "get_weather"
|
||||
), "Function call should be for get_weather"
|
||||
|
||||
# Verify thought signature is present in provider_specific_fields
|
||||
provider_specific_fields = function_call_item.get("provider_specific_fields")
|
||||
assert (
|
||||
provider_specific_fields is not None
|
||||
), "Function call should have provider_specific_fields"
|
||||
assert (
|
||||
"thought_signature" in provider_specific_fields
|
||||
), "provider_specific_fields should contain thought_signature"
|
||||
assert isinstance(
|
||||
provider_specific_fields["thought_signature"], str
|
||||
), "thought_signature should be a string"
|
||||
assert (
|
||||
len(provider_specific_fields["thought_signature"]) > 0
|
||||
), "thought_signature should not be empty"
|
||||
|
||||
print(
|
||||
f"✅ Thought signature preserved: {provider_specific_fields['thought_signature'][:50]}..."
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_3_responses_api_streaming_with_thought_signatures():
|
||||
"""
|
||||
Test that Gemini 3 Responses API preserves thought signatures in streaming mode.
|
||||
This test verifies that provider_specific_fields with thought_signature are correctly
|
||||
preserved when using streaming Responses API with Gemini 3.
|
||||
"""
|
||||
if not os.getenv("GEMINI_API_KEY"):
|
||||
pytest.skip("GEMINI_API_KEY not set")
|
||||
|
||||
litellm.set_verbose = False
|
||||
request_model = "gemini/gemini-3.1-pro-preview"
|
||||
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"name": "get_weather",
|
||||
"description": "Get the current weather for a location",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {
|
||||
"type": "string",
|
||||
"description": "City and country e.g. Mumbai, India",
|
||||
},
|
||||
"units": {
|
||||
"type": "string",
|
||||
"enum": ["celsius", "fahrenheit"],
|
||||
"description": "Units the temperature will be returned in.",
|
||||
},
|
||||
},
|
||||
"required": ["location", "units"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
"strict": True,
|
||||
}
|
||||
]
|
||||
|
||||
# Step 1: Streaming request with tools
|
||||
response_stream = await litellm.aresponses(
|
||||
model=request_model,
|
||||
input="What is the weather in Mumbai?",
|
||||
tools=tools,
|
||||
stream=True,
|
||||
reasoning_effort="low",
|
||||
)
|
||||
|
||||
# Collect all chunks
|
||||
chunks = []
|
||||
completed_response = None
|
||||
|
||||
async for chunk in response_stream:
|
||||
chunks.append(chunk)
|
||||
# Check if this is the completed response event
|
||||
if hasattr(chunk, "type") and chunk.type == "response.completed":
|
||||
completed_response = chunk.response
|
||||
elif isinstance(chunk, dict) and chunk.get("type") == "response.completed":
|
||||
completed_response = chunk.get("response")
|
||||
|
||||
# Verify we got chunks
|
||||
assert len(chunks) > 0, "Should receive at least one chunk"
|
||||
|
||||
# If we have a completed response, check for thought signatures
|
||||
if completed_response:
|
||||
output = completed_response.get("output", [])
|
||||
function_call_item = None
|
||||
for item in output:
|
||||
if isinstance(item, dict) and item.get("type") == "function_call":
|
||||
function_call_item = item
|
||||
break
|
||||
|
||||
if function_call_item:
|
||||
provider_specific_fields = function_call_item.get(
|
||||
"provider_specific_fields"
|
||||
)
|
||||
if provider_specific_fields:
|
||||
thought_signature = provider_specific_fields.get("thought_signature")
|
||||
if thought_signature:
|
||||
assert isinstance(
|
||||
thought_signature, str
|
||||
), "thought_signature should be a string"
|
||||
assert (
|
||||
len(thought_signature) > 0
|
||||
), "thought_signature should not be empty"
|
||||
print(
|
||||
f"✅ Streaming thought signature preserved: {thought_signature[:50]}..."
|
||||
)
|
||||
|
||||
print(f"✅ Collected {len(chunks)} streaming chunks")
|
||||
|
||||
|
||||
class TestGoogleAIStudioResponsesAPITest(BaseResponsesAPITest):
|
||||
test_basic_openai_responses_delete_endpoint = None
|
||||
test_basic_openai_responses_streaming_delete_endpoint = None
|
||||
test_basic_openai_responses_get_endpoint = None
|
||||
test_basic_openai_responses_cancel_endpoint = None
|
||||
|
||||
def get_base_completion_call_args(self):
|
||||
# litellm.turn_on_debug()
|
||||
return {"model": "gemini/gemini-2.5-flash-lite"}
|
||||
|
||||
async def test_basic_openai_responses_delete_endpoint(self, sync_mode=False):
|
||||
pytest.skip("DELETE responses is not supported for Google AI Studio")
|
||||
|
||||
async def test_basic_openai_responses_streaming_delete_endpoint(
|
||||
self, sync_mode=False
|
||||
):
|
||||
pytest.skip("DELETE responses is not supported for Google AI Studio")
|
||||
|
||||
async def test_basic_openai_responses_get_endpoint(self, sync_mode=False):
|
||||
pytest.skip("GET responses is not supported for Google AI Studio")
|
||||
|
||||
async def test_basic_openai_responses_cancel_endpoint(self, sync_mode=False):
|
||||
pytest.skip("CANCEL responses is not supported for Google AI Studio")
|
||||
|
||||
async def test_cancel_responses_invalid_response_id(self, sync_mode=False):
|
||||
pytest.skip("CANCEL responses is not supported for Google AI Studio")
|
||||
|
|
|
|||
|
|
@ -1,20 +1,9 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from typing import Optional, cast
|
||||
|
||||
import pytest
|
||||
from base_responses_api import BaseResponsesAPITest, validate_responses_api_response
|
||||
from base_responses_api import BaseResponsesAPITest
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseAPIUsage,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
)
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
|
||||
class TestOpenAIResponsesAPITest(BaseResponsesAPITest):
|
||||
|
|
@ -29,736 +18,6 @@ class TestOpenAIResponsesAPITest(BaseResponsesAPITest):
|
|||
return "openai/gpt-5.2"
|
||||
|
||||
|
||||
class TestCustomLogger(CustomLogger):
|
||||
def __init__(
|
||||
self,
|
||||
):
|
||||
self.standard_logging_object: Optional[StandardLoggingPayload] = None
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
print("in async_log_success_event")
|
||||
print("kwargs=", json.dumps(kwargs, indent=4, default=str))
|
||||
self.standard_logging_object = kwargs["standard_logging_object"]
|
||||
pass
|
||||
|
||||
|
||||
def validate_standard_logging_payload(
|
||||
slp: StandardLoggingPayload, response: ResponsesAPIResponse, request_model: str
|
||||
):
|
||||
"""
|
||||
Validate that a StandardLoggingPayload object matches the expected response
|
||||
|
||||
Args:
|
||||
slp (StandardLoggingPayload): The standard logging payload object to validate
|
||||
response (dict): The litellm response to compare against
|
||||
request_model (str): The model name that was requested
|
||||
"""
|
||||
# Validate payload exists
|
||||
assert slp is not None, "Standard logging payload should not be None"
|
||||
|
||||
# Validate token counts
|
||||
print(
|
||||
"VALIDATING STANDARD LOGGING PAYLOAD. response=",
|
||||
json.dumps(response, indent=4, default=str),
|
||||
)
|
||||
print("FIELDS IN SLP=", json.dumps(slp, indent=4, default=str))
|
||||
print("SLP PROMPT TOKENS=", slp["prompt_tokens"])
|
||||
print("RESPONSE PROMPT TOKENS=", response["usage"]["input_tokens"])
|
||||
assert (
|
||||
slp["prompt_tokens"] == response["usage"]["input_tokens"]
|
||||
), "Prompt tokens mismatch"
|
||||
assert (
|
||||
slp["completion_tokens"] == response["usage"]["output_tokens"]
|
||||
), "Completion tokens mismatch"
|
||||
assert (
|
||||
slp["total_tokens"]
|
||||
== response["usage"]["input_tokens"] + response["usage"]["output_tokens"]
|
||||
), "Total tokens mismatch"
|
||||
|
||||
# Validate spend and response metadata
|
||||
assert slp["response_cost"] > 0, "Response cost should be greater than 0"
|
||||
assert slp["id"] == response["id"], "Response ID mismatch"
|
||||
assert slp["model"] == request_model, "Model name mismatch"
|
||||
|
||||
# Validate messages
|
||||
assert slp["messages"] == [{"content": "hi", "role": "user"}], "Messages mismatch"
|
||||
|
||||
# Validate complete response structure
|
||||
validate_responses_match(slp["response"], response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
def test_basic_openai_responses_api_streaming_with_logging():
|
||||
litellm.turn_on_debug()
|
||||
litellm.set_verbose = True
|
||||
test_custom_logger = TestCustomLogger()
|
||||
litellm.callbacks = [test_custom_logger]
|
||||
request_model = "gpt-5.5"
|
||||
response = litellm.responses(
|
||||
model=request_model,
|
||||
input="hi",
|
||||
stream=True,
|
||||
)
|
||||
final_response: Optional[ResponseCompletedEvent] = None
|
||||
for event in response:
|
||||
if event.type == "response.completed":
|
||||
final_response = event
|
||||
print("litellm response=", json.dumps(event, indent=4, default=str))
|
||||
|
||||
print("sleeping for 2 seconds...")
|
||||
time.sleep(2)
|
||||
print(
|
||||
"standard logging payload=",
|
||||
json.dumps(test_custom_logger.standard_logging_object, indent=4, default=str),
|
||||
)
|
||||
|
||||
assert final_response is not None
|
||||
assert test_custom_logger.standard_logging_object is not None
|
||||
|
||||
validate_standard_logging_payload(
|
||||
slp=test_custom_logger.standard_logging_object,
|
||||
response=final_response.response,
|
||||
request_model=request_model,
|
||||
)
|
||||
|
||||
|
||||
def validate_responses_match(slp_response, litellm_response):
|
||||
"""Validate that the standard logging payload OpenAI response matches the litellm response"""
|
||||
# Validate core fields
|
||||
assert slp_response["id"] == litellm_response["id"], "ID mismatch"
|
||||
assert slp_response["model"] == litellm_response["model"], "Model mismatch"
|
||||
assert (
|
||||
slp_response["created_at"] == litellm_response["created_at"]
|
||||
), "Created at mismatch"
|
||||
|
||||
# Validate usage
|
||||
assert (
|
||||
slp_response["usage"]["prompt_tokens"]
|
||||
== litellm_response["usage"]["input_tokens"]
|
||||
), "Input tokens mismatch"
|
||||
assert (
|
||||
slp_response["usage"]["completion_tokens"]
|
||||
== litellm_response["usage"]["output_tokens"]
|
||||
), "Output tokens mismatch"
|
||||
assert (
|
||||
slp_response["usage"]["total_tokens"]
|
||||
== litellm_response["usage"]["total_tokens"]
|
||||
), "Total tokens mismatch"
|
||||
|
||||
# Validate output/messages
|
||||
assert len(slp_response["output"]) == len(
|
||||
litellm_response["output"]
|
||||
), "Output length mismatch"
|
||||
for slp_msg, litellm_msg in zip(slp_response["output"], litellm_response["output"]):
|
||||
assert slp_msg["role"] == litellm_msg.role, "Message role mismatch"
|
||||
# Access the content's text field for the litellm response
|
||||
litellm_content = litellm_msg.content[0].text if litellm_msg.content else ""
|
||||
assert (
|
||||
slp_msg["content"][0]["text"] == litellm_content
|
||||
), f"Message content mismatch. Expected {litellm_content}, Got {slp_msg['content']}"
|
||||
assert slp_msg["status"] == litellm_msg.status, "Message status mismatch"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_basic_openai_responses_api_non_streaming_with_logging():
|
||||
litellm.turn_on_debug()
|
||||
litellm.set_verbose = True
|
||||
test_custom_logger = TestCustomLogger()
|
||||
litellm.callbacks = [test_custom_logger]
|
||||
request_model = "gpt-5.5"
|
||||
response = await litellm.aresponses(
|
||||
model=request_model,
|
||||
input="hi",
|
||||
)
|
||||
|
||||
print("litellm response=", json.dumps(response, indent=4, default=str))
|
||||
print("response hidden params=", response._hidden_params)
|
||||
|
||||
print("sleeping for 2 seconds...")
|
||||
await asyncio.sleep(5)
|
||||
print(
|
||||
"standard logging payload=",
|
||||
json.dumps(test_custom_logger.standard_logging_object, indent=4, default=str),
|
||||
)
|
||||
print("response usage=", response.usage)
|
||||
|
||||
assert response is not None
|
||||
assert test_custom_logger.standard_logging_object is not None
|
||||
|
||||
validate_standard_logging_payload(
|
||||
test_custom_logger.standard_logging_object, response, request_model
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_responses_api_returns_headers(sync_mode):
|
||||
"""
|
||||
Test that OpenAI responses API returns OpenAI headers in _hidden_params.
|
||||
This ensures the proxy can forward these headers to clients.
|
||||
|
||||
Related issue: LiteLLM responses API should return OpenAI headers like chat completions does
|
||||
"""
|
||||
litellm.turn_on_debug()
|
||||
litellm.set_verbose = True
|
||||
|
||||
if sync_mode:
|
||||
response = litellm.responses(
|
||||
model="gpt-5.5",
|
||||
input="Say hello",
|
||||
max_output_tokens=20,
|
||||
)
|
||||
else:
|
||||
response = await litellm.aresponses(
|
||||
model="gpt-5.5",
|
||||
input="Say hello",
|
||||
max_output_tokens=20,
|
||||
)
|
||||
|
||||
# Verify response is valid
|
||||
assert response is not None
|
||||
assert isinstance(response, ResponsesAPIResponse)
|
||||
|
||||
# Verify _hidden_params exists
|
||||
assert hasattr(
|
||||
response, "_hidden_params"
|
||||
), "Response should have _hidden_params attribute"
|
||||
assert response._hidden_params is not None, "_hidden_params should not be None"
|
||||
|
||||
# Verify additional_headers exists in _hidden_params
|
||||
assert (
|
||||
"additional_headers" in response._hidden_params
|
||||
), "_hidden_params should contain 'additional_headers' key"
|
||||
|
||||
additional_headers = response._hidden_params["additional_headers"]
|
||||
assert isinstance(
|
||||
additional_headers, dict
|
||||
), "additional_headers should be a dictionary"
|
||||
assert len(additional_headers) > 0, "additional_headers should not be empty"
|
||||
|
||||
# Check for expected OpenAI rate limit headers
|
||||
# These can be either direct (x-ratelimit-*) or prefixed (llm_provider-x-ratelimit-*)
|
||||
rate_limit_headers = [
|
||||
"x-ratelimit-remaining-tokens",
|
||||
"x-ratelimit-limit-tokens",
|
||||
"x-ratelimit-remaining-requests",
|
||||
"x-ratelimit-limit-requests",
|
||||
]
|
||||
|
||||
found_headers = []
|
||||
for header_name in rate_limit_headers:
|
||||
if header_name in additional_headers:
|
||||
found_headers.append(header_name)
|
||||
elif f"llm_provider-{header_name}" in additional_headers:
|
||||
found_headers.append(f"llm_provider-{header_name}")
|
||||
|
||||
assert (
|
||||
len(found_headers) > 0
|
||||
), f"Should find at least one OpenAI rate limit header. Headers found: {list(additional_headers.keys())}"
|
||||
|
||||
# Verify headers key also exists (raw headers)
|
||||
assert (
|
||||
"headers" in response._hidden_params
|
||||
), "_hidden_params should contain 'headers' key with raw response headers"
|
||||
|
||||
print(
|
||||
f"✓ Successfully validated OpenAI headers in {'sync' if sync_mode else 'async'} mode"
|
||||
)
|
||||
print(f" Found {len(additional_headers)} headers total")
|
||||
print(f" Rate limit headers found: {found_headers}")
|
||||
|
||||
|
||||
def validate_stream_event(event):
|
||||
"""
|
||||
Validate that a streaming event from litellm.responses() or litellm.aresponses()
|
||||
with stream=True conforms to the expected structure based on its event type.
|
||||
|
||||
Args:
|
||||
event: The streaming event object to validate
|
||||
|
||||
Raises:
|
||||
AssertionError: If the event doesn't match the expected structure for its type
|
||||
"""
|
||||
# Common validation for all event types
|
||||
assert hasattr(event, "type"), "Event should have a 'type' attribute"
|
||||
|
||||
# Type-specific validation
|
||||
if event.type == "response.created" or event.type == "response.in_progress":
|
||||
assert hasattr(
|
||||
event, "response"
|
||||
), f"{event.type} event should have a 'response' attribute"
|
||||
validate_responses_api_response(event.response, final_chunk=False)
|
||||
|
||||
elif event.type == "response.completed":
|
||||
assert hasattr(
|
||||
event, "response"
|
||||
), "response.completed event should have a 'response' attribute"
|
||||
validate_responses_api_response(event.response, final_chunk=True)
|
||||
# Usage is guaranteed only on the completed event
|
||||
assert (
|
||||
"usage" in event.response
|
||||
), "response.completed event should have usage information"
|
||||
print("Usage in event.response=", event.response["usage"])
|
||||
assert isinstance(event.response["usage"], ResponseAPIUsage)
|
||||
elif event.type == "response.failed" or event.type == "response.incomplete":
|
||||
assert hasattr(
|
||||
event, "response"
|
||||
), f"{event.type} event should have a 'response' attribute"
|
||||
|
||||
elif (
|
||||
event.type == "response.output_item.added"
|
||||
or event.type == "response.output_item.done"
|
||||
):
|
||||
assert hasattr(
|
||||
event, "output_index"
|
||||
), f"{event.type} event should have an 'output_index' attribute"
|
||||
assert hasattr(
|
||||
event, "item"
|
||||
), f"{event.type} event should have an 'item' attribute"
|
||||
|
||||
elif (
|
||||
event.type == "response.content_part.added"
|
||||
or event.type == "response.content_part.done"
|
||||
):
|
||||
assert hasattr(
|
||||
event, "item_id"
|
||||
), f"{event.type} event should have an 'item_id' attribute"
|
||||
assert hasattr(
|
||||
event, "output_index"
|
||||
), f"{event.type} event should have an 'output_index' attribute"
|
||||
assert hasattr(
|
||||
event, "content_index"
|
||||
), f"{event.type} event should have a 'content_index' attribute"
|
||||
assert hasattr(
|
||||
event, "part"
|
||||
), f"{event.type} event should have a 'part' attribute"
|
||||
|
||||
elif event.type == "response.output_text.delta":
|
||||
assert hasattr(
|
||||
event, "item_id"
|
||||
), f"{event.type} event should have an 'item_id' attribute"
|
||||
assert hasattr(
|
||||
event, "output_index"
|
||||
), f"{event.type} event should have an 'output_index' attribute"
|
||||
assert hasattr(
|
||||
event, "content_index"
|
||||
), f"{event.type} event should have a 'content_index' attribute"
|
||||
assert hasattr(
|
||||
event, "delta"
|
||||
), f"{event.type} event should have a 'delta' attribute"
|
||||
|
||||
elif event.type == "response.output_text.annotation.added":
|
||||
assert hasattr(
|
||||
event, "item_id"
|
||||
), f"{event.type} event should have an 'item_id' attribute"
|
||||
assert hasattr(
|
||||
event, "output_index"
|
||||
), f"{event.type} event should have an 'output_index' attribute"
|
||||
assert hasattr(
|
||||
event, "content_index"
|
||||
), f"{event.type} event should have a 'content_index' attribute"
|
||||
assert hasattr(
|
||||
event, "annotation_index"
|
||||
), f"{event.type} event should have an 'annotation_index' attribute"
|
||||
assert hasattr(
|
||||
event, "annotation"
|
||||
), f"{event.type} event should have an 'annotation' attribute"
|
||||
|
||||
elif event.type == "response.output_text.done":
|
||||
assert hasattr(
|
||||
event, "item_id"
|
||||
), f"{event.type} event should have an 'item_id' attribute"
|
||||
assert hasattr(
|
||||
event, "output_index"
|
||||
), f"{event.type} event should have an 'output_index' attribute"
|
||||
assert hasattr(
|
||||
event, "content_index"
|
||||
), f"{event.type} event should have a 'content_index' attribute"
|
||||
assert hasattr(
|
||||
event, "text"
|
||||
), f"{event.type} event should have a 'text' attribute"
|
||||
|
||||
elif event.type == "response.refusal.delta":
|
||||
assert hasattr(
|
||||
event, "item_id"
|
||||
), f"{event.type} event should have an 'item_id' attribute"
|
||||
assert hasattr(
|
||||
event, "output_index"
|
||||
), f"{event.type} event should have an 'output_index' attribute"
|
||||
assert hasattr(
|
||||
event, "content_index"
|
||||
), f"{event.type} event should have a 'content_index' attribute"
|
||||
assert hasattr(
|
||||
event, "delta"
|
||||
), f"{event.type} event should have a 'delta' attribute"
|
||||
|
||||
elif event.type == "response.refusal.done":
|
||||
assert hasattr(
|
||||
event, "item_id"
|
||||
), f"{event.type} event should have an 'item_id' attribute"
|
||||
assert hasattr(
|
||||
event, "output_index"
|
||||
), f"{event.type} event should have an 'output_index' attribute"
|
||||
assert hasattr(
|
||||
event, "content_index"
|
||||
), f"{event.type} event should have a 'content_index' attribute"
|
||||
assert hasattr(
|
||||
event, "refusal"
|
||||
), f"{event.type} event should have a 'refusal' attribute"
|
||||
|
||||
elif event.type == "response.function_call_arguments.delta":
|
||||
assert hasattr(
|
||||
event, "item_id"
|
||||
), f"{event.type} event should have an 'item_id' attribute"
|
||||
assert hasattr(
|
||||
event, "output_index"
|
||||
), f"{event.type} event should have an 'output_index' attribute"
|
||||
assert hasattr(
|
||||
event, "delta"
|
||||
), f"{event.type} event should have a 'delta' attribute"
|
||||
|
||||
elif event.type == "response.function_call_arguments.done":
|
||||
assert hasattr(
|
||||
event, "item_id"
|
||||
), f"{event.type} event should have an 'item_id' attribute"
|
||||
assert hasattr(
|
||||
event, "output_index"
|
||||
), f"{event.type} event should have an 'output_index' attribute"
|
||||
assert hasattr(
|
||||
event, "arguments"
|
||||
), f"{event.type} event should have an 'arguments' attribute"
|
||||
|
||||
elif event.type in [
|
||||
"response.file_search_call.in_progress",
|
||||
"response.file_search_call.searching",
|
||||
"response.file_search_call.completed",
|
||||
"response.web_search_call.in_progress",
|
||||
"response.web_search_call.searching",
|
||||
"response.web_search_call.completed",
|
||||
]:
|
||||
assert hasattr(
|
||||
event, "output_index"
|
||||
), f"{event.type} event should have an 'output_index' attribute"
|
||||
assert hasattr(
|
||||
event, "item_id"
|
||||
), f"{event.type} event should have an 'item_id' attribute"
|
||||
|
||||
elif event.type == "error":
|
||||
assert hasattr(
|
||||
event, "message"
|
||||
), "Error event should have a 'message' attribute"
|
||||
return True # Return True if validation passes
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_responses_api_streaming_validation(sync_mode):
|
||||
"""Test that validates each streaming event from the responses API"""
|
||||
litellm.turn_on_debug()
|
||||
|
||||
event_types_seen = set()
|
||||
|
||||
if sync_mode:
|
||||
response = litellm.responses(
|
||||
model="gpt-5.5",
|
||||
input="Tell me about artificial intelligence in 3 sentences.",
|
||||
stream=True,
|
||||
)
|
||||
for event in response:
|
||||
print(f"Validating event type: {event.type}")
|
||||
validate_stream_event(event)
|
||||
event_types_seen.add(event.type)
|
||||
else:
|
||||
response = await litellm.aresponses(
|
||||
model="gpt-5.5",
|
||||
input="Tell me about artificial intelligence in 3 sentences.",
|
||||
stream=True,
|
||||
)
|
||||
async for event in response:
|
||||
print(f"Validating event type: {event.type}")
|
||||
validate_stream_event(event)
|
||||
event_types_seen.add(event.type)
|
||||
|
||||
# At minimum, we should see these core event types
|
||||
required_events = {"response.created", "response.completed"}
|
||||
|
||||
missing_events = required_events - event_types_seen
|
||||
assert not missing_events, f"Missing required event types: {missing_events}"
|
||||
|
||||
print(f"Successfully validated all event types: {event_types_seen}")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_responses_litellm_router(sync_mode):
|
||||
"""
|
||||
Test the OpenAI responses API with LiteLLM Router in both sync and async modes
|
||||
"""
|
||||
litellm.turn_on_debug()
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt4o-special-alias",
|
||||
"litellm_params": {
|
||||
"model": "gpt-5.5",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
# Call the handler
|
||||
if sync_mode:
|
||||
response = router.responses(
|
||||
model="gpt4o-special-alias",
|
||||
input="Hello, can you tell me a short joke?",
|
||||
max_output_tokens=100,
|
||||
)
|
||||
print("SYNC MODE RESPONSE=", response)
|
||||
else:
|
||||
response = await router.aresponses(
|
||||
model="gpt4o-special-alias",
|
||||
input="Hello, can you tell me a short joke?",
|
||||
max_output_tokens=100,
|
||||
)
|
||||
|
||||
print(
|
||||
f"Router {'sync' if sync_mode else 'async'} response=",
|
||||
json.dumps(response, indent=4, default=str),
|
||||
)
|
||||
|
||||
# Use the helper function to validate the response
|
||||
validate_responses_api_response(response, final_chunk=True)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_responses_litellm_router_streaming(sync_mode):
|
||||
"""
|
||||
Test the OpenAI responses API with streaming through LiteLLM Router
|
||||
"""
|
||||
litellm.turn_on_debug()
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt4o-special-alias",
|
||||
"litellm_params": {
|
||||
"model": "gpt-5.5",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
event_types_seen = set()
|
||||
|
||||
if sync_mode:
|
||||
response = router.responses(
|
||||
model="gpt4o-special-alias",
|
||||
input="Tell me about artificial intelligence in 2 sentences.",
|
||||
stream=True,
|
||||
)
|
||||
for event in response:
|
||||
print(f"Validating event type: {event.type}")
|
||||
validate_stream_event(event)
|
||||
event_types_seen.add(event.type)
|
||||
else:
|
||||
response = await router.aresponses(
|
||||
model="gpt4o-special-alias",
|
||||
input="Tell me about artificial intelligence in 2 sentences.",
|
||||
stream=True,
|
||||
)
|
||||
async for event in response:
|
||||
print(f"Validating event type: {event.type}")
|
||||
validate_stream_event(event)
|
||||
event_types_seen.add(event.type)
|
||||
|
||||
# At minimum, we should see these core event types
|
||||
required_events = {"response.created", "response.completed"}
|
||||
|
||||
missing_events = required_events - event_types_seen
|
||||
assert not missing_events, f"Missing required event types: {missing_events}"
|
||||
|
||||
print(f"Successfully validated all event types: {event_types_seen}")
|
||||
|
||||
|
||||
def test_mcp_tools_with_responses_api():
|
||||
litellm.turn_on_debug()
|
||||
MCP_TOOLS = [
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_label": "zapier",
|
||||
"server_url": "https://mcp.zapier.com/api/mcp/mcp",
|
||||
"headers": {
|
||||
"Authorization": f"Bearer {os.getenv('ZAPIER_CI_CD_MCP_TOKEN')}"
|
||||
},
|
||||
}
|
||||
]
|
||||
MODEL = "openai/gpt-4.1"
|
||||
USER_QUERY = "how does tiktoken work?"
|
||||
#########################################################
|
||||
# Step 1: OpenAI will use MCP LIST, and return a list of MCP calls for our approval
|
||||
try:
|
||||
response = litellm.responses(model=MODEL, tools=MCP_TOOLS, input=USER_QUERY)
|
||||
print(response)
|
||||
|
||||
response = cast(ResponsesAPIResponse, response)
|
||||
|
||||
mcp_approval_id: Optional[str] = None
|
||||
for output in response.output:
|
||||
if output.type == "mcp_approval_request":
|
||||
mcp_approval_id = output.id
|
||||
break
|
||||
|
||||
# Step 2: Send followup with approval for the MCP call
|
||||
if mcp_approval_id:
|
||||
response_with_mcp_call = litellm.responses(
|
||||
model=MODEL,
|
||||
tools=MCP_TOOLS,
|
||||
input=[
|
||||
{
|
||||
"type": "mcp_approval_response",
|
||||
"approve": True,
|
||||
"approval_request_id": mcp_approval_id,
|
||||
}
|
||||
],
|
||||
previous_response_id=response.id,
|
||||
)
|
||||
print(response_with_mcp_call)
|
||||
except litellm.APIError as e:
|
||||
if (
|
||||
"424" in str(e)
|
||||
or "Failed Dependency" in str(e)
|
||||
or "external_connector_error" in str(e)
|
||||
):
|
||||
pytest.skip(f"Skipping test due to external MCP server error: {e}")
|
||||
else:
|
||||
raise e
|
||||
except litellm.InternalServerError as e:
|
||||
if "500" in str(e) or "server_error" in str(e):
|
||||
pytest.skip(
|
||||
f"Skipping test due to OpenAI server error (likely MCP server unavailable): {e}"
|
||||
)
|
||||
else:
|
||||
raise e
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_responses_api_field_types():
|
||||
"""Test that specific fields in the response have the correct types"""
|
||||
litellm.turn_on_debug()
|
||||
litellm.set_verbose = True
|
||||
|
||||
# Test with store=True
|
||||
response = await litellm.aresponses(
|
||||
model="gpt-5.5",
|
||||
input="hi",
|
||||
)
|
||||
|
||||
# Verify created_at is an integer
|
||||
assert isinstance(response.created_at, int), "created_at should be an integer"
|
||||
|
||||
# Verify store field is present and matches input
|
||||
assert hasattr(response, "store"), "store field should be present"
|
||||
assert response.store is True, "store field should match input value"
|
||||
|
||||
# Test without store parameter
|
||||
response_without_store = await litellm.aresponses(model="gpt-5.5", input="hi")
|
||||
|
||||
# Verify created_at is still an integer
|
||||
assert isinstance(
|
||||
response_without_store.created_at, int
|
||||
), "created_at should be an integer"
|
||||
|
||||
# Verify store field is present but None when not specified
|
||||
assert hasattr(response_without_store, "store"), "store field should be present"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_responses_api_token_limit_error():
|
||||
"""
|
||||
Relevant issue: https://github.com/BerriAI/litellm/issues/15785
|
||||
|
||||
Parsing the in-stream ErrorEvent must not raise
|
||||
"pydantic_core._pydantic_core.ValidationError: 3 validation errors for ErrorEvent".
|
||||
The iterator routes the event through litellm.exception_type, so it surfaces as
|
||||
the typed 400 client error the non-streaming path raises (litellm.BadRequestError)
|
||||
carrying the provider's message. invalid_request_error is a non-retriable client
|
||||
error, so there is no MidStreamFallbackError wrapping.
|
||||
"""
|
||||
litellm.turn_on_debug()
|
||||
|
||||
# Generate text with >400k tokens to trigger token limit error
|
||||
oversized_text = "This is a test sentence. " * 50000 # ~400k tokens
|
||||
|
||||
response = await litellm.aresponses(
|
||||
model="gpt-5-mini", input=oversized_text, stream=True
|
||||
)
|
||||
|
||||
async def _drain():
|
||||
async for event in response:
|
||||
print(event)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError) as exc_info:
|
||||
await _drain()
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "exceeds the context window" in str(exc_info.value)
|
||||
|
||||
|
||||
async def test_openai_streaming_logging():
|
||||
"""Test that OpenAI Responses API streaming logging is working correctly."""
|
||||
litellm.turn_on_debug()
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
class TestCustomLogger(CustomLogger):
|
||||
validate_usage = False
|
||||
|
||||
def __init__(self):
|
||||
self.standard_logging_object: Optional[StandardLoggingPayload] = None
|
||||
|
||||
async def async_log_success_event(
|
||||
self, kwargs, response_obj, start_time, end_time
|
||||
):
|
||||
print(f"response_obj: {response_obj.usage}")
|
||||
assert isinstance(
|
||||
response_obj.usage, (Usage, dict)
|
||||
), f"Expected response_obj.usage to be of type Usage or dict, but got {type(response_obj.usage)}"
|
||||
# Verify it has the chat completion format fields
|
||||
if isinstance(response_obj.usage, dict):
|
||||
assert (
|
||||
"prompt_tokens" in response_obj.usage
|
||||
), "Usage dict should have prompt_tokens"
|
||||
assert (
|
||||
"completion_tokens" in response_obj.usage
|
||||
), "Usage dict should have completion_tokens"
|
||||
print("\n\nVALIDATED USAGE\n\n")
|
||||
self.validate_usage = True
|
||||
|
||||
tcl = TestCustomLogger()
|
||||
litellm.callbacks = [tcl]
|
||||
request_model = "gpt-5-mini"
|
||||
response = await litellm.aresponses(
|
||||
model=request_model,
|
||||
input="What is the capital of France?",
|
||||
stream=True,
|
||||
)
|
||||
print("response=", json.dumps(response, indent=4, default=str))
|
||||
|
||||
async for event in response:
|
||||
if event.type == "response.completed":
|
||||
final_response = event
|
||||
print("litellm response=", json.dumps(event, indent=4, default=str))
|
||||
|
||||
await asyncio.sleep(2)
|
||||
assert tcl.validate_usage, "Usage should be validated"
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
async def test_openai_compact_responses_api(sync_mode):
|
||||
|
|
|
|||
Binary file not shown.
|
|
@ -1,261 +0,0 @@
|
|||
import os
|
||||
import json
|
||||
import traceback
|
||||
from typing import Optional
|
||||
from dotenv import load_dotenv
|
||||
from fastapi import Request
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, patch, MagicMock
|
||||
|
||||
from litellm import Router, CustomLogger
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
## Get the current directory of the file being run
|
||||
pwd = os.path.dirname(os.path.realpath(__file__))
|
||||
print(pwd)
|
||||
|
||||
file_path = os.path.join(pwd, "gettysburg.wav")
|
||||
|
||||
audio_file = open(file_path, "rb")
|
||||
from pathlib import Path
|
||||
import litellm
|
||||
import pytest
|
||||
import asyncio
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def model_list():
|
||||
return [
|
||||
{
|
||||
"model_name": "gpt-5-mini",
|
||||
"litellm_params": {
|
||||
"model": "gpt-5-mini",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.5",
|
||||
"litellm_params": {
|
||||
"model": "gpt-5.5",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-image-1",
|
||||
"litellm_params": {
|
||||
"model": "gpt-image-1",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "cohere-rerank",
|
||||
"litellm_params": {
|
||||
"model": "cohere/rerank-english-v3.0",
|
||||
"api_key": os.getenv("COHERE_API_KEY"),
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "claude-sonnet-4-5-20250929",
|
||||
"litellm_params": {
|
||||
"model": "gpt-5-mini",
|
||||
"mock_response": "hi this is macintosh.",
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
# This file includes the custom callbacks for LiteLLM Proxy
|
||||
# Once defined, these can be passed in proxy_config.yaml
|
||||
class MyCustomHandler(CustomLogger):
|
||||
def __init__(self):
|
||||
self.openai_client = None
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
# init logging config
|
||||
print("logging a transcript kwargs: ", kwargs)
|
||||
print("openai client=", kwargs.get("client"))
|
||||
self.openai_client = kwargs.get("client")
|
||||
self.standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object"
|
||||
)
|
||||
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
# Set litellm.callbacks = [proxy_handler_instance] on the proxy
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.flaky(retries=6, delay=10)
|
||||
async def test_transcription_on_router():
|
||||
proxy_handler_instance = MyCustomHandler()
|
||||
litellm.set_verbose = True
|
||||
litellm.callbacks = [proxy_handler_instance]
|
||||
print("\n Testing async transcription on router\n")
|
||||
try:
|
||||
model_list = [
|
||||
{
|
||||
"model_name": "whisper",
|
||||
"litellm_params": {
|
||||
"model": "whisper-1",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "whisper",
|
||||
"litellm_params": {
|
||||
"model": "azure/azure-whisper",
|
||||
"api_base": "https://my-endpoint-europe-berri-992.openai.azure.com/",
|
||||
"api_key": os.getenv("AZURE_EUROPE_API_KEY"),
|
||||
"api_version": "2024-02-15-preview",
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
router = Router(model_list=model_list)
|
||||
|
||||
router_level_clients = []
|
||||
for deployment in router.model_list:
|
||||
_deployment_openai_client = router._get_client(
|
||||
deployment=deployment,
|
||||
kwargs={"model": "whisper-1"},
|
||||
client_type="async",
|
||||
)
|
||||
|
||||
router_level_clients.append(str(_deployment_openai_client))
|
||||
|
||||
## test 1: user facing function
|
||||
response = await router.atranscription(
|
||||
model="whisper",
|
||||
file=audio_file,
|
||||
)
|
||||
|
||||
## test 2: underlying function
|
||||
response = await router._atranscription(
|
||||
model="whisper",
|
||||
file=audio_file,
|
||||
)
|
||||
print(response)
|
||||
|
||||
# PROD Test
|
||||
# Ensure we ONLY use OpenAI/Azure client initialized on the router level
|
||||
await asyncio.sleep(5)
|
||||
print("OpenAI Client used= ", proxy_handler_instance.openai_client)
|
||||
print("all router level clients= ", router_level_clients)
|
||||
assert proxy_handler_instance.openai_client in router_level_clients
|
||||
except Exception as e:
|
||||
traceback.print_exc()
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", ["iterator"]) # "file",
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_speech_router(mode):
|
||||
litellm.set_verbose = True
|
||||
test_logger = MyCustomHandler()
|
||||
litellm.callbacks = [test_logger]
|
||||
from litellm import Router
|
||||
|
||||
client = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "tts",
|
||||
"litellm_params": {
|
||||
"model": "openai/tts-1",
|
||||
},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
response = await client.aspeech(
|
||||
model="tts",
|
||||
voice="alloy",
|
||||
input="the quick brown fox jumped over the lazy dogs",
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
organization=None,
|
||||
project=None,
|
||||
max_retries=1,
|
||||
timeout=600,
|
||||
client=None,
|
||||
optional_params={},
|
||||
)
|
||||
|
||||
await asyncio.sleep(3)
|
||||
|
||||
from litellm.llms.openai.openai import HttpxBinaryResponseContent
|
||||
|
||||
assert isinstance(response, HttpxBinaryResponseContent)
|
||||
|
||||
assert test_logger.standard_logging_object is not None
|
||||
print(
|
||||
"standard_logging_object=",
|
||||
json.dumps(test_logger.standard_logging_object, indent=4),
|
||||
)
|
||||
assert test_logger.standard_logging_object["model_group"] == "tts"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
async def test_rerank_endpoint(model_list):
|
||||
from litellm.types.utils import RerankResponse
|
||||
|
||||
router = Router(model_list=model_list)
|
||||
|
||||
## Test 1: user facing function
|
||||
response = await router.arerank(
|
||||
model="cohere-rerank",
|
||||
query="hello",
|
||||
documents=["hello", "world"],
|
||||
top_n=3,
|
||||
)
|
||||
|
||||
## Test 2: underlying function
|
||||
response = await router._arerank(
|
||||
model="cohere-rerank",
|
||||
query="hello",
|
||||
documents=["hello", "world"],
|
||||
top_n=3,
|
||||
)
|
||||
|
||||
print("async re rank response: ", response)
|
||||
|
||||
assert response.id is not None
|
||||
assert response.results is not None
|
||||
|
||||
RerankResponse.model_validate(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
@pytest.mark.parametrize(
|
||||
"model", ["omni-moderation-latest", "openai/omni-moderation-latest", None]
|
||||
)
|
||||
async def test_moderation_endpoint(model):
|
||||
litellm.set_verbose = True
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "openai/*",
|
||||
"litellm_params": {
|
||||
"model": "openai/*",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "*",
|
||||
"litellm_params": {
|
||||
"model": "openai/*",
|
||||
},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
if model is None:
|
||||
response = await router.amoderation(input="hello this is a test")
|
||||
else:
|
||||
response = await router.amoderation(model=model, input="hello this is a test")
|
||||
|
||||
print("moderation response: ", response)
|
||||
|
|
@ -1,526 +0,0 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import traceback
|
||||
from dotenv import load_dotenv
|
||||
from fastapi import Request
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from litellm import Router
|
||||
import pytest
|
||||
import litellm
|
||||
from unittest.mock import patch, MagicMock, AsyncMock
|
||||
from litellm.types.utils import ModelResponse, StandardLoggingPayload
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute
|
||||
from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo
|
||||
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS, ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def model_list():
|
||||
return [
|
||||
{
|
||||
"model_name": "gpt-5-mini",
|
||||
"litellm_params": {
|
||||
"model": "gpt-5-mini",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
"tpm": 1000, # Add TPM limit so async method doesn't return early
|
||||
"rpm": 100, # Add RPM limit so async method doesn't return early
|
||||
},
|
||||
"model_info": {
|
||||
"access_groups": ["group1", "group2"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.5",
|
||||
"litellm_params": {
|
||||
"model": "gpt-5.5",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-image-1",
|
||||
"litellm_params": {
|
||||
"model": "gpt-image-1",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "*",
|
||||
"litellm_params": {
|
||||
"model": "openai/*",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "claude-*",
|
||||
"litellm_params": {
|
||||
"model": "anthropic/*",
|
||||
"api_key": os.getenv("ANTHROPIC_API_KEY"),
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def test_validate_fallbacks(model_list):
|
||||
router = Router(model_list=model_list, fallbacks=[{"gpt-5.5": "gpt-5-mini"}])
|
||||
router.validate_fallbacks(fallback_param=[{"gpt-5.5": "gpt-5-mini"}])
|
||||
|
||||
|
||||
def test_routing_strategy_init(model_list):
|
||||
"""Test if all routing strategies are initialized correctly"""
|
||||
from litellm.types.router import RoutingStrategy
|
||||
|
||||
router = Router(model_list=model_list)
|
||||
for strategy in RoutingStrategy:
|
||||
router.routing_strategy_init(
|
||||
routing_strategy=strategy, routing_strategy_args={}
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
def test_routing_strategy_init_valid_string_strategies(model_list):
|
||||
"""Test that all valid string routing strategies work without error.
|
||||
|
||||
Valid strategies are derived from RoutingStrategy enum values plus 'simple-shuffle'.
|
||||
"""
|
||||
from litellm.types.router import RoutingStrategy
|
||||
|
||||
router = Router(model_list=model_list)
|
||||
|
||||
# All strategies from enum + simple-shuffle (default, not in enum)
|
||||
valid_strategies = ["simple-shuffle"] + [s.value for s in RoutingStrategy]
|
||||
|
||||
for strategy in valid_strategies:
|
||||
# Should not raise
|
||||
router.routing_strategy_init(
|
||||
routing_strategy=strategy, routing_strategy_args={}
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.flaky(retries=6, delay=1)
|
||||
@pytest.mark.asyncio
|
||||
async def test_image_generation(model_list, sync_mode):
|
||||
"""Test if the underlying '_image_generation' function is working correctly"""
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
router = Router(model_list=model_list)
|
||||
if sync_mode:
|
||||
response = router._image_generation(
|
||||
model="gpt-image-1",
|
||||
prompt="A cute baby sea otter",
|
||||
)
|
||||
else:
|
||||
response = await router._aimage_generation(
|
||||
model="gpt-image-1",
|
||||
prompt="A cute baby sea otter",
|
||||
)
|
||||
|
||||
ImageResponse.model_validate(response)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def _rpm_tpm_router(model_id: str) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-5-mini",
|
||||
"litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake", "tpm": 1000, "rpm": 100},
|
||||
"model_info": {"id": model_id},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def router_minute_pinned(monkeypatch):
|
||||
pinned = datetime(2026, 1, 1, 12, 0, 30, tzinfo=timezone.utc)
|
||||
monkeypatch.setattr("litellm.router.get_utc_datetime", lambda: pinned)
|
||||
|
||||
|
||||
def _ratelimit_headers(response: ModelResponse | CustomStreamWrapper) -> dict[str, int]:
|
||||
return {k: v for k, v in response._hidden_params["additional_headers"].items() if k.startswith("x-ratelimit-")}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("router_minute_pinned")
|
||||
async def test_acompletion_headers_read_post_increment_counter_and_count_once():
|
||||
router = _rpm_tpm_router("lit-3058-async")
|
||||
|
||||
response = await router.acompletion(
|
||||
model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong"
|
||||
)
|
||||
total_tokens = response.usage.total_tokens
|
||||
assert total_tokens > 0
|
||||
|
||||
headers = _ratelimit_headers(response)
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 1000 - total_tokens
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("router_minute_pinned")
|
||||
async def test_acompletion_stream_counts_request_before_headers_and_tokens_once_on_completion():
|
||||
router = _rpm_tpm_router("lit-3058-stream")
|
||||
|
||||
stream = await router.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="pong pong pong",
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
headers = _ratelimit_headers(stream)
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 1000
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (0, 1)
|
||||
|
||||
chunks = [chunk async for chunk in stream]
|
||||
total_tokens = chunks[-1].usage.total_tokens
|
||||
assert total_tokens > 0
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
|
||||
|
||||
|
||||
|
||||
|
||||
class _GatedIncrementCache(DualCache):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(in_memory_cache=InMemoryCache())
|
||||
self.first_increment_started = asyncio.Event()
|
||||
self.release_first_increment = asyncio.Event()
|
||||
self.increment_calls = 0
|
||||
|
||||
async def async_increment_cache_pipeline(
|
||||
self,
|
||||
increment_list: list[RedisPipelineIncrementOperation],
|
||||
local_only: bool = False,
|
||||
parent_otel_span: object = None,
|
||||
**kwargs: object,
|
||||
) -> list[float] | None:
|
||||
self.increment_calls += 1
|
||||
if self.increment_calls == 1:
|
||||
self.first_increment_started.set()
|
||||
await self.release_first_increment.wait()
|
||||
return await super().async_increment_cache_pipeline(
|
||||
increment_list, local_only=local_only, parent_otel_span=parent_otel_span, **kwargs
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_callback_running_during_pre_header_increment_does_not_double_count():
|
||||
router = _rpm_tpm_router("lit-3058-race")
|
||||
cache = _GatedIncrementCache()
|
||||
router.cache = cache
|
||||
|
||||
request = asyncio.ensure_future(
|
||||
router.acompletion(model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong")
|
||||
)
|
||||
await asyncio.wait_for(cache.first_increment_started.wait(), timeout=5)
|
||||
for _ in range(50):
|
||||
if get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1:
|
||||
break
|
||||
await asyncio.sleep(0.1)
|
||||
assert get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1
|
||||
assert cache.increment_calls == 1
|
||||
|
||||
cache.release_first_increment.set()
|
||||
response = await request
|
||||
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (response.usage.total_tokens, 1)
|
||||
|
||||
|
||||
class _UnavailableIncrementCache(DualCache):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(in_memory_cache=InMemoryCache())
|
||||
self.first_increment_started = asyncio.Event()
|
||||
self.release_first_increment = asyncio.Event()
|
||||
self.increment_calls = 0
|
||||
|
||||
async def async_increment_cache_pipeline(
|
||||
self,
|
||||
increment_list: list[RedisPipelineIncrementOperation],
|
||||
local_only: bool = False,
|
||||
parent_otel_span: object = None,
|
||||
**kwargs: object,
|
||||
) -> list[float] | None:
|
||||
self.increment_calls += 1
|
||||
if self.increment_calls == 1:
|
||||
self.first_increment_started.set()
|
||||
await self.release_first_increment.wait()
|
||||
raise RuntimeError("cache unavailable")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_observing_stamp_before_pre_header_increment_fails_leaves_no_stamp_behind():
|
||||
router = _rpm_tpm_router("lit-3058-fail")
|
||||
cache = _UnavailableIncrementCache()
|
||||
router.cache = cache
|
||||
metadata: dict[str, object] = {}
|
||||
|
||||
request = asyncio.ensure_future(
|
||||
router.acompletion(
|
||||
model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong", metadata=metadata
|
||||
)
|
||||
)
|
||||
await asyncio.wait_for(cache.first_increment_started.wait(), timeout=5)
|
||||
assert metadata[ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY] == 30
|
||||
for _ in range(50):
|
||||
if get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1:
|
||||
break
|
||||
await asyncio.sleep(0.1)
|
||||
assert get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1
|
||||
assert cache.increment_calls == 1
|
||||
|
||||
cache.release_first_increment.set()
|
||||
response = await request
|
||||
|
||||
assert response.usage.total_tokens == 30
|
||||
assert ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY not in metadata
|
||||
assert _ratelimit_headers(response)["x-ratelimit-remaining-requests"] == 100
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (None, None)
|
||||
|
||||
|
||||
def test_track_deployment_metrics(model_list):
|
||||
"""Test if the 'track_deployment_metrics' function is working correctly"""
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
router = Router(model_list=model_list)
|
||||
router._track_deployment_metrics(
|
||||
deployment=router.get_deployment_by_model_group_name(
|
||||
model_group_name="gpt-5-mini"
|
||||
),
|
||||
response=ModelResponse(
|
||||
model="gpt-5-mini",
|
||||
usage={"total_tokens": 100},
|
||||
),
|
||||
parent_otel_span=None,
|
||||
)
|
||||
|
||||
|
||||
def test_pass_through_assistants_endpoint_factory(model_list):
|
||||
"""Test if the 'pass_through_assistants_endpoint_factory' function is working correctly"""
|
||||
router = Router(model_list=model_list)
|
||||
router._pass_through_assistants_endpoint_factory(
|
||||
original_function=litellm.acreate_assistants,
|
||||
custom_llm_provider="openai",
|
||||
client=None,
|
||||
**{},
|
||||
)
|
||||
|
||||
|
||||
def test_factory_function(model_list):
|
||||
"""Test if the 'factory_function' function is working correctly"""
|
||||
router = Router(model_list=model_list)
|
||||
router.factory_function(litellm.acreate_assistants)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# def test_pattern_match_deployments(model_list):
|
||||
# from litellm.router_utils.pattern_match_deployments import PatternMatchRouter
|
||||
# import re
|
||||
|
||||
# patter_router = PatternMatchRouter()
|
||||
|
||||
# request = "fo::hi::static::hello"
|
||||
# model_name = "fo::*:static::*"
|
||||
|
||||
# model_name_regex = patter_router._pattern_to_regex(model_name)
|
||||
|
||||
# # Match against the request
|
||||
# match = re.match(model_name_regex, request)
|
||||
|
||||
# print(f"match: {match}")
|
||||
# print(f"match.end: {match.end()}")
|
||||
# if match is None:
|
||||
# raise ValueError("Match not found")
|
||||
# updated_model = patter_router.set_deployment_model_name(
|
||||
# matched_pattern=match, litellm_deployment_litellm_model="openai/*"
|
||||
# )
|
||||
# assert updated_model == "openai/fo::hi:static::hello"
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_moderation_endpoint_factory(model_list):
|
||||
router = Router(model_list=model_list)
|
||||
response = await router._pass_through_moderation_endpoint_factory(
|
||||
original_function=litellm.amoderation,
|
||||
input="this is valid good text",
|
||||
model=None,
|
||||
)
|
||||
assert response is not None
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_handle_clientside_credential_no_metadata(model_list):
|
||||
"""Test that _handle_clientside_credential handles cases where no metadata is provided"""
|
||||
router = Router(model_list=model_list)
|
||||
|
||||
# Mock deployment
|
||||
deployment = {
|
||||
"model_name": "gpt-4.1",
|
||||
"litellm_params": {"model": "gpt-4.1", "api_key": "test_key"},
|
||||
"model_info": {"id": "original-id-789"},
|
||||
}
|
||||
|
||||
# Mock kwargs with clientside credentials but NO metadata
|
||||
kwargs = {
|
||||
"api_key": "client_side_key",
|
||||
"api_base": "https://api.openai.com/v1",
|
||||
# No metadata key at all
|
||||
}
|
||||
|
||||
# This should fail because there's no model_group in metadata
|
||||
# The function expects to find model_group in the metadata
|
||||
try:
|
||||
result_deployment = router._handle_clientside_credential(
|
||||
deployment=deployment, kwargs=kwargs, function_name="acompletion"
|
||||
)
|
||||
# If we get here, the function should have used deployment.model_name as fallback
|
||||
assert result_deployment.model_name == "gpt-4.1"
|
||||
print("✓ Success with no metadata - used deployment.model_name as fallback")
|
||||
except Exception as e:
|
||||
# This is expected behavior - the function needs model_group to generate model_id
|
||||
print(f"✓ Correctly handled no metadata case: {e}")
|
||||
|
||||
# Test with empty metadata
|
||||
kwargs_with_empty_metadata = {
|
||||
"api_key": "client_side_key",
|
||||
"api_base": "https://api.openai.com/v1",
|
||||
"metadata": {}, # Empty metadata
|
||||
}
|
||||
|
||||
try:
|
||||
result_deployment = router._handle_clientside_credential(
|
||||
deployment=deployment,
|
||||
kwargs=kwargs_with_empty_metadata,
|
||||
function_name="acompletion",
|
||||
)
|
||||
# Should fail because empty metadata has no model_group
|
||||
pytest.fail("Expected failure with empty metadata")
|
||||
except Exception as e:
|
||||
print(f"✓ Correctly handled empty metadata case: {e}")
|
||||
|
|
@ -2197,6 +2197,58 @@ def test_log_event_returns_the_v2_dict_shape_for_the_alerting_trace_id_cache():
|
|||
assert returned["generation_id"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("metadata", "expected_source"),
|
||||
[
|
||||
({}, None),
|
||||
({"trace_id": "my-unique-trace-id"}, "my-unique-trace-id"),
|
||||
({"existing_trace_id": "my-unique-existing-trace-id"}, "my-unique-existing-trace-id"),
|
||||
(
|
||||
{"trace_id": "my-unique-trace-id", "existing_trace_id": "my-unique-existing-trace-id"},
|
||||
"my-unique-existing-trace-id",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_logging_get_trace_id_reports_the_langfuse_trace_that_won_precedence(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
metadata: dict[str, str],
|
||||
expected_source: str | None,
|
||||
) -> None:
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
logger, exporter = _steering_logger()
|
||||
monkeypatch.setattr(litellm_logging, "langFuseLogger", logger)
|
||||
monkeypatch.setattr(litellm, "success_callback", ["langfuse"])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
call_id: Final = f"trace-precedence-{len(metadata)}-{expected_source}"
|
||||
logging_obj: Final = Logging(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
litellm_call_id=call_id,
|
||||
start_time=datetime.datetime.now(),
|
||||
function_id="trace-precedence",
|
||||
)
|
||||
|
||||
litellm.completion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "trace precedence"}],
|
||||
mock_response="ok",
|
||||
litellm_logging_obj=logging_obj,
|
||||
metadata=dict(metadata),
|
||||
)
|
||||
deadline: Final = time.monotonic() + 5
|
||||
while logging_obj.get_trace_id(service_name="langfuse") is None and time.monotonic() < deadline:
|
||||
time.sleep(0.01)
|
||||
|
||||
expected_trace_id: Final = resolve_trace_id(expected_source or logging_obj.litellm_trace_id)
|
||||
assert logging_obj.get_trace_id(service_name="langfuse") == expected_trace_id
|
||||
assert _span_trace_id(_exported_span(logger, exporter)) == expected_trace_id
|
||||
|
||||
|
||||
def test_parse_langfuse_debug_only_enables_on_true_strings():
|
||||
"""v4 treats any truthy value as debug=on, so the raw env string "false" would enable debug."""
|
||||
assert langfuse_module.parse_langfuse_debug("true") is True
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
"""Test health check helper functions"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
import struct
|
||||
import zlib
|
||||
|
|
@ -1235,3 +1234,193 @@ async def test_health_check_with_custom_llm_provider(
|
|||
assert "error" not in response, response
|
||||
assert upstream.called
|
||||
assert json.loads(upstream.calls[0].request.content)["model"] == "deepseek-r1-distill-qwen-1.5B-q4"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_chat_health_check_surfaces_provider_rate_limit_headers(
|
||||
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
upstream: Final = respx_mock.post(
|
||||
url__regex=r"https://resource\.example/openai/deployments/gpt-4\.1-mini/chat/completions.*"
|
||||
).respond(
|
||||
json={
|
||||
"id": "chatcmpl-health",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-4.1-mini",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
},
|
||||
headers={"x-ratelimit-remaining-tokens": "42"},
|
||||
)
|
||||
|
||||
response: Final = await ahealth_check(
|
||||
{
|
||||
"model": "azure/gpt-4.1-mini",
|
||||
"api_key": "fake-key",
|
||||
"api_base": "https://resource.example",
|
||||
"api_version": "2024-06-01",
|
||||
},
|
||||
mode="chat",
|
||||
)
|
||||
|
||||
assert response["x-ratelimit-remaining-tokens"] == "42"
|
||||
assert upstream.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_embedding_health_check_surfaces_provider_rate_limit_headers(
|
||||
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
upstream: Final = respx_mock.post(
|
||||
url__regex=r"https://resource\.example/openai/deployments/text-embedding-ada-002/embeddings.*"
|
||||
).respond(
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}],
|
||||
"model": "text-embedding-ada-002",
|
||||
"usage": {"prompt_tokens": 1, "total_tokens": 1},
|
||||
},
|
||||
headers={"x-ratelimit-remaining-tokens": "84"},
|
||||
)
|
||||
|
||||
response: Final = await ahealth_check(
|
||||
{
|
||||
"model": "azure/text-embedding-ada-002",
|
||||
"api_key": "fake-key",
|
||||
"api_base": "https://resource.example",
|
||||
"api_version": "2024-06-01",
|
||||
},
|
||||
input=["health check"],
|
||||
mode="embedding",
|
||||
)
|
||||
|
||||
assert response["x-ratelimit-remaining-tokens"] == "84"
|
||||
assert upstream.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_image_generation_health_check_returns_a_successful_response(
|
||||
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
upstream: Final = respx_mock.post("https://api.openai.com/v1/images/generations").respond(
|
||||
json={"created": 1, "data": [{"b64_json": "AA=="}]}
|
||||
)
|
||||
|
||||
response: Final = await ahealth_check(
|
||||
{"model": "gpt-image-1", "api_key": "fake-key"},
|
||||
mode="image_generation",
|
||||
prompt="health check",
|
||||
)
|
||||
|
||||
assert "error" not in response
|
||||
assert upstream.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_groq_wildcard_health_check_uses_a_concrete_model(
|
||||
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"models_by_provider",
|
||||
{"groq": ["groq/openai/gpt-oss-20b"]},
|
||||
)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
upstream: Final = respx_mock.post("https://api.groq.com/openai/v1/chat/completions").respond(
|
||||
json={
|
||||
"id": "chatcmpl-health",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "groq/openai/gpt-oss-20b",
|
||||
"service_tier": "on_demand",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "2"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
}
|
||||
)
|
||||
|
||||
response: Final = await ahealth_check(
|
||||
{
|
||||
"model": "groq/*",
|
||||
"api_key": "fake-key",
|
||||
"messages": [{"role": "user", "content": "What is 1 + 1?"}],
|
||||
}
|
||||
)
|
||||
|
||||
assert upstream.called
|
||||
assert json.loads(upstream.calls.last.request.content)["model"] == "openai/gpt-oss-20b"
|
||||
assert response == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cohere_rerank_health_check_returns_a_successful_response(
|
||||
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
upstream: Final = respx_mock.post("https://api.cohere.com/v2/rerank").respond(
|
||||
json={
|
||||
"id": "rerank-health",
|
||||
"results": [{"index": 0, "relevance_score": 0.7}],
|
||||
"meta": {"billed_units": {"search_units": 1}},
|
||||
}
|
||||
)
|
||||
|
||||
response: Final = await ahealth_check(
|
||||
{"model": "cohere/rerank-english-v3.0", "api_key": "fake-key"},
|
||||
mode="rerank",
|
||||
prompt="health check",
|
||||
)
|
||||
|
||||
assert "error" not in response
|
||||
assert upstream.called
|
||||
assert json.loads(upstream.calls.last.request.content)["query"] == "health check"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_speech_health_check_returns_audio(
|
||||
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
upstream: Final = respx_mock.post("https://api.openai.com/v1/audio/speech").respond(
|
||||
content=b"audio",
|
||||
headers={"content-type": "audio/mpeg"},
|
||||
)
|
||||
|
||||
response: Final = await ahealth_check(
|
||||
{"model": "openai/tts-1", "api_key": "fake-key"},
|
||||
mode="audio_speech",
|
||||
prompt="health check",
|
||||
)
|
||||
|
||||
assert "error" not in response
|
||||
assert upstream.called
|
||||
assert json.loads(upstream.calls.last.request.content)["input"] == "health check"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_transcription_health_check_returns_transcribed_text(
|
||||
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
upstream: Final = respx_mock.post("https://api.openai.com/v1/audio/transcriptions").respond(
|
||||
json={"text": "health check audio"}
|
||||
)
|
||||
|
||||
response: Final = await ahealth_check(
|
||||
{"model": "openai/whisper-1", "api_key": "fake-key"},
|
||||
mode="audio_transcription",
|
||||
)
|
||||
|
||||
assert "error" not in response
|
||||
assert upstream.called
|
||||
assert b'name="file"' in upstream.calls.last.request.content
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ Regression tests for https://github.com/BerriAI/litellm/issues/22040 and for the
|
|||
|
||||
import os
|
||||
import sys
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -21,10 +22,15 @@ sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../.
|
|||
import litellm
|
||||
from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
from litellm.llms.anthropic.count_tokens.token_counter import AnthropicTokenCounter
|
||||
from litellm.llms.anthropic.count_tokens.transformation import (
|
||||
AnthropicCountTokensConfig,
|
||||
)
|
||||
from litellm.llms.azure_ai.anthropic.count_tokens.token_counter import (
|
||||
AzureAIAnthropicTokenCounter,
|
||||
)
|
||||
from litellm.types.llms.anthropic import ANTHROPIC_OAUTH_BETA_HEADER
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
# Fake tokens for testing (not real secrets)
|
||||
FAKE_OAUTH_TOKEN = "sk-ant-oat01-fake-token-for-testing-123456789abcdef"
|
||||
|
|
@ -269,3 +275,80 @@ class TestCountTokensUsesWorkloadIdentity:
|
|||
assert result is not None
|
||||
assert result.total_tokens == 7
|
||||
assert seen["auth_header"] == {"x-api-key": vault_key}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("counter_type", "api_base", "endpoint", "tokenizer_type"),
|
||||
(
|
||||
(
|
||||
AnthropicTokenCounter,
|
||||
"https://gateway.example",
|
||||
"https://gateway.example/v1/messages/count_tokens",
|
||||
"anthropic_api",
|
||||
),
|
||||
(
|
||||
AzureAIAnthropicTokenCounter,
|
||||
"https://resource.example",
|
||||
"https://resource.example/anthropic/v1/messages/count_tokens",
|
||||
"azure_ai_anthropic_api",
|
||||
),
|
||||
),
|
||||
)
|
||||
@pytest.mark.parametrize(("status_code", "expected_error"), ((200, False), (401, True)))
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_token_counters_return_typed_success_and_error_responses(
|
||||
counter_type: type[AnthropicTokenCounter] | type[AzureAIAnthropicTokenCounter],
|
||||
api_base: str,
|
||||
endpoint: str,
|
||||
tokenizer_type: str,
|
||||
status_code: int,
|
||||
expected_error: bool,
|
||||
httpx_transport_clients: None,
|
||||
) -> None:
|
||||
response_body: Final = {"input_tokens": 17} if status_code == 200 else {"error": {"message": "invalid key"}}
|
||||
router: Final = respx.mock
|
||||
|
||||
with router:
|
||||
count_route: Final = router.post(endpoint).mock(return_value=httpx.Response(status_code, json=response_body))
|
||||
result: Final = await counter_type().count_tokens(
|
||||
model_to_use="claude-test",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"api_key": "sk-ant-api03-test-key", "api_base": api_base}},
|
||||
request_model="claude-test",
|
||||
)
|
||||
requests: Final = tuple(router.calls)
|
||||
|
||||
assert count_route.called
|
||||
assert len(requests) == 1
|
||||
assert requests[0].request.url == httpx.URL(endpoint)
|
||||
assert isinstance(result, TokenCountResponse)
|
||||
assert result.request_model == "claude-test"
|
||||
assert result.model_used == "claude-test"
|
||||
assert result.tokenizer_type == tokenizer_type
|
||||
|
||||
if expected_error:
|
||||
assert result.error is True
|
||||
assert result.status_code == status_code
|
||||
assert result.total_tokens == 0
|
||||
return
|
||||
|
||||
assert result.error is not True
|
||||
assert result.total_tokens == 17
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("counter_type", "provider"),
|
||||
(
|
||||
(AnthropicTokenCounter, "anthropic"),
|
||||
(AzureAIAnthropicTokenCounter, "azure_ai"),
|
||||
),
|
||||
)
|
||||
def test_count_token_counters_select_their_own_provider(
|
||||
counter_type: type[AnthropicTokenCounter] | type[AzureAIAnthropicTokenCounter],
|
||||
provider: str,
|
||||
) -> None:
|
||||
counter: Final = counter_type()
|
||||
|
||||
assert counter.should_use_token_counting_api(custom_llm_provider=provider) is True
|
||||
assert counter.should_use_token_counting_api(custom_llm_provider="unknown") is False
|
||||
|
|
|
|||
531
tests/unit/llms/openai/responses/test_openai_responses_http.py
Normal file
531
tests/unit/llms/openai/responses/test_openai_responses_http.py
Normal file
|
|
@ -0,0 +1,531 @@
|
|||
import asyncio
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import Iterable
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.llms.openai import (
|
||||
IncompleteDetails,
|
||||
ResponseAPIUsage,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
)
|
||||
from litellm.types.utils import StandardLoggingPayload, Usage
|
||||
from tests.unit.proxy.conftest import httpx_transport
|
||||
|
||||
pytestmark: Final = pytest.mark.usefixtures(httpx_transport.__name__)
|
||||
_OPENAI_URL: Final = "https://api.openai.com/v1/responses"
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
|
||||
_STANDARD_LOGGING_PAYLOAD: Final = TypeAdapter(StandardLoggingPayload)
|
||||
_OUTPUT_TEXT: Final = "Hello from the mocked response"
|
||||
_CREATED_AT: Final = 1750000000
|
||||
_AsyncLoggingMode: TypeAlias = Literal["non_stream", "stream"]
|
||||
_RESPONSE_FIELD_TYPES: Final = MappingProxyType(
|
||||
{
|
||||
"error": (dict, type(None)),
|
||||
"incomplete_details": (IncompleteDetails, type(None)),
|
||||
"instructions": (str, type(None)),
|
||||
"metadata": (dict,),
|
||||
"model": (str,),
|
||||
"object": (str,),
|
||||
"parallel_tool_calls": (bool, type(None)),
|
||||
"temperature": (int, float, type(None)),
|
||||
"tool_choice": (dict, str, type(None)),
|
||||
"tools": (list, type(None)),
|
||||
"top_p": (int, float, type(None)),
|
||||
"max_output_tokens": (int, type(None)),
|
||||
"previous_response_id": (str, type(None)),
|
||||
"reasoning": (dict, type(None)),
|
||||
"status": (str,),
|
||||
"text": (dict,),
|
||||
"truncation": (str, type(None)),
|
||||
"user": (str, type(None)),
|
||||
"store": (bool, type(None)),
|
||||
}
|
||||
)
|
||||
_STREAM_EVENT_FIELDS: Final = MappingProxyType(
|
||||
{
|
||||
"response.created": ("response",),
|
||||
"response.in_progress": ("response",),
|
||||
"response.output_item.added": ("output_index", "item"),
|
||||
"response.content_part.added": ("item_id", "output_index", "content_index", "part"),
|
||||
"response.output_text.delta": ("item_id", "output_index", "content_index", "delta"),
|
||||
"response.output_text.done": ("item_id", "output_index", "content_index", "text"),
|
||||
"response.content_part.done": ("item_id", "output_index", "content_index", "part"),
|
||||
"response.output_item.done": ("output_index", "item"),
|
||||
"response.completed": ("response",),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class _ResponsesLoggingCapture(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.completed: Final = threading.Event()
|
||||
self.payload: StandardLoggingPayload | None = None
|
||||
self.usage: object = None
|
||||
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: dict[str, object],
|
||||
response_obj: object,
|
||||
start_time: object,
|
||||
end_time: object,
|
||||
) -> None:
|
||||
assert isinstance(response_obj, ResponsesAPIResponse)
|
||||
self.payload = _STANDARD_LOGGING_PAYLOAD.validate_python(kwargs["standard_logging_object"])
|
||||
self.usage = response_obj.usage
|
||||
self.completed.set()
|
||||
|
||||
|
||||
def _response_body(response_id: str, store: bool = False, created_at: float = _CREATED_AT) -> dict[str, object]:
|
||||
return {
|
||||
"id": response_id,
|
||||
"object": "response",
|
||||
"created_at": created_at,
|
||||
"status": "completed",
|
||||
"error": None,
|
||||
"incomplete_details": None,
|
||||
"instructions": None,
|
||||
"max_output_tokens": None,
|
||||
"model": "gpt-4o",
|
||||
"output": [
|
||||
{
|
||||
"id": f"msg_{response_id}",
|
||||
"type": "message",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": _OUTPUT_TEXT, "annotations": []}],
|
||||
}
|
||||
],
|
||||
"parallel_tool_calls": True,
|
||||
"previous_response_id": None,
|
||||
"reasoning": {"effort": None, "summary": None},
|
||||
"store": store,
|
||||
"temperature": 1.0,
|
||||
"text": {"format": {"type": "text"}},
|
||||
"tool_choice": "auto",
|
||||
"tools": [],
|
||||
"top_p": 1.0,
|
||||
"truncation": "disabled",
|
||||
"usage": {
|
||||
"input_tokens": 2,
|
||||
"input_tokens_details": {"cached_tokens": 0},
|
||||
"output_tokens": 3,
|
||||
"output_tokens_details": {"reasoning_tokens": 0},
|
||||
"total_tokens": 5,
|
||||
},
|
||||
"user": None,
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
|
||||
def _sse(events: Iterable[dict[str, object]]) -> str:
|
||||
return "".join(f"data: {json.dumps(event)}\n\n" for event in events)
|
||||
|
||||
|
||||
def _response_events(response_id: str) -> tuple[dict[str, object], ...]:
|
||||
completed_response: Final = _response_body(response_id)
|
||||
in_progress_response: Final = {**completed_response, "status": "in_progress", "output": [], "usage": None}
|
||||
item_id: Final = f"msg_{response_id}"
|
||||
text_part: Final = {"type": "output_text", "text": _OUTPUT_TEXT, "annotations": []}
|
||||
position: Final = {"item_id": item_id, "output_index": 0, "content_index": 0}
|
||||
return (
|
||||
{"type": "response.created", "sequence_number": 0, "response": in_progress_response},
|
||||
{"type": "response.in_progress", "sequence_number": 1, "response": in_progress_response},
|
||||
{
|
||||
"type": "response.output_item.added",
|
||||
"sequence_number": 2,
|
||||
"output_index": 0,
|
||||
"item": {"id": item_id, "type": "message", "status": "in_progress", "role": "assistant", "content": []},
|
||||
},
|
||||
{
|
||||
"type": "response.content_part.added",
|
||||
"sequence_number": 3,
|
||||
**position,
|
||||
"part": {"type": "output_text", "text": "", "annotations": []},
|
||||
},
|
||||
{"type": "response.output_text.delta", "sequence_number": 4, **position, "delta": _OUTPUT_TEXT},
|
||||
{"type": "response.output_text.done", "sequence_number": 5, **position, "text": _OUTPUT_TEXT},
|
||||
{"type": "response.content_part.done", "sequence_number": 6, **position, "part": text_part},
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"sequence_number": 7,
|
||||
"output_index": 0,
|
||||
"item": completed_response["output"][0],
|
||||
},
|
||||
{"type": "response.completed", "sequence_number": 8, "response": completed_response},
|
||||
)
|
||||
|
||||
|
||||
def _response_sse(response_id: str) -> str:
|
||||
return _sse(_response_events(response_id))
|
||||
|
||||
|
||||
def _sse_reply(body: str) -> httpx.Response:
|
||||
return httpx.Response(status_code=200, content=body, headers={"content-type": "text/event-stream"})
|
||||
|
||||
|
||||
def _assert_valid_response(response: object, final_chunk: bool) -> None:
|
||||
assert isinstance(response, ResponsesAPIResponse)
|
||||
assert isinstance(response.id, str)
|
||||
assert isinstance(response.created_at, int)
|
||||
assert isinstance(response.usage, ResponseAPIUsage if final_chunk else type(None))
|
||||
mistyped: Final = {
|
||||
field: type(response[field]).__name__
|
||||
for field, expected in _RESPONSE_FIELD_TYPES.items()
|
||||
if not isinstance(response[field], expected)
|
||||
}
|
||||
assert mistyped == {}
|
||||
if final_chunk and response.status == "completed":
|
||||
assert len(response.output) > 0
|
||||
|
||||
|
||||
def _json_object(value: object) -> dict[str, object]:
|
||||
return _JSON_OBJECT.validate_python(value.model_dump(mode="json") if isinstance(value, BaseModel) else value)
|
||||
|
||||
|
||||
def _assert_valid_stream(events: tuple[object, ...], item_id: str) -> None:
|
||||
event_types: Final = tuple(getattr(event, "type", None) for event in events)
|
||||
assert event_types == tuple(_STREAM_EVENT_FIELDS)
|
||||
missing_fields: Final = {
|
||||
event_type: tuple(name for name in fields if getattr(event, name, None) is None)
|
||||
for (event_type, fields), event in zip(_STREAM_EVENT_FIELDS.items(), events)
|
||||
}
|
||||
assert all(fields == () for fields in missing_fields.values()), missing_fields
|
||||
created: Final = getattr(events[0], "response", None)
|
||||
_assert_valid_response(created, final_chunk=False)
|
||||
_assert_valid_response(getattr(events[1], "response", None), final_chunk=False)
|
||||
completed: Final = events[-1]
|
||||
assert isinstance(completed, ResponseCompletedEvent)
|
||||
_assert_valid_response(completed.response, final_chunk=True)
|
||||
assert completed.response.id == getattr(created, "id", None)
|
||||
assert completed.response.output_text == _OUTPUT_TEXT
|
||||
assert tuple(getattr(event, "sequence_number", None) for event in events) == tuple(range(len(events)))
|
||||
item_ids: Final = (
|
||||
getattr(getattr(events[2], "item", None), "id", None),
|
||||
*(getattr(event, "item_id", None) for event in events[3:7]),
|
||||
getattr(getattr(events[7], "item", None), "id", None),
|
||||
)
|
||||
assert item_ids == (item_id,) * 6
|
||||
assert tuple(getattr(event, "output_index", None) for event in events[2:8]) == (0,) * 6
|
||||
assert tuple(getattr(event, "content_index", None) for event in events[3:7]) == (0,) * 4
|
||||
assert _json_object(getattr(events[3], "part", None)) == {
|
||||
"type": "output_text",
|
||||
"text": "",
|
||||
"annotations": [],
|
||||
}
|
||||
assert getattr(events[4], "delta", None) == _OUTPUT_TEXT
|
||||
assert getattr(events[5], "text", None) == _OUTPUT_TEXT
|
||||
assert _json_object(getattr(events[6], "part", None)) == {
|
||||
"type": "output_text",
|
||||
"text": _OUTPUT_TEXT,
|
||||
"annotations": [],
|
||||
}
|
||||
assert getattr(getattr(events[7], "item", None), "type", None) == "message"
|
||||
assert tuple(item.id for item in completed.response.output) == (item_id,)
|
||||
|
||||
|
||||
def _completed_response(events: tuple[object, ...]) -> ResponsesAPIResponse:
|
||||
completed_events: Final = tuple(event for event in events if isinstance(event, ResponseCompletedEvent))
|
||||
assert len(completed_events) == 1
|
||||
return completed_events[0].response
|
||||
|
||||
|
||||
def _install_capture(monkeypatch: pytest.MonkeyPatch) -> _ResponsesLoggingCapture:
|
||||
capture: Final = _ResponsesLoggingCapture()
|
||||
monkeypatch.setattr(litellm, "callbacks", [capture])
|
||||
monkeypatch.setattr(litellm, "success_callback", [capture])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [capture])
|
||||
return capture
|
||||
|
||||
|
||||
def _assert_logged_payload_matches(capture: _ResponsesLoggingCapture, response: ResponsesAPIResponse) -> None:
|
||||
payload: Final = capture.payload
|
||||
usage: Final = capture.usage
|
||||
assert payload is not None
|
||||
assert response.usage is not None
|
||||
assert payload["prompt_tokens"] == response.usage.input_tokens
|
||||
assert payload["completion_tokens"] == response.usage.output_tokens
|
||||
assert payload["total_tokens"] == response.usage.input_tokens + response.usage.output_tokens
|
||||
assert payload["response_cost"] > 0
|
||||
assert payload["id"] == response.id
|
||||
assert payload["model"] == "gpt-4o-mini"
|
||||
assert payload["messages"] == [{"content": "hi", "role": "user"}]
|
||||
callback_usage: Final = usage.model_dump() if isinstance(usage, Usage) else _JSON_OBJECT.validate_python(usage)
|
||||
assert callback_usage["prompt_tokens"] == response.usage.input_tokens
|
||||
assert callback_usage["completion_tokens"] == response.usage.output_tokens
|
||||
logged_response: Final = _JSON_OBJECT.validate_python(payload["response"])
|
||||
final_response: Final = response.model_dump(mode="json")
|
||||
assert {key: value for key, value in logged_response.items() if key != "usage"} == {
|
||||
key: value for key, value in final_response.items() if key != "usage"
|
||||
}
|
||||
logged_usage: Final = _JSON_OBJECT.validate_python(logged_response["usage"])
|
||||
assert logged_usage["prompt_tokens"] == response.usage.input_tokens
|
||||
assert logged_usage["completion_tokens"] == response.usage.output_tokens
|
||||
assert logged_usage["total_tokens"] == response.usage.total_tokens
|
||||
|
||||
|
||||
async def _async_logged_openai_response(mode: _AsyncLoggingMode) -> ResponsesAPIResponse:
|
||||
match mode:
|
||||
case "stream":
|
||||
response_stream: Final = await litellm.aresponses(
|
||||
model="openai/gpt-4o-mini", api_key="sk-test", input="hi", stream=True
|
||||
)
|
||||
return _completed_response(tuple([event async for event in response_stream]))
|
||||
case "non_stream":
|
||||
response: Final = await litellm.aresponses(model="openai/gpt-4o-mini", api_key="sk-test", input="hi")
|
||||
assert isinstance(response, ResponsesAPIResponse)
|
||||
return response
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", (True, False))
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_exposes_provider_rate_limit_headers(sync_mode: bool) -> None:
|
||||
response_headers: Final = {
|
||||
"x-ratelimit-limit-requests": "500",
|
||||
"x-ratelimit-remaining-requests": "499",
|
||||
"x-ratelimit-reset-requests": "1s",
|
||||
}
|
||||
|
||||
with respx.mock() as mock_router:
|
||||
mock_router.post(_OPENAI_URL).mock(
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json=_response_body("resp_rate_limit"),
|
||||
headers=response_headers,
|
||||
)
|
||||
)
|
||||
response: Final = (
|
||||
litellm.responses(model="openai/gpt-4o", api_key="sk-test", input="hi")
|
||||
if sync_mode
|
||||
else await litellm.aresponses(model="openai/gpt-4o", api_key="sk-test", input="hi")
|
||||
)
|
||||
requests: Final = tuple(mock_router.calls)
|
||||
|
||||
assert isinstance(response, ResponsesAPIResponse)
|
||||
assert len(requests) == 1
|
||||
additional_headers: Final = _JSON_OBJECT.validate_python(response.hidden_params["additional_headers"])
|
||||
raw_headers: Final = _JSON_OBJECT.validate_python(response.hidden_params["headers"])
|
||||
assert {name: additional_headers[f"llm_provider-{name}"] for name in response_headers} == response_headers
|
||||
assert {name: raw_headers[name] for name in response_headers} == response_headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_converts_created_at_to_int_and_forwards_store() -> None:
|
||||
response_body: Final = _response_body("resp_store", store=True, created_at=_CREATED_AT + 0.75)
|
||||
|
||||
with respx.mock() as mock_router:
|
||||
mock_router.post(_OPENAI_URL).mock(return_value=httpx.Response(status_code=200, json=response_body))
|
||||
response: Final = await litellm.aresponses(
|
||||
model="openai/gpt-4o",
|
||||
api_key="sk-test",
|
||||
input="hi",
|
||||
store=True,
|
||||
)
|
||||
requests: Final = tuple(mock_router.calls)
|
||||
|
||||
assert isinstance(response, ResponsesAPIResponse)
|
||||
assert type(response.created_at) is int
|
||||
assert response.created_at == _CREATED_AT
|
||||
assert response.store is True
|
||||
assert len(requests) == 1
|
||||
request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content)
|
||||
assert request_body["store"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_mcp_followup_forwards_approval_and_previous_id() -> None:
|
||||
mcp_tools: Final = [
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_label": "weather",
|
||||
"server_url": "https://mcp.example.test",
|
||||
"headers": {"Authorization": "Bearer mcp-test-token"},
|
||||
}
|
||||
]
|
||||
approval: Final = [{"type": "mcp_approval_response", "approve": True, "approval_request_id": "approval_123"}]
|
||||
|
||||
with respx.mock() as mock_router:
|
||||
route: Final = mock_router.post(_OPENAI_URL).mock(
|
||||
return_value=httpx.Response(status_code=200, json=_response_body("resp_mcp"))
|
||||
)
|
||||
first_response: Final = await litellm.aresponses(
|
||||
model="openai/gpt-4o",
|
||||
api_key="sk-test",
|
||||
input="Search for recent weather",
|
||||
tools=mcp_tools,
|
||||
)
|
||||
assert isinstance(first_response, ResponsesAPIResponse)
|
||||
second_response: Final = await litellm.aresponses(
|
||||
model="openai/gpt-4o",
|
||||
api_key="sk-test",
|
||||
input=approval,
|
||||
tools=mcp_tools,
|
||||
previous_response_id=first_response.id,
|
||||
)
|
||||
requests: Final = tuple(route.calls)
|
||||
|
||||
assert isinstance(second_response, ResponsesAPIResponse)
|
||||
assert len(requests) == 2
|
||||
first_request: Final = _JSON_OBJECT.validate_json(requests[0].request.content)
|
||||
second_request: Final = _JSON_OBJECT.validate_json(requests[1].request.content)
|
||||
assert first_request["tools"] == mcp_tools
|
||||
assert second_request["tools"] == mcp_tools
|
||||
assert second_request["input"] == approval
|
||||
assert second_request["previous_response_id"] == "resp_mcp"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", ("non_stream", "stream"))
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_standard_logging_matches_final_response(
|
||||
mode: _AsyncLoggingMode, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
capture: Final = _install_capture(monkeypatch)
|
||||
|
||||
with respx.mock() as mock_router:
|
||||
mock_router.post(_OPENAI_URL).mock(
|
||||
return_value=(
|
||||
_sse_reply(_response_sse("resp_logging"))
|
||||
if mode == "stream"
|
||||
else httpx.Response(status_code=200, json=_response_body("resp_logging"))
|
||||
)
|
||||
)
|
||||
response: Final = await _async_logged_openai_response(mode)
|
||||
requests: Final = tuple(mock_router.calls)
|
||||
logged: Final = await asyncio.to_thread(capture.completed.wait, 10)
|
||||
|
||||
assert logged
|
||||
assert len(requests) == 1
|
||||
_assert_logged_payload_matches(capture, response)
|
||||
|
||||
|
||||
def test_sync_stream_standard_logging_matches_final_response(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
capture: Final = _install_capture(monkeypatch)
|
||||
|
||||
with respx.mock() as mock_router:
|
||||
mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_response_sse("resp_sync_logging")))
|
||||
stream: Final = litellm.responses(model="openai/gpt-4o-mini", api_key="sk-test", input="hi", stream=True)
|
||||
response: Final = _completed_response(tuple(stream))
|
||||
requests: Final = tuple(mock_router.calls)
|
||||
logged: Final = capture.completed.wait(10)
|
||||
|
||||
assert logged
|
||||
assert len(requests) == 1
|
||||
_assert_logged_payload_matches(capture, response)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", (True, False))
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_stream_emits_valid_events(sync_mode: bool) -> None:
|
||||
with respx.mock() as mock_router:
|
||||
route: Final = mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_response_sse("resp_stream")))
|
||||
events: Final = (
|
||||
tuple(litellm.responses(model="openai/gpt-4o", api_key="sk-test", input="hi", stream=True))
|
||||
if sync_mode
|
||||
else tuple(
|
||||
[
|
||||
event
|
||||
async for event in await litellm.aresponses(
|
||||
model="openai/gpt-4o", api_key="sk-test", input="hi", stream=True
|
||||
)
|
||||
]
|
||||
)
|
||||
)
|
||||
requests: Final = tuple(route.calls)
|
||||
|
||||
assert len(requests) == 1
|
||||
_assert_valid_stream(events, "msg_resp_stream")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_stream_error_event_raises_bad_request() -> None:
|
||||
message: Final = "Your input exceeds the context window of this model."
|
||||
created: Final = _response_events("resp_too_long")[0]
|
||||
error_event: Final = {
|
||||
"type": "error",
|
||||
"sequence_number": 1,
|
||||
"error": {
|
||||
"type": "invalid_request_error",
|
||||
"code": "context_length_exceeded",
|
||||
"message": message,
|
||||
"param": "input",
|
||||
},
|
||||
}
|
||||
|
||||
with respx.mock() as mock_router:
|
||||
route: Final = mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_sse((created, error_event))))
|
||||
stream: Final = await litellm.aresponses(
|
||||
model="openai/gpt-5-mini", api_key="sk-test", input="oversized prompt", stream=True
|
||||
)
|
||||
with pytest.raises(litellm.BadRequestError) as exc_info:
|
||||
_ = [event async for event in stream]
|
||||
requests: Final = tuple(route.calls)
|
||||
|
||||
assert len(requests) == 1
|
||||
assert exc_info.value.status_code == 400
|
||||
assert message in str(exc_info.value)
|
||||
|
||||
|
||||
def _alias_router() -> litellm.Router:
|
||||
return litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "openai-offline-alias",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("sync_mode", (True, False))
|
||||
async def test_router_responses_alias_uses_underlying_model(sync_mode: bool) -> None:
|
||||
router: Final = _alias_router()
|
||||
with respx.mock() as mock_router:
|
||||
route: Final = mock_router.post(_OPENAI_URL).mock(
|
||||
return_value=httpx.Response(status_code=200, json=_response_body("resp_router_alias"))
|
||||
)
|
||||
response: Final = (
|
||||
router.responses(model="openai-offline-alias", input="hi")
|
||||
if sync_mode
|
||||
else await router.aresponses(model="openai-offline-alias", input="hi")
|
||||
)
|
||||
requests: Final = tuple(route.calls)
|
||||
|
||||
_assert_valid_response(response, final_chunk=True)
|
||||
assert len(requests) == 1
|
||||
request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content)
|
||||
assert request_body["model"] == "gpt-4o"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("sync_mode", (True, False))
|
||||
async def test_router_responses_alias_stream_uses_underlying_model(sync_mode: bool) -> None:
|
||||
router: Final = _alias_router()
|
||||
with respx.mock() as mock_router:
|
||||
route: Final = mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_response_sse("resp_router_stream")))
|
||||
events: Final = (
|
||||
tuple(router.responses(model="openai-offline-alias", input="hi", stream=True))
|
||||
if sync_mode
|
||||
else tuple(
|
||||
[
|
||||
event
|
||||
async for event in await router.aresponses(model="openai-offline-alias", input="hi", stream=True)
|
||||
]
|
||||
)
|
||||
)
|
||||
requests: Final = tuple(route.calls)
|
||||
|
||||
_assert_valid_stream(events, "msg_resp_router_stream")
|
||||
assert len(requests) == 1
|
||||
request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content)
|
||||
assert request_body["model"] == "gpt-4o"
|
||||
|
|
@ -6,6 +6,8 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -24,6 +26,7 @@ from litellm.types.llms.openai import (
|
|||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
from tests.unit.proxy.conftest import httpx_transport
|
||||
import time
|
||||
|
||||
_ARTIFACT_FIELD_PATTERN: Final = r'^(?!__.*__$)[^\p{Cc}\p{Cf}\p{Zl}\p{Zp}"\\./[\]]{1,200}$'
|
||||
|
|
@ -3445,3 +3448,127 @@ async def test_extra_body_merges_with_request_data(extra_body_mock_response_data
|
|||
assert "temperature" in request_body
|
||||
assert "custom_field" in request_body
|
||||
assert request_body["custom_field"] == "custom_value"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures(httpx_transport.__name__)
|
||||
async def test_aresponses_forwards_previous_response_id_to_openai() -> None:
|
||||
first_input: Final = "remember the first turn"
|
||||
second_input: Final = "continue the conversation"
|
||||
first_id: Final = "resp_previous_turn"
|
||||
second_id: Final = "resp_follow_up"
|
||||
response_payloads: Final = (
|
||||
{
|
||||
"id": first_id,
|
||||
"object": "response",
|
||||
"created_at": 1734366691,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_first",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "first answer", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"input_tokens": 1,
|
||||
"output_tokens": 1,
|
||||
"total_tokens": 2,
|
||||
"output_tokens_details": {"reasoning_tokens": 0},
|
||||
},
|
||||
},
|
||||
{
|
||||
"id": second_id,
|
||||
"object": "response",
|
||||
"created_at": 1734366692,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_second",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "second answer", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"input_tokens": 1,
|
||||
"output_tokens": 1,
|
||||
"total_tokens": 2,
|
||||
"output_tokens_details": {"reasoning_tokens": 0},
|
||||
},
|
||||
},
|
||||
)
|
||||
provider_url: Final = "https://api.openai.com/v1/responses"
|
||||
|
||||
with respx.mock() as router:
|
||||
route: Final = router.post(provider_url).mock(
|
||||
side_effect=[
|
||||
httpx.Response(status_code=200, json=response_payloads[0]),
|
||||
httpx.Response(status_code=200, json=response_payloads[1]),
|
||||
]
|
||||
)
|
||||
first_response: Final = await litellm.aresponses(
|
||||
model="openai/gpt-4o",
|
||||
api_key="sk-test",
|
||||
input=first_input,
|
||||
)
|
||||
assert isinstance(first_response, ResponsesAPIResponse)
|
||||
second_response: Final = await litellm.aresponses(
|
||||
model="openai/gpt-4o",
|
||||
api_key="sk-test",
|
||||
input=second_input,
|
||||
previous_response_id=first_response.id,
|
||||
)
|
||||
calls: Final = route.calls
|
||||
|
||||
assert isinstance(second_response, ResponsesAPIResponse)
|
||||
assert first_response.output[0].content[0].text == "first answer"
|
||||
assert second_response.output[0].content[0].text == "second answer"
|
||||
assert len(calls) == 2
|
||||
request_adapter: Final = TypeAdapter(dict[str, JsonValue])
|
||||
request_bodies: Final = tuple(request_adapter.validate_json(call.request.content) for call in calls)
|
||||
assert request_bodies[0]["input"] == first_input
|
||||
assert request_bodies[1]["input"] == second_input
|
||||
assert request_bodies[1]["previous_response_id"] == first_id
|
||||
|
||||
|
||||
def test_dict_responses_input_filters_unset_reasoning_fields() -> None:
|
||||
test_input: Final = [
|
||||
{"role": "user", "content": "test"},
|
||||
{
|
||||
"id": "rs_123",
|
||||
"summary": [{"text": "test", "type": "summary_text"}],
|
||||
"type": "reasoning",
|
||||
"content": None,
|
||||
"encrypted_content": None,
|
||||
"status": None,
|
||||
},
|
||||
{
|
||||
"arguments": "{}",
|
||||
"call_id": "call_123",
|
||||
"name": "get_today",
|
||||
"type": "function_call",
|
||||
"id": "fc_123",
|
||||
"status": "completed",
|
||||
},
|
||||
]
|
||||
|
||||
validated_input: Final = OpenAIResponsesAPIConfig()._validate_input_param(test_input)
|
||||
|
||||
assert len(validated_input) == 3
|
||||
reasoning_item: Final = validated_input[1]
|
||||
assert reasoning_item["type"] == "reasoning"
|
||||
assert "status" not in reasoning_item
|
||||
assert "content" not in reasoning_item
|
||||
assert "encrypted_content" not in reasoning_item
|
||||
assert reasoning_item["id"] == "rs_123"
|
||||
assert reasoning_item["summary"] == [{"text": "test", "type": "summary_text"}]
|
||||
|
||||
function_call_item: Final = validated_input[2]
|
||||
assert function_call_item["type"] == "function_call"
|
||||
assert function_call_item["status"] == "completed"
|
||||
|
|
|
|||
|
|
@ -1,7 +1,10 @@
|
|||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
|
|
@ -1131,3 +1134,166 @@ def test_transitive_probe_expansion_terminates_on_a_router_cycle():
|
|||
probes = hc_module._dependency_deployments_to_probe(a_only, router.model_list, router)
|
||||
|
||||
assert {d["model_info"]["id"] for d in probes} == {"b-1"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_audio_speech_health_check_uses_model_info_voice(
|
||||
httpx_transport: None, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
upstream: Final = respx_mock.post("https://speech.example/v1/audio/speech").respond(
|
||||
content=b"audio",
|
||||
headers={"content-type": "audio/mpeg"},
|
||||
)
|
||||
|
||||
healthy, unhealthy, _ = await hc_module.perform_health_check(
|
||||
[
|
||||
{
|
||||
"litellm_params": {
|
||||
"model": "openai/tts-1",
|
||||
"api_key": "fake-key",
|
||||
"api_base": "https://speech.example/v1",
|
||||
},
|
||||
"model_info": {"id": "speech", "mode": "audio_speech", "health_check_voice": "nova"},
|
||||
}
|
||||
],
|
||||
max_concurrency=1,
|
||||
)
|
||||
|
||||
assert len(healthy) == 1
|
||||
assert unhealthy == []
|
||||
assert upstream.called
|
||||
assert json.loads(upstream.calls.last.request.content)["voice"] == "nova"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_health_check_observes_the_concurrency_limit_and_queue(
|
||||
httpx_transport: None, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
request_started: Final = asyncio.Queue[None]()
|
||||
release: Final = asyncio.Event()
|
||||
|
||||
async def complete_request(_: httpx.Request) -> httpx.Response:
|
||||
request_started.put_nowait(None)
|
||||
await release.wait()
|
||||
return httpx.Response(
|
||||
200,
|
||||
content=b"audio",
|
||||
headers={"content-type": "audio/mpeg"},
|
||||
)
|
||||
|
||||
upstream: Final = respx_mock.post("https://health.example/v1/audio/speech").mock(side_effect=complete_request)
|
||||
model_list: Final = [
|
||||
{
|
||||
"litellm_params": {
|
||||
"model": "openai/tts-1",
|
||||
"api_key": "fake-key",
|
||||
"api_base": "https://health.example/v1",
|
||||
},
|
||||
"model_info": {"id": f"audio-{index}", "mode": "audio_speech"},
|
||||
}
|
||||
for index in range(10)
|
||||
]
|
||||
tasks_before: Final = len(asyncio.all_tasks())
|
||||
perform_task: Final = asyncio.create_task(hc_module.perform_health_check(model_list, max_concurrency=2))
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(request_started.get(), timeout=1)
|
||||
await asyncio.wait_for(request_started.get(), timeout=1)
|
||||
for _ in range(20):
|
||||
await asyncio.sleep(0)
|
||||
extra_requests_started: Final = request_started.qsize()
|
||||
tasks_while_blocked: Final = len(asyncio.all_tasks()) - tasks_before
|
||||
finally:
|
||||
release.set()
|
||||
healthy, unhealthy, _ = await perform_task
|
||||
|
||||
assert extra_requests_started == 0
|
||||
assert tasks_while_blocked <= 5
|
||||
assert upstream.call_count == 10
|
||||
assert len(healthy) == 10
|
||||
assert unhealthy == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_health_check_timeout_marks_a_blocked_provider_unhealthy(
|
||||
httpx_transport: None, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
never_release: Final = asyncio.Event()
|
||||
request_started: Final = asyncio.Event()
|
||||
|
||||
async def blocked_response(_: httpx.Request) -> httpx.Response:
|
||||
request_started.set()
|
||||
await never_release.wait()
|
||||
return httpx.Response(200, content=b"audio", headers={"content-type": "audio/mpeg"})
|
||||
|
||||
respx_mock.post("https://health.example/v1/audio/speech").mock(side_effect=blocked_response)
|
||||
model_list: Final = [
|
||||
{
|
||||
"litellm_params": {
|
||||
"model": "openai/tts-1",
|
||||
"api_key": "fake-key",
|
||||
"api_base": "https://health.example/v1",
|
||||
},
|
||||
"model_info": {"id": "blocked", "mode": "audio_speech", "health_check_timeout": 2},
|
||||
}
|
||||
]
|
||||
|
||||
healthy, unhealthy, _ = await asyncio.wait_for(
|
||||
hc_module.perform_health_check(model_list),
|
||||
timeout=4,
|
||||
)
|
||||
|
||||
assert request_started.is_set()
|
||||
assert unhealthy[0]["error"] == "Timeout exceeded"
|
||||
assert healthy == []
|
||||
assert len(unhealthy) == 1
|
||||
assert unhealthy[0]["model"] == "openai/tts-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_health_check_timeout_does_not_cancel_a_sibling(
|
||||
httpx_transport: None, respx_mock: respx.MockRouter
|
||||
) -> None:
|
||||
never_release: Final = asyncio.Event()
|
||||
slow_request_started: Final = asyncio.Event()
|
||||
|
||||
async def blocked_response(_: httpx.Request) -> httpx.Response:
|
||||
slow_request_started.set()
|
||||
await never_release.wait()
|
||||
return httpx.Response(200, content=b"audio", headers={"content-type": "audio/mpeg"})
|
||||
|
||||
respx_mock.post("https://slow.example/v1/audio/speech").mock(side_effect=blocked_response)
|
||||
fast_upstream: Final = respx_mock.post("https://fast.example/v1/audio/speech").respond(
|
||||
content=b"audio",
|
||||
headers={"content-type": "audio/mpeg"},
|
||||
)
|
||||
model_list: Final = [
|
||||
{
|
||||
"litellm_params": {
|
||||
"model": "openai/tts-1",
|
||||
"api_key": "fake-key",
|
||||
"api_base": "https://slow.example/v1",
|
||||
},
|
||||
"model_info": {"id": "slow", "mode": "audio_speech", "health_check_timeout": 1},
|
||||
},
|
||||
{
|
||||
"litellm_params": {
|
||||
"model": "openai/tts-1",
|
||||
"api_key": "fake-key",
|
||||
"api_base": "https://fast.example/v1",
|
||||
},
|
||||
"model_info": {"id": "fast", "mode": "audio_speech", "health_check_timeout": 2},
|
||||
},
|
||||
]
|
||||
|
||||
healthy, unhealthy, _ = await asyncio.wait_for(
|
||||
hc_module.perform_health_check(model_list, max_concurrency=1),
|
||||
timeout=4,
|
||||
)
|
||||
healthy_model_ids: Final = {endpoint["model_id"] for endpoint in healthy}
|
||||
unhealthy_model_ids: Final = {endpoint["model_id"] for endpoint in unhealthy}
|
||||
|
||||
assert slow_request_started.is_set()
|
||||
assert fast_upstream.called
|
||||
assert healthy_model_ids == {"fast"}
|
||||
assert unhealthy_model_ids == {"slow"}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,158 @@
|
|||
import json
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.types.llms.openai import ResponseCompletedEvent, ResponsesAPIResponse
|
||||
from tests.unit.proxy.conftest import httpx_transport
|
||||
|
||||
pytestmark: Final = pytest.mark.usefixtures(httpx_transport.__name__)
|
||||
_GEMINI_URL: Final = (
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/"
|
||||
"gemini-2.5-flash:(?:generateContent|streamGenerateContent).*"
|
||||
)
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def _function_call_response(signature: str) -> dict[str, object]:
|
||||
return {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"role": "model",
|
||||
"parts": [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": {"location": "San Francisco"},
|
||||
},
|
||||
"thoughtSignature": signature,
|
||||
}
|
||||
],
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 4,
|
||||
"candidatesTokenCount": 2,
|
||||
"totalTokenCount": 6,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_web_search_preview_is_sent_as_google_search() -> None:
|
||||
response_body: Final = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"role": "model", "parts": [{"text": "Search completed"}]},
|
||||
"finishReason": "STOP",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 2,
|
||||
"candidatesTokenCount": 2,
|
||||
"totalTokenCount": 4,
|
||||
},
|
||||
}
|
||||
|
||||
with respx.mock() as mock_router:
|
||||
route: Final = mock_router.post(url__regex=_GEMINI_URL).mock(
|
||||
return_value=httpx.Response(status_code=200, json=response_body)
|
||||
)
|
||||
response: Final = await litellm.aresponses(
|
||||
model="gemini/gemini-2.5-flash",
|
||||
api_key="test-key",
|
||||
input="Find current weather",
|
||||
tools=[{"type": "web_search_preview", "search_context_size": "low"}],
|
||||
)
|
||||
requests: Final = tuple(route.calls)
|
||||
|
||||
assert isinstance(response, ResponsesAPIResponse)
|
||||
assert response.output_text == "Search completed"
|
||||
assert len(requests) == 1
|
||||
request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content)
|
||||
assert request_body["tools"] == [{"googleSearch": {}}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_function_call_preserves_thought_signature() -> None:
|
||||
signature: Final = "gemini-thought-signature"
|
||||
response_body: Final = _function_call_response(signature)
|
||||
|
||||
with respx.mock() as mock_router:
|
||||
route: Final = mock_router.post(url__regex=_GEMINI_URL).mock(
|
||||
return_value=httpx.Response(status_code=200, json=response_body)
|
||||
)
|
||||
response: Final = await litellm.aresponses(
|
||||
model="gemini/gemini-2.5-flash",
|
||||
api_key="test-key",
|
||||
input="What is the weather in San Francisco?",
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"name": "get_weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"location": {"type": "string"}},
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
requests: Final = tuple(route.calls)
|
||||
|
||||
assert isinstance(response, ResponsesAPIResponse)
|
||||
function_calls: Final = tuple(item for item in response.output if item.type == "function_call")
|
||||
assert len(function_calls) == 1
|
||||
assert function_calls[0].name == "get_weather"
|
||||
assert function_calls[0].provider_specific_fields == {"thought_signature": signature}
|
||||
assert len(requests) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_streaming_function_call_preserves_thought_signature() -> None:
|
||||
signature: Final = "gemini-stream-thought-signature"
|
||||
response_body: Final = _function_call_response(signature)
|
||||
event_stream: Final = f"data: {json.dumps(response_body)}\n\n"
|
||||
|
||||
with respx.mock() as mock_router:
|
||||
route: Final = mock_router.post(url__regex=_GEMINI_URL).mock(
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
content=event_stream,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
)
|
||||
)
|
||||
stream: Final = await litellm.aresponses(
|
||||
model="gemini/gemini-2.5-flash",
|
||||
api_key="test-key",
|
||||
input="What is the weather in San Francisco?",
|
||||
stream=True,
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"name": "get_weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"location": {"type": "string"}},
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
events: Final = tuple([event async for event in stream])
|
||||
requests: Final = tuple(route.calls)
|
||||
|
||||
completed_events: Final = tuple(event for event in events if isinstance(event, ResponseCompletedEvent))
|
||||
assert len(completed_events) == 1
|
||||
function_calls: Final = tuple(item for item in completed_events[0].response.output if item.type == "function_call")
|
||||
assert len(function_calls) == 1
|
||||
assert function_calls[0].name == "get_weather"
|
||||
assert function_calls[0].provider_specific_fields == {"thought_signature": signature}
|
||||
assert len(requests) == 1
|
||||
109
tests/unit/responses/test_responses_api_lifecycle.py
Normal file
109
tests/unit/responses/test_responses_api_lifecycle.py
Normal file
|
|
@ -0,0 +1,109 @@
|
|||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from tests.unit.proxy.conftest import httpx_transport
|
||||
|
||||
pytestmark: Final = pytest.mark.usefixtures(httpx_transport.__name__)
|
||||
Provider: TypeAlias = Literal["anthropic", "gemini"]
|
||||
Operation: TypeAlias = Literal["delete", "get", "cancel"]
|
||||
CancelProvider: TypeAlias = Literal["openai", "azure"]
|
||||
|
||||
|
||||
def _invoke_unsupported_response_operation(provider: Provider, operation: Operation) -> object:
|
||||
match operation:
|
||||
case "delete":
|
||||
return litellm.delete_responses(
|
||||
response_id="resp_unsupported",
|
||||
custom_llm_provider=provider,
|
||||
api_key="sk-test",
|
||||
)
|
||||
case "get":
|
||||
return litellm.get_responses(
|
||||
response_id="resp_unsupported",
|
||||
custom_llm_provider=provider,
|
||||
api_key="sk-test",
|
||||
)
|
||||
case "cancel":
|
||||
return litellm.cancel_responses(
|
||||
response_id="resp_unsupported",
|
||||
custom_llm_provider=provider,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
|
||||
async def _invoke_cancel_response(sync_mode: bool, call_kwargs: dict[str, object]) -> object:
|
||||
if sync_mode:
|
||||
return litellm.cancel_responses(**call_kwargs)
|
||||
return await litellm.acancel_responses(**call_kwargs)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("provider", "operation"),
|
||||
(
|
||||
("anthropic", "delete"),
|
||||
("anthropic", "get"),
|
||||
("anthropic", "cancel"),
|
||||
("gemini", "delete"),
|
||||
("gemini", "get"),
|
||||
("gemini", "cancel"),
|
||||
),
|
||||
)
|
||||
def test_unsupported_response_lifecycle_operation_fails_before_http(provider: Provider, operation: Operation) -> None:
|
||||
with respx.mock() as mock_router:
|
||||
with pytest.raises(litellm.APIConnectionError) as exc_info:
|
||||
_invoke_unsupported_response_operation(provider, operation)
|
||||
calls: Final = tuple(mock_router.calls)
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert f"not supported for {provider}" in str(exc_info.value)
|
||||
assert calls == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("provider", "sync_mode"),
|
||||
(("openai", True), ("openai", False), ("azure", True), ("azure", False)),
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel_responses_404_surfaces_openai_api_error_with_exact_url(
|
||||
provider: CancelProvider, sync_mode: bool
|
||||
) -> None:
|
||||
response_id: Final = "resp_missing"
|
||||
error_body: Final = {
|
||||
"error": {
|
||||
"message": "Response was not found",
|
||||
"type": "invalid_request_error",
|
||||
"code": "not_found",
|
||||
}
|
||||
}
|
||||
api_base: Final = (
|
||||
"https://api.openai.com/v1" if provider == "openai" else "https://example-resource.openai.azure.com"
|
||||
)
|
||||
expected_url: Final = (
|
||||
f"{api_base}/responses/{response_id}/cancel"
|
||||
if provider == "openai"
|
||||
else f"{api_base}/openai/responses/{response_id}/cancel?api-version=2025-03-01-preview"
|
||||
)
|
||||
|
||||
with respx.mock() as mock_router:
|
||||
route: Final = mock_router.post(expected_url).mock(
|
||||
return_value=httpx.Response(status_code=404, json=error_body)
|
||||
)
|
||||
call_kwargs: Final = {
|
||||
"custom_llm_provider": provider,
|
||||
"api_key": "sk-test",
|
||||
"api_base": api_base,
|
||||
"api_version": "2025-03-01-preview",
|
||||
"response_id": response_id,
|
||||
}
|
||||
with pytest.raises(openai.APIError) as exc_info:
|
||||
await _invoke_cancel_response(sync_mode, call_kwargs)
|
||||
requests: Final = tuple(route.calls)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
assert len(requests) == 1
|
||||
assert str(requests[0].request.url) == expected_url
|
||||
83
tests/unit/secret_managers/test_google_secret_manager.py
Normal file
83
tests/unit/secret_managers/test_google_secret_manager.py
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
import base64
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
from litellm.secret_managers.google_secret_manager import GoogleSecretManager
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _CachedVertexCredentials:
|
||||
token: str
|
||||
quota_project_id: str | None
|
||||
expired: bool = False
|
||||
|
||||
def refresh(self, request: object) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def _google_secret_manager(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
project_id: str,
|
||||
) -> GoogleSecretManager:
|
||||
monkeypatch.setenv("GOOGLE_SECRET_MANAGER_PROJECT_ID", project_id)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
|
||||
secret_manager: Final = GoogleSecretManager()
|
||||
credentials: Final = _CachedVertexCredentials(
|
||||
token="test-gsm-token",
|
||||
quota_project_id=project_id,
|
||||
)
|
||||
vertex_chat_completion: Final = litellm.vertex_chat_completion
|
||||
monkeypatch.setitem(
|
||||
vertex_chat_completion._credentials_project_mapping,
|
||||
(None, project_id),
|
||||
(credentials, project_id),
|
||||
)
|
||||
return secret_manager
|
||||
|
||||
|
||||
def test_google_secret_manager_decodes_secret_and_requests_latest_version(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
project_id: Final = "test-secret-project"
|
||||
secret_manager: Final = _google_secret_manager(monkeypatch, project_id)
|
||||
secret_url: Final = (
|
||||
f"https://secretmanager.googleapis.com/v1/projects/{project_id}/secrets/OPENAI_API_KEY/versions/latest:access"
|
||||
)
|
||||
encoded_secret: Final = base64.b64encode(b"anything").decode("ascii")
|
||||
upstream: Final[respx.MockRouter]
|
||||
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
secret_route: Final = upstream.get(secret_url).respond(
|
||||
200,
|
||||
json={"payload": {"data": encoded_secret}},
|
||||
)
|
||||
|
||||
result: Final = secret_manager.get_secret_from_google_secret_manager("OPENAI_API_KEY")
|
||||
|
||||
assert result == "anything"
|
||||
assert secret_route.called
|
||||
assert len(upstream.calls) == 1
|
||||
assert str(upstream.calls.last.request.url) == secret_url
|
||||
assert upstream.calls.last.request.headers["Authorization"] == "Bearer test-gsm-token"
|
||||
|
||||
|
||||
def test_google_secret_manager_returns_cached_values_without_http(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
project_id: Final = "test-secret-project"
|
||||
secret_manager: Final = _google_secret_manager(monkeypatch, project_id)
|
||||
secret_manager.cache.set_cache("cached-none", None)
|
||||
secret_manager.cache.set_cache("cached-string", "lite-llm")
|
||||
upstream: Final[respx.MockRouter]
|
||||
|
||||
with respx.mock() as upstream:
|
||||
missing_value: Final = secret_manager.get_secret_from_google_secret_manager("cached-none")
|
||||
cached_value: Final = secret_manager.get_secret_from_google_secret_manager("cached-string")
|
||||
|
||||
assert missing_value is None
|
||||
assert cached_value == "lite-llm"
|
||||
assert upstream.calls.call_count == 0
|
||||
|
|
@ -3,6 +3,7 @@ import json
|
|||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Final
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -230,6 +231,13 @@ def test_oidc_circleci_success(monkeypatch):
|
|||
assert result == "circleci_token"
|
||||
|
||||
|
||||
def test_oidc_circleci_v2_returns_the_environment_token(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
token: Final = "circleci-v2-token"
|
||||
monkeypatch.setenv("CIRCLE_OIDC_TOKEN_V2", token)
|
||||
|
||||
assert get_secret("oidc/circleci_v2/test-audience") == token
|
||||
|
||||
|
||||
def test_oidc_circleci_failure(monkeypatch):
|
||||
monkeypatch.delenv("CIRCLE_OIDC_TOKEN", raising=False)
|
||||
secret_name = "oidc/circleci/test-audience"
|
||||
|
|
|
|||
732
tests/unit/test_router/test_router_provider_endpoints.py
Normal file
732
tests/unit/test_router/test_router_provider_endpoints.py
Normal file
|
|
@ -0,0 +1,732 @@
|
|||
import asyncio
|
||||
import contextlib
|
||||
import io
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import litellm
|
||||
import pytest
|
||||
import respx
|
||||
from litellm import CustomLogger, Router
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.router_utils.router_callbacks.track_deployment_metrics import (
|
||||
get_deployment_successes_for_current_minute,
|
||||
)
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
from litellm.types.router import RoutingStrategy
|
||||
from litellm.types.utils import ImageResponse, ModelResponse, RerankResponse
|
||||
from openai import AsyncAzureOpenAI, AsyncOpenAI
|
||||
|
||||
EVENT_TIMEOUT_SECONDS: Final = 5
|
||||
ROUTING_SELECTIONS: Final = 40
|
||||
ROUTING_MESSAGES: Final = ({"role": "user", "content": "route this request"},)
|
||||
EXPENSIVE_COSTS: Final = {"input_cost_per_token": 1.0, "output_cost_per_token": 1.0}
|
||||
CHEAP_COSTS: Final = {"input_cost_per_token": 1e-9, "output_cost_per_token": 1e-9}
|
||||
CHAT_RESPONSE: Final = {
|
||||
"id": "chatcmpl-1",
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 1, "total_tokens": 6},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def httpx_transport(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
yield
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def router_minute_pinned(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
pinned: Final = datetime(2026, 1, 1, 12, 0, 30, tzinfo=timezone.utc)
|
||||
monkeypatch.setattr("litellm.router.get_utc_datetime", lambda: pinned)
|
||||
|
||||
|
||||
class _RouterLoggingCapture(CustomLogger):
|
||||
def __init__(self, model_id: str) -> None:
|
||||
super().__init__()
|
||||
self.model_id: Final = model_id
|
||||
self.success_events: asyncio.Queue[tuple[object | None, object | None]] = asyncio.Queue()
|
||||
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: dict[str, object],
|
||||
response_obj: object,
|
||||
start_time: object,
|
||||
end_time: object,
|
||||
) -> None:
|
||||
standard_logging_object: Final = kwargs.get("standard_logging_object")
|
||||
if not isinstance(standard_logging_object, dict) or standard_logging_object.get("model_id") != self.model_id:
|
||||
return
|
||||
self.success_events.put_nowait((kwargs.get("client"), standard_logging_object))
|
||||
|
||||
async def next_event(self) -> tuple[object | None, object | None]:
|
||||
return await asyncio.wait_for(self.success_events.get(), timeout=EVENT_TIMEOUT_SECONDS)
|
||||
|
||||
|
||||
def _rpm_tpm_router(model_id: str) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-5-mini",
|
||||
"litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake", "tpm": 1000, "rpm": 100},
|
||||
"model_info": {"id": model_id},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _ratelimit_headers(response: ModelResponse | CustomStreamWrapper) -> dict[str, int]:
|
||||
return {k: v for k, v in response._hidden_params["additional_headers"].items() if k.startswith("x-ratelimit-")}
|
||||
|
||||
|
||||
def _openai_router_client() -> AsyncOpenAI:
|
||||
return AsyncOpenAI(api_key="sk-fake")
|
||||
|
||||
|
||||
def _azure_router_client() -> AsyncAzureOpenAI:
|
||||
return AsyncAzureOpenAI(api_key="sk-fake", azure_endpoint="https://azure.test", api_version="2025-02-01-preview")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("litellm_params", "build_client", "route_url"),
|
||||
[
|
||||
(
|
||||
{"model": "whisper-1", "api_key": "sk-fake"},
|
||||
_openai_router_client,
|
||||
"https://api.openai.com/v1/audio/transcriptions",
|
||||
),
|
||||
(
|
||||
{
|
||||
"model": "azure/whisper",
|
||||
"api_base": "https://azure.test",
|
||||
"api_key": "sk-fake",
|
||||
"api_version": "2025-02-01-preview",
|
||||
},
|
||||
_azure_router_client,
|
||||
"https://azure.test/openai/deployments/whisper/audio/transcriptions?api-version=2025-02-01-preview",
|
||||
),
|
||||
],
|
||||
ids=["openai", "azure"],
|
||||
)
|
||||
async def test_router_transcription_reuses_router_level_client_for_each_deployment(
|
||||
litellm_params: dict[str, str],
|
||||
build_client: Callable[[], AsyncOpenAI],
|
||||
route_url: str,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
model_id: Final = f"whisper-{uuid.uuid4().hex}"
|
||||
capture: Final = _RouterLoggingCapture(model_id)
|
||||
monkeypatch.setattr(litellm, "callbacks", [capture])
|
||||
router: Final = Router(
|
||||
model_list=[{"model_name": "whisper", "litellm_params": dict(litellm_params), "model_info": {"id": model_id}}]
|
||||
)
|
||||
router_level_client: Final = build_client()
|
||||
router.cache.set_cache(key=f"{model_id}_async_client", value=router_level_client, local_only=True)
|
||||
route: Final = respx_mock.post(route_url).respond(200, json={"text": "hello"})
|
||||
|
||||
response: Final = await router.atranscription(
|
||||
model="whisper", file=("speech.wav", io.BytesIO(b"offline audio"), "audio/wav")
|
||||
)
|
||||
public_client, public_logging = await capture.next_event()
|
||||
internal_response: Final = await router._atranscription(
|
||||
model="whisper", file=("speech.wav", io.BytesIO(b"offline audio"), "audio/wav")
|
||||
)
|
||||
internal_client, _ = await capture.next_event()
|
||||
upstream_requests: Final = tuple(call.request for call in route.calls)
|
||||
|
||||
assert public_client == str(router_level_client)
|
||||
assert internal_client == str(router_level_client)
|
||||
assert tuple(str(request.url) for request in upstream_requests) == (route_url, route_url)
|
||||
assert all(
|
||||
request.headers.get("content-type", "").startswith("multipart/form-data")
|
||||
and b'filename="speech.wav"' in request.content
|
||||
and b"offline audio" in request.content
|
||||
for request in upstream_requests
|
||||
)
|
||||
assert isinstance(public_logging, dict)
|
||||
assert public_logging.get("model_group") == "whisper"
|
||||
assert response.text == "hello"
|
||||
assert internal_response.text == "hello"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_speech_returns_binary_content_and_logs_model_group(
|
||||
respx_mock: respx.MockRouter,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
model_id: Final = f"tts-{uuid.uuid4().hex}"
|
||||
capture: Final = _RouterLoggingCapture(model_id)
|
||||
monkeypatch.setattr(litellm, "callbacks", [capture])
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "tts",
|
||||
"litellm_params": {"model": "openai/tts-1", "api_key": "sk-fake"},
|
||||
"model_info": {"id": model_id},
|
||||
}
|
||||
]
|
||||
)
|
||||
route: Final = respx_mock.post("https://api.openai.com/v1/audio/speech").respond(200, content=b"audio")
|
||||
|
||||
response: Final = await router.aspeech(model="tts", input="hello", voice="alloy")
|
||||
_, standard_logging_object = await capture.next_event()
|
||||
|
||||
assert route.call_count == 1
|
||||
assert json.loads(route.calls[0].request.content) == {"model": "tts-1", "input": "hello", "voice": "alloy"}
|
||||
assert isinstance(response, HttpxBinaryResponseContent)
|
||||
assert response.content == b"audio"
|
||||
assert isinstance(standard_logging_object, dict)
|
||||
assert standard_logging_object["model_group"] == "tts"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_rerank_returns_valid_response_from_public_and_underlying_calls(
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "cohere-rerank",
|
||||
"litellm_params": {"model": "cohere/rerank-english-v3.0", "api_key": "sk-fake"},
|
||||
}
|
||||
]
|
||||
)
|
||||
route: Final = respx_mock.post("https://api.cohere.com/v2/rerank").respond(
|
||||
200,
|
||||
json={
|
||||
"id": "rerank-1",
|
||||
"results": [{"index": 0, "relevance_score": 0.9}],
|
||||
"meta": {"api_version": {"version": "2"}},
|
||||
},
|
||||
)
|
||||
|
||||
public_response: Final = await router.arerank(
|
||||
model="cohere-rerank",
|
||||
query="hello",
|
||||
documents=["hello", "world"],
|
||||
top_n=1,
|
||||
)
|
||||
underlying_response: Final = await router._arerank(
|
||||
model="cohere-rerank",
|
||||
query="hello",
|
||||
documents=["hello", "world"],
|
||||
top_n=1,
|
||||
)
|
||||
|
||||
assert route.call_count == 2
|
||||
request_bodies: Final = tuple(json.loads(call.request.content) for call in route.calls)
|
||||
assert all(
|
||||
body["model"] == "rerank-english-v3.0"
|
||||
and body["query"] == "hello"
|
||||
and body["documents"] == ["hello", "world"]
|
||||
and body["top_n"] == 1
|
||||
for body in request_bodies
|
||||
)
|
||||
public_validated: Final = RerankResponse.model_validate(public_response)
|
||||
assert public_validated.id == "rerank-1"
|
||||
assert public_validated.results[0]["relevance_score"] == 0.9
|
||||
underlying_validated: Final = RerankResponse.model_validate(underlying_response)
|
||||
assert underlying_validated.id == "rerank-1"
|
||||
assert underlying_validated.results[0]["relevance_score"] == 0.9
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("model", "expected_model", "expected_api_key"),
|
||||
[
|
||||
("omni-moderation-latest", "omni-moderation-latest", "sk-catch-all"),
|
||||
("openai/omni-moderation-latest", "omni-moderation-latest", "sk-openai-wildcard"),
|
||||
(None, None, "sk-env"),
|
||||
],
|
||||
)
|
||||
async def test_router_moderation_routes_through_wildcard_deployments(
|
||||
model: str | None,
|
||||
expected_model: str | None,
|
||||
expected_api_key: str,
|
||||
respx_mock: respx.MockRouter,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-env")
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "openai/*",
|
||||
"litellm_params": {"model": "openai/*", "api_key": "sk-openai-wildcard"},
|
||||
},
|
||||
{
|
||||
"model_name": "*",
|
||||
"litellm_params": {"model": "openai/*", "api_key": "sk-catch-all"},
|
||||
},
|
||||
]
|
||||
)
|
||||
route: Final = respx_mock.post("https://api.openai.com/v1/moderations").respond(
|
||||
200,
|
||||
json={
|
||||
"id": "modr-1",
|
||||
"model": "omni-moderation-latest",
|
||||
"results": [{"flagged": False, "categories": {}, "category_scores": {}}],
|
||||
},
|
||||
)
|
||||
|
||||
response: Final = await router.amoderation(model=model, input="hello")
|
||||
|
||||
assert route.call_count == 1
|
||||
upstream_request: Final = route.calls[0].request
|
||||
expected_body: Final = {"input": "hello"} if expected_model is None else {"input": "hello", "model": expected_model}
|
||||
assert json.loads(upstream_request.content) == expected_body
|
||||
assert upstream_request.headers["authorization"] == f"Bearer {expected_api_key}"
|
||||
assert response.id == "modr-1"
|
||||
assert response.model == "omni-moderation-latest"
|
||||
assert response.results[0].flagged is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
async def test_router_image_generation_returns_valid_image_response(
|
||||
sync_mode: bool,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-image-1",
|
||||
"litellm_params": {"model": "openai/gpt-image-1", "api_key": "sk-fake"},
|
||||
}
|
||||
]
|
||||
)
|
||||
route: Final = respx_mock.post("https://api.openai.com/v1/images/generations").respond(
|
||||
200,
|
||||
json={"created": 1700000000, "data": [{"url": "https://images.test/result.png"}]},
|
||||
)
|
||||
|
||||
response: Final = (
|
||||
router._image_generation(model="gpt-image-1", prompt="a cat")
|
||||
if sync_mode
|
||||
else await router._aimage_generation(model="gpt-image-1", prompt="a cat")
|
||||
)
|
||||
|
||||
assert route.call_count == 1
|
||||
request_body: Final = json.loads(route.calls[0].request.content)
|
||||
assert request_body["model"] == "gpt-image-1"
|
||||
assert request_body["prompt"] == "a cat"
|
||||
validated: Final = ImageResponse.model_validate(response)
|
||||
assert validated.data[0].url == "https://images.test/result.png"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("router_minute_pinned")
|
||||
async def test_router_acompletion_headers_read_post_increment_counter_and_count_once(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
capture: Final = _RouterLoggingCapture("lit-3058-async")
|
||||
monkeypatch.setattr(litellm, "callbacks", [capture])
|
||||
router: Final = _rpm_tpm_router("lit-3058-async")
|
||||
|
||||
response: Final = await router.acompletion(
|
||||
model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong"
|
||||
)
|
||||
total_tokens: Final = response.usage.total_tokens
|
||||
headers: Final = _ratelimit_headers(response)
|
||||
|
||||
assert total_tokens > 0
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 1000 - total_tokens
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
|
||||
|
||||
await capture.next_event()
|
||||
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("router_minute_pinned")
|
||||
async def test_router_stream_counts_request_before_headers_and_tokens_once_on_completion(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
capture: Final = _RouterLoggingCapture("lit-3058-stream")
|
||||
monkeypatch.setattr(litellm, "callbacks", [capture])
|
||||
router: Final = _rpm_tpm_router("lit-3058-stream")
|
||||
stream: Final = await router.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="pong pong pong",
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
headers: Final = _ratelimit_headers(stream)
|
||||
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 1000
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (0, 1)
|
||||
|
||||
chunks: Final = [chunk async for chunk in stream]
|
||||
total_tokens: Final = chunks[-1].usage.total_tokens
|
||||
await capture.next_event()
|
||||
|
||||
assert total_tokens > 0
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
|
||||
|
||||
|
||||
def test_router_validate_fallbacks_accepts_well_formed_and_rejects_malformed_entries() -> None:
|
||||
router: Final = Router(model_list=[])
|
||||
|
||||
assert router.validate_fallbacks([{"gpt-5.5": ["gpt-5-mini"]}, {"gpt-5-mini": ["gpt-5.5"]}]) is None
|
||||
with pytest.raises(ValueError, match="must have exactly one key"):
|
||||
router.validate_fallbacks([{"primary": "fallback", "other": "fallback"}])
|
||||
with pytest.raises(ValueError, match="is not a dictionary"):
|
||||
router.validate_fallbacks(["primary"])
|
||||
|
||||
|
||||
def _routing_deployment(deployment_id: str, extra_params: dict[str, float]) -> dict[str, object]:
|
||||
return {
|
||||
"model_name": "gpt",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "sk-fake",
|
||||
"api_base": f"https://{deployment_id}.test/v1",
|
||||
**extra_params,
|
||||
},
|
||||
"model_info": {"id": deployment_id},
|
||||
}
|
||||
|
||||
|
||||
def _routing_router(
|
||||
strategy: RoutingStrategy | str,
|
||||
a_params: dict[str, float] | None = None,
|
||||
b_params: dict[str, float] | None = None,
|
||||
) -> Router:
|
||||
return Router(
|
||||
model_list=[_routing_deployment("a", a_params or {}), _routing_deployment("b", b_params or {})],
|
||||
routing_strategy=strategy,
|
||||
disable_cooldowns=True,
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
|
||||
async def _selected_deployment_ids(router: Router) -> frozenset[str]:
|
||||
deployments: Final = [
|
||||
await router.async_get_available_deployment(model="gpt", messages=list(ROUTING_MESSAGES), request_kwargs={})
|
||||
for _ in range(ROUTING_SELECTIONS)
|
||||
]
|
||||
return frozenset(deployment["model_info"]["id"] for deployment in deployments)
|
||||
|
||||
|
||||
def _enum_and_string(strategy: RoutingStrategy) -> list[RoutingStrategy | str]:
|
||||
return [strategy, strategy.value]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.COST_BASED))
|
||||
async def test_router_cost_based_routing_selects_cheapest_deployment(strategy: RoutingStrategy | str) -> None:
|
||||
router: Final = _routing_router(strategy, EXPENSIVE_COSTS, CHEAP_COSTS)
|
||||
|
||||
assert await _selected_deployment_ids(router) == frozenset({"b"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"strategy",
|
||||
[*_enum_and_string(RoutingStrategy.USAGE_BASED_ROUTING), *_enum_and_string(RoutingStrategy.USAGE_BASED_ROUTING_V2)],
|
||||
)
|
||||
async def test_router_usage_based_routing_skips_deployment_over_tpm_limit(strategy: RoutingStrategy | str) -> None:
|
||||
router: Final = _routing_router(strategy, {"tpm": 1}, {"tpm": 1_000_000})
|
||||
|
||||
assert await _selected_deployment_ids(router) == frozenset({"b"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.LATENCY_BASED))
|
||||
async def test_router_latency_based_routing_avoids_deployment_that_timed_out(
|
||||
strategy: RoutingStrategy | str,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
router: Final = _routing_router(strategy)
|
||||
timed_out: Final = respx_mock.post("https://a.test/v1/chat/completions").mock(
|
||||
side_effect=httpx.ReadTimeout("upstream timed out")
|
||||
)
|
||||
respx_mock.post("https://b.test/v1/chat/completions").respond(200, json=CHAT_RESPONSE)
|
||||
|
||||
for _ in range(5):
|
||||
if timed_out.called:
|
||||
break
|
||||
with contextlib.suppress(litellm.Timeout):
|
||||
await router.acompletion(model="gpt", messages=list(ROUTING_MESSAGES), max_retries=0)
|
||||
|
||||
assert timed_out.called
|
||||
assert await _selected_deployment_ids(router) == frozenset({"b"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.LEAST_BUSY))
|
||||
async def test_router_least_busy_routing_avoids_deployment_with_request_in_flight(
|
||||
strategy: RoutingStrategy | str,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
router: Final = _routing_router(strategy)
|
||||
upstream_hosts: Final = asyncio.Queue[str]()
|
||||
release_upstream: Final = asyncio.Event()
|
||||
|
||||
async def hold_request(request: httpx.Request) -> httpx.Response:
|
||||
upstream_hosts.put_nowait(request.url.host.split(".")[0])
|
||||
await release_upstream.wait()
|
||||
return httpx.Response(200, json=CHAT_RESPONSE)
|
||||
|
||||
respx_mock.post(url__regex=r"https://[ab]\.test/v1/chat/completions").mock(side_effect=hold_request)
|
||||
in_flight: Final = asyncio.create_task(router.acompletion(model="gpt", messages=list(ROUTING_MESSAGES)))
|
||||
try:
|
||||
busy_id: Final = await asyncio.wait_for(upstream_hosts.get(), timeout=EVENT_TIMEOUT_SECONDS)
|
||||
selected: Final = await _selected_deployment_ids(router)
|
||||
finally:
|
||||
release_upstream.set()
|
||||
await in_flight
|
||||
|
||||
assert selected == frozenset({"a", "b"} - {busy_id})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_simple_shuffle_ignores_cost_and_spreads_across_deployments() -> None:
|
||||
router: Final = _routing_router("simple-shuffle", EXPENSIVE_COSTS, CHEAP_COSTS)
|
||||
|
||||
assert await _selected_deployment_ids(router) == frozenset({"a", "b"})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.PROVIDER_BUDGET_LIMITING))
|
||||
def test_router_routing_strategy_init_accepts_provider_budget_strategy(strategy: RoutingStrategy | str) -> None:
|
||||
router: Final = _routing_router(strategy)
|
||||
|
||||
router.routing_strategy_init(routing_strategy=strategy, routing_strategy_args={})
|
||||
|
||||
assert router.get_settings()["routing_strategy"] == "provider-budget-routing"
|
||||
|
||||
|
||||
def test_router_track_deployment_metrics_updates_observable_usage() -> None:
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-5-mini",
|
||||
"litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake"},
|
||||
"model_info": {"id": "metrics-deployment"},
|
||||
}
|
||||
]
|
||||
)
|
||||
deployment: Final = router.model_list[0]
|
||||
|
||||
router._track_deployment_metrics(deployment=deployment, parent_otel_span=None)
|
||||
router._track_deployment_metrics(deployment=deployment, parent_otel_span=None)
|
||||
|
||||
assert router.cache.get_cache(key="metrics-deployment", local_only=True) == 2
|
||||
|
||||
|
||||
class _GatedIncrementCache(DualCache):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(in_memory_cache=InMemoryCache())
|
||||
self.first_increment_started = asyncio.Event()
|
||||
self.release_first_increment = asyncio.Event()
|
||||
self.deployment_success_incremented = asyncio.Event()
|
||||
self.increment_calls = 0
|
||||
|
||||
def increment_cache(self, key: str, value: int, local_only: bool = False, **kwargs: object) -> int:
|
||||
result: Final = super().increment_cache(key=key, value=value, local_only=local_only, **kwargs)
|
||||
if key.endswith(":successes"):
|
||||
self.deployment_success_incremented.set()
|
||||
return result
|
||||
|
||||
async def async_increment_cache_pipeline(
|
||||
self,
|
||||
increment_list: list[RedisPipelineIncrementOperation],
|
||||
local_only: bool = False,
|
||||
parent_otel_span: object = None,
|
||||
**kwargs: object,
|
||||
) -> list[float] | None:
|
||||
self.increment_calls += 1
|
||||
if self.increment_calls == 1:
|
||||
self.first_increment_started.set()
|
||||
await self.release_first_increment.wait()
|
||||
return await super().async_increment_cache_pipeline(
|
||||
increment_list=increment_list,
|
||||
local_only=local_only,
|
||||
parent_otel_span=parent_otel_span,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class _UnavailableIncrementCache(DualCache):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(in_memory_cache=InMemoryCache())
|
||||
self.first_increment_started = asyncio.Event()
|
||||
self.release_first_increment = asyncio.Event()
|
||||
self.deployment_success_incremented = asyncio.Event()
|
||||
self.increment_calls = 0
|
||||
|
||||
def increment_cache(self, key: str, value: int, local_only: bool = False, **kwargs: object) -> int:
|
||||
result: Final = super().increment_cache(key=key, value=value, local_only=local_only, **kwargs)
|
||||
if key.endswith(":successes"):
|
||||
self.deployment_success_incremented.set()
|
||||
return result
|
||||
|
||||
async def async_increment_cache_pipeline(
|
||||
self,
|
||||
increment_list: list[RedisPipelineIncrementOperation],
|
||||
local_only: bool = False,
|
||||
parent_otel_span: object = None,
|
||||
**kwargs: object,
|
||||
) -> list[float] | None:
|
||||
self.increment_calls += 1
|
||||
if self.increment_calls == 1:
|
||||
self.first_increment_started.set()
|
||||
await self.release_first_increment.wait()
|
||||
raise RuntimeError("cache unavailable")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("router_minute_pinned")
|
||||
async def test_router_success_callback_during_pre_header_increment_does_not_double_count() -> None:
|
||||
router: Final = _rpm_tpm_router("lit-3058-race")
|
||||
cache: Final = _GatedIncrementCache()
|
||||
router.cache = cache
|
||||
request: Final = asyncio.create_task(
|
||||
router.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="pong",
|
||||
)
|
||||
)
|
||||
|
||||
await asyncio.wait_for(cache.first_increment_started.wait(), timeout=EVENT_TIMEOUT_SECONDS)
|
||||
await asyncio.wait_for(cache.deployment_success_incremented.wait(), timeout=EVENT_TIMEOUT_SECONDS)
|
||||
assert get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1
|
||||
assert cache.increment_calls == 1
|
||||
cache.release_first_increment.set()
|
||||
response: Final = await request
|
||||
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (response.usage.total_tokens, 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("router_minute_pinned")
|
||||
async def test_router_failed_pre_header_increment_clears_counted_tokens_stamp() -> None:
|
||||
router: Final = _rpm_tpm_router("lit-3058-fail")
|
||||
cache: Final = _UnavailableIncrementCache()
|
||||
router.cache = cache
|
||||
metadata: Final[dict[str, object]] = {}
|
||||
request: Final = asyncio.create_task(
|
||||
router.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="pong",
|
||||
metadata=metadata,
|
||||
)
|
||||
)
|
||||
|
||||
await asyncio.wait_for(cache.first_increment_started.wait(), timeout=EVENT_TIMEOUT_SECONDS)
|
||||
await asyncio.wait_for(cache.deployment_success_incremented.wait(), timeout=EVENT_TIMEOUT_SECONDS)
|
||||
assert metadata[ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY] == 30
|
||||
assert get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1
|
||||
assert cache.increment_calls == 1
|
||||
cache.release_first_increment.set()
|
||||
response: Final = await request
|
||||
|
||||
assert response.usage.total_tokens == 30
|
||||
assert ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY not in metadata
|
||||
assert _ratelimit_headers(response)["x-ratelimit-remaining-requests"] == 100
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (None, None)
|
||||
|
||||
|
||||
ASSISTANT_RESPONSE: Final = {
|
||||
"object": "assistant",
|
||||
"created_at": 1700000000,
|
||||
"name": "offline",
|
||||
"description": None,
|
||||
"model": "gpt-4o-mini",
|
||||
"instructions": "hello",
|
||||
"tools": [],
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_assistants_endpoint_factory_invokes_provider(
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
router: Final = Router(model_list=[])
|
||||
route: Final = respx_mock.post("https://api.openai.com/v1/assistants").respond(
|
||||
200, json={**ASSISTANT_RESPONSE, "id": "asst-1"}
|
||||
)
|
||||
|
||||
response: Final = await router._pass_through_assistants_endpoint_factory(
|
||||
original_function=litellm.acreate_assistants,
|
||||
custom_llm_provider="openai",
|
||||
model="gpt-4o-mini",
|
||||
api_key="sk-fake",
|
||||
name="offline",
|
||||
)
|
||||
|
||||
assert route.call_count == 1
|
||||
assert json.loads(route.calls[0].request.content) == {"model": "gpt-4o-mini", "name": "offline"}
|
||||
assert response.id == "asst-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_factory_function_returns_invokable_assistants_wrapper(
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
router: Final = Router(model_list=[])
|
||||
route: Final = respx_mock.post("https://api.openai.com/v1/assistants").respond(
|
||||
200, json={**ASSISTANT_RESPONSE, "id": "asst-2"}
|
||||
)
|
||||
wrapper: Final = router.factory_function(litellm.acreate_assistants, call_type="assistants")
|
||||
|
||||
response: Final = await wrapper(
|
||||
custom_llm_provider="openai",
|
||||
model="gpt-4o-mini",
|
||||
api_key="sk-fake",
|
||||
name="offline",
|
||||
)
|
||||
|
||||
assert route.call_count == 1
|
||||
assert json.loads(route.calls[0].request.content) == {"model": "gpt-4o-mini", "name": "offline"}
|
||||
assert response.id == "asst-2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_moderation_endpoint_factory_invokes_default_model(
|
||||
respx_mock: respx.MockRouter,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-fake")
|
||||
router: Final = Router(model_list=[])
|
||||
route: Final = respx_mock.post("https://api.openai.com/v1/moderations").respond(
|
||||
200,
|
||||
json={
|
||||
"id": "modr-2",
|
||||
"model": "omni-moderation-latest",
|
||||
"results": [{"flagged": False, "categories": {}, "category_scores": {}}],
|
||||
},
|
||||
)
|
||||
|
||||
response: Final = await router._pass_through_moderation_endpoint_factory(
|
||||
original_function=litellm.amoderation,
|
||||
custom_llm_provider="openai",
|
||||
input="hello",
|
||||
model=None,
|
||||
api_key="sk-fake",
|
||||
)
|
||||
|
||||
assert route.call_count == 1
|
||||
assert json.loads(route.calls[0].request.content) == {"input": "hello"}
|
||||
assert response.id == "modr-2"
|
||||
|
|
@ -37,6 +37,10 @@ from litellm.caching.in_memory_cache import InMemoryCache
|
|||
from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.humanloop import HumanloopLogger
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import BaseAudioTranscriptionConfig
|
||||
from litellm.integrations.langfuse.langfuse_prompt_management import LangfusePromptManagement
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
from litellm.litellm_core_utils.duration_parser import (
|
||||
_extract_from_regex,
|
||||
duration_in_seconds,
|
||||
|
|
@ -50,7 +54,7 @@ from litellm.proxy.utils import is_valid_api_key
|
|||
from litellm.types.caching import CachingSupportedCallTypes
|
||||
from litellm.types.integrations.custom_logger import HEADROOM_CONVERTED_STREAM_KEY
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams
|
||||
from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams, LiteLLM_Params
|
||||
from litellm.types.utils import (
|
||||
ADDRESSED_RESPONSE_ID_FIELD,
|
||||
CallTypes,
|
||||
|
|
@ -74,6 +78,7 @@ from litellm.types.utils import (
|
|||
from litellm.types.videos.main import VideoObject
|
||||
from litellm.utils import (
|
||||
_invalidate_model_cost_lowercase_map,
|
||||
add_custom_logger_callback_to_specific_event,
|
||||
check_valid_key,
|
||||
CustomStreamWrapper,
|
||||
filter_out_litellm_params,
|
||||
|
|
@ -6802,6 +6807,7 @@ def setup_and_teardown():
|
|||
|
||||
MODEL: Final = "anthropic/claude-haiku-4-5"
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown")
|
||||
def test_validate_tool_choice_none():
|
||||
"""Test that None is returned as-is."""
|
||||
|
|
@ -6946,6 +6952,7 @@ _SCALAR_DEFAULTS = {
|
|||
"api_key": getattr(litellm, "api_key", None),
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def setup_and_teardown_local_testing():
|
||||
"""
|
||||
|
|
@ -8786,6 +8793,407 @@ def test_get_valid_models_from_provider():
|
|||
assert "gpt-5-mini" in valid_models
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("provider", "api_base", "api_key", "model_id"),
|
||||
[
|
||||
("anthropic", "https://anthropic.models.test", "anthropic-test-key", "claude-test-model"),
|
||||
("xai", "https://xai.models.test", "xai-test-key", "grok-test-model"),
|
||||
],
|
||||
)
|
||||
def test_get_valid_models_discovers_provider_models_from_http(
|
||||
provider: str,
|
||||
api_base: str,
|
||||
api_key: str,
|
||||
model_id: str,
|
||||
) -> None:
|
||||
response_body: Final = {
|
||||
"data": [
|
||||
{
|
||||
"id": model_id,
|
||||
"type": "model",
|
||||
"display_name": "Test model",
|
||||
"created_at": "2024-01-01T00:00:00Z",
|
||||
}
|
||||
],
|
||||
"has_more": False,
|
||||
"first_id": model_id,
|
||||
"last_id": model_id,
|
||||
}
|
||||
models_url: Final = f"{api_base}/v1/models"
|
||||
upstream: Final[respx.MockRouter]
|
||||
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
model_list_route: Final = upstream.get(models_url).respond(200, json=response_body)
|
||||
|
||||
discovered_models: Final = get_valid_models(
|
||||
check_provider_endpoint=True,
|
||||
custom_llm_provider=provider,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
assert discovered_models == [f"{provider}/{model_id}"]
|
||||
assert model_list_route.called
|
||||
assert len(upstream.calls) == 1
|
||||
|
||||
|
||||
def test_check_valid_key_returns_false_for_http_unauthorized() -> None:
|
||||
response_body: Final = {
|
||||
"error": {
|
||||
"message": "Invalid API key",
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": "invalid_api_key",
|
||||
}
|
||||
}
|
||||
upstream: Final[respx.MockRouter]
|
||||
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
invalid_key_route: Final = upstream.post("https://api.openai.com/v1/chat/completions").respond(
|
||||
401,
|
||||
json=response_body,
|
||||
)
|
||||
|
||||
valid_key: Final = check_valid_key(model="gpt-5-mini", api_key="invalid-test-key")
|
||||
|
||||
assert valid_key is False
|
||||
assert invalid_key_route.called
|
||||
assert upstream.calls.last.request.headers["Authorization"] == "Bearer invalid-test-key"
|
||||
|
||||
|
||||
def test_check_valid_key_returns_true_for_successful_completion() -> None:
|
||||
response_body: Final = {
|
||||
"id": "chatcmpl-test",
|
||||
"object": "chat.completion",
|
||||
"created": 0,
|
||||
"model": "gpt-5-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "ok"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
}
|
||||
upstream: Final[respx.MockRouter]
|
||||
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
valid_key_route: Final = upstream.post("https://api.openai.com/v1/chat/completions").respond(
|
||||
200,
|
||||
json=response_body,
|
||||
)
|
||||
|
||||
valid_key: Final = check_valid_key(model="gpt-5-mini", api_key="valid-test-key")
|
||||
|
||||
assert valid_key is True
|
||||
assert valid_key_route.called
|
||||
assert upstream.calls.last.request.headers["Authorization"] == "Bearer valid-test-key"
|
||||
|
||||
|
||||
def test_function_to_dict_parses_numpy_docstring_schema() -> None:
|
||||
pytest.importorskip("numpydoc")
|
||||
|
||||
def get_current_weather(location: str, unit: str) -> str:
|
||||
"""Get the current weather in a given location
|
||||
|
||||
Parameters
|
||||
----------
|
||||
location : str
|
||||
The city and state, e.g. San Francisco, CA
|
||||
unit : {'celsius', 'fahrenheit'}
|
||||
Temperature unit
|
||||
|
||||
Returns
|
||||
-------
|
||||
str
|
||||
A sentence indicating the weather
|
||||
"""
|
||||
return f"Weather for {location} in {unit}"
|
||||
|
||||
schema: Final = litellm.utils.function_to_dict(get_current_weather)
|
||||
|
||||
assert schema["name"] == "get_current_weather"
|
||||
assert schema["description"] == "Get the current weather in a given location"
|
||||
assert schema["parameters"]["type"] == "object"
|
||||
assert schema["parameters"]["properties"]["location"] == {
|
||||
"type": "string",
|
||||
"description": "The city and state, e.g. San Francisco, CA",
|
||||
}
|
||||
assert schema["parameters"]["properties"]["unit"]["type"] == "string"
|
||||
assert schema["parameters"]["properties"]["unit"]["description"] == "Temperature unit"
|
||||
assert schema["parameters"]["required"] == ["location", "unit"]
|
||||
|
||||
|
||||
def test_duration_in_seconds_one_month_uses_the_fixed_calendar_interval(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
fixed_start: Final = datetime(2025, 2, 15, 12, 0, 0, 123456)
|
||||
fixed_timestamp: Final = fixed_start.timestamp()
|
||||
monkeypatch.setattr(
|
||||
"litellm.litellm_core_utils.duration_parser.time_module.time",
|
||||
lambda: fixed_timestamp,
|
||||
)
|
||||
|
||||
assert duration_in_seconds("1mo") == 28 * 24 * 60 * 60
|
||||
|
||||
|
||||
def test_prompt_caching_image_check_uses_default_image_dimensions() -> None:
|
||||
image_bytes: Final = base64.b64decode(
|
||||
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/x+AAwMCAO+ip1sAAAAASUVORK5CYII="
|
||||
)
|
||||
image_url: Final = "https://93.184.216.34/test.png"
|
||||
messages: Final = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "What is in this image?"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": image_url, "detail": "high"},
|
||||
},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
image_route: Final = upstream.get(image_url).respond(200, content=image_bytes)
|
||||
cacheable: Final = is_prompt_caching_valid_prompt(
|
||||
model="gpt-4o-mini",
|
||||
messages=messages,
|
||||
custom_llm_provider="openai",
|
||||
min_token_count=100_000,
|
||||
)
|
||||
|
||||
assert cacheable is False
|
||||
assert image_route.called is False
|
||||
assert len(upstream.calls) == 0
|
||||
|
||||
|
||||
def test_get_valid_models_discovers_fireworks_models_from_http(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
api_base: Final = "https://fireworks.models.test/v1"
|
||||
api_key: Final = "fireworks-test-key"
|
||||
account_id: Final = "fireworks-test-account"
|
||||
model_name: Final = "accounts/fireworks/models/llama-test-model"
|
||||
models_url: Final = f"https://fireworks.models.test/v1/accounts/{account_id}/models"
|
||||
monkeypatch.setenv("FIREWORKS_ACCOUNT_ID", account_id)
|
||||
upstream: Final[respx.MockRouter]
|
||||
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
model_list_route: Final = upstream.get(models_url).respond(
|
||||
200,
|
||||
json={"models": [{"name": model_name}]},
|
||||
)
|
||||
|
||||
discovered_models: Final = get_valid_models(
|
||||
check_provider_endpoint=True,
|
||||
custom_llm_provider="fireworks_ai",
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
assert discovered_models == [f"fireworks_ai/{model_name}"]
|
||||
assert model_list_route.called
|
||||
assert upstream.calls.last.request.headers["Authorization"] == f"Bearer {api_key}"
|
||||
|
||||
|
||||
def test_get_valid_models_returns_static_fireworks_models_without_endpoint_check(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "check_provider_endpoint", False)
|
||||
monkeypatch.setenv("FIREWORKS_AI_API_KEY", "fireworks-test-key")
|
||||
expected_models: Final = litellm.models_by_provider["fireworks_ai"]
|
||||
|
||||
actual_models: Final = get_valid_models(
|
||||
check_provider_endpoint=False,
|
||||
custom_llm_provider="fireworks_ai",
|
||||
)
|
||||
env_inferred_models: Final = get_valid_models()
|
||||
|
||||
assert set(actual_models) == expected_models
|
||||
assert actual_models
|
||||
assert expected_models <= set(env_inferred_models)
|
||||
|
||||
|
||||
def test_get_valid_models_uses_the_litellm_params_anthropic_api_key() -> None:
|
||||
model_id: Final = "claude-test-model"
|
||||
models_url: Final = "https://api.anthropic.com/v1/models"
|
||||
response_body: Final = {
|
||||
"data": [
|
||||
{
|
||||
"id": model_id,
|
||||
"type": "model",
|
||||
"display_name": "Test Claude",
|
||||
"created_at": "2024-01-01T00:00:00Z",
|
||||
}
|
||||
],
|
||||
"has_more": False,
|
||||
"first_id": model_id,
|
||||
"last_id": model_id,
|
||||
}
|
||||
upstream: Final[respx.MockRouter]
|
||||
|
||||
def response_for_api_key(request: httpx.Request) -> httpx.Response:
|
||||
if request.headers["x-api-key"] == "bad-test-key":
|
||||
return httpx.Response(401, json={"error": {"message": "invalid key"}}, request=request)
|
||||
return httpx.Response(200, json=response_body, request=request)
|
||||
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
model_list_route: Final = upstream.get(models_url).mock(side_effect=response_for_api_key)
|
||||
|
||||
bad_key_models: Final = get_valid_models(
|
||||
check_provider_endpoint=True,
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=LiteLLM_Params(model="anthropic/*", api_key="bad-test-key"),
|
||||
)
|
||||
good_key_models: Final = get_valid_models(
|
||||
check_provider_endpoint=True,
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=LiteLLM_Params(model="anthropic/*", api_key="good-test-key"),
|
||||
)
|
||||
|
||||
assert bad_key_models == []
|
||||
assert good_key_models == [f"anthropic/{model_id}"]
|
||||
assert model_list_route.called
|
||||
assert [call.request.headers["x-api-key"] for call in upstream.calls] == [
|
||||
"bad-test-key",
|
||||
"good-test-key",
|
||||
]
|
||||
|
||||
|
||||
def test_add_custom_logger_to_success_callback_registers_once(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "success_callback", [])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
monkeypatch.setattr(litellm_logging, "_in_memory_loggers", [])
|
||||
|
||||
add_custom_logger_callback_to_specific_event("langfuse", "success")
|
||||
|
||||
assert len(litellm.success_callback) == 1
|
||||
assert isinstance(litellm.success_callback[0], LangfusePromptManagement)
|
||||
assert len(litellm._async_success_callback) == 1
|
||||
assert isinstance(litellm._async_success_callback[0], LangfusePromptManagement)
|
||||
assert litellm.failure_callback == []
|
||||
assert litellm._async_failure_callback == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"registered_lists",
|
||||
[
|
||||
("success_callback", "_async_success_callback"),
|
||||
("success_callback",),
|
||||
],
|
||||
)
|
||||
def test_add_custom_logger_callback_does_not_duplicate_existing_success_logger(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
registered_lists: tuple[str, ...],
|
||||
) -> None:
|
||||
logger: Final = HumanloopLogger()
|
||||
async_success_callbacks: Final = [logger] if "_async_success_callback" in registered_lists else []
|
||||
success_callbacks: Final = [logger] if "success_callback" in registered_lists else []
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", async_success_callbacks)
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "success_callback", success_callbacks)
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
monkeypatch.setattr(litellm_logging, "_in_memory_loggers", [])
|
||||
|
||||
add_custom_logger_callback_to_specific_event("humanloop", "success")
|
||||
|
||||
assert sum(type(callback) is HumanloopLogger for callback in litellm.success_callback) == int(
|
||||
"success_callback" in registered_lists
|
||||
)
|
||||
assert sum(type(callback) is HumanloopLogger for callback in litellm._async_success_callback) == int(
|
||||
"_async_success_callback" in registered_lists
|
||||
)
|
||||
assert litellm.failure_callback == []
|
||||
assert litellm._async_failure_callback == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("registered_lists", "expected_async_success_callback_count"),
|
||||
[
|
||||
(("success_callback", "_async_success_callback"), 1),
|
||||
(("success_callback",), 0),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_does_not_duplicate_a_logger_already_in_success_callbacks(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
registered_lists: tuple[str, ...],
|
||||
expected_async_success_callback_count: int,
|
||||
) -> None:
|
||||
logger: Final = HumanloopLogger()
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm, "input_callback", [])
|
||||
monkeypatch.setattr(litellm, "success_callback", [logger])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_input_callback", [])
|
||||
monkeypatch.setattr(
|
||||
litellm, "_async_success_callback", [logger] if "_async_success_callback" in registered_lists else []
|
||||
)
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
|
||||
await litellm.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "callback registration"}],
|
||||
mock_response="ok",
|
||||
)
|
||||
|
||||
assert litellm.success_callback == [logger]
|
||||
assert litellm._async_success_callback == [logger] * expected_async_success_callback_count
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_logger_in_global_callbacks_registers_once_across_completion_calls(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
logger: Final = HumanloopLogger()
|
||||
monkeypatch.setattr(litellm, "callbacks", [logger])
|
||||
monkeypatch.setattr(litellm, "input_callback", [])
|
||||
monkeypatch.setattr(litellm, "success_callback", [])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_input_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
|
||||
for _ in range(11):
|
||||
await litellm.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "callback registration"}],
|
||||
mock_response="ok",
|
||||
)
|
||||
|
||||
assert litellm.callbacks == [logger]
|
||||
assert litellm.input_callback == [logger]
|
||||
assert litellm.success_callback == [logger]
|
||||
assert litellm.failure_callback == [logger]
|
||||
assert litellm._async_input_callback == []
|
||||
assert litellm._async_success_callback == [logger]
|
||||
assert litellm._async_failure_callback == [logger]
|
||||
|
||||
|
||||
def test_get_provider_audio_transcription_config_resolves_for_every_provider() -> None:
|
||||
configs: Final = {
|
||||
provider: ProviderConfigManager.get_provider_audio_transcription_config(model="whisper-1", provider=provider)
|
||||
for provider in LlmProviders
|
||||
}
|
||||
unexpected: Final = {
|
||||
provider: config
|
||||
for provider, config in configs.items()
|
||||
if config is not None and not isinstance(config, BaseAudioTranscriptionConfig)
|
||||
}
|
||||
|
||||
assert unexpected == {}
|
||||
assert isinstance(configs[LlmProviders.OPENAI], litellm.OpenAIWhisperAudioTranscriptionConfig)
|
||||
|
||||
|
||||
def test_get_valid_models_from_provider_cache_invalidation(monkeypatch):
|
||||
"""
|
||||
Test that get_valid_models returns the correct models for a given provider
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue