mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
* fix(ci): stop five stale or flaky CI reds and retry CyberArk policy-load conflicts The Langfuse redaction unit test exports to a local OTLP capture instead of polling Langfuse Cloud through a recorded lookup. The passthrough worker-kill test only requires spend rows for requests the surviving worker served. The spend-routes sweep treats the intentional /spend/capture_rate 503 as expected. CyberArk retries a 409 policy load in Python, Rust and the e2e Conjur helper instead of reading it as "variable exists". The integration egress guard now matches the script's own cgroup, so it no longer blocks the CircleCI agent, which runs as the same user. * fix(ci): keep the policy-load backoff typed as float * fix(ci): retry CyberArk policy loads without blocking the event loop and tighten the worker-kill and Langfuse tests * fix(secrets): load CyberArk policy one request at a time per manager * test(secrets): pin that non-conflict CyberArk policy failures are not retried * test(unit): run tests/unit with only an allowlisted host environment CircleCI's unit job inherits every project env var, so real provider keys, REDIS_HOST, DATABASE_URL and AWS or Azure credentials reached tests that assume none are set. Locally, litellm's import-time load_dotenv did the same from any .env up the tree. The unit conftest now drops every variable outside a small allowlist and disables dotenv before litellm is imported. * test(e2e): name a failed search and the stuck batch status instead of misattributing them The websearch session test read an empty web_search_tool_result_error block as a successful search, so a failing search tool surfaced as a session billing bug. The batch cancellation timeout now reports the last status the proxy returned. * fix(ci): scrub the host environment per unit test instead of for the whole pytest process GHA shards run tests/unit next to other suites in one process, so the import-time scrub deleted MCP_TEST_PEER_PYTHON before tests/mcp_tests read it and the MCP upstream fell back to the SDK2 interpreter. The two websearch tests that called OpenAI and Perplexity live are removed: tests/unit no longer sees their keys. * fix(ci): scrub only the host variables present before litellm is imported The per-test scrub also deleted TIKTOKEN_CACHE_DIR, which litellm sets at import to its bundled encodings, so tokenizer paths tried to download them and hit the socket guard. The prisma setup test now passes its own database URL instead of reading one another test leaked into the process environment. * fix(ci): stop the order-dependent unit reds and settle logging tasks on their own queue LoggingWorker marked a task done on whichever queue was current when the callback finished, so a callback that outlived an event-loop change raised "task_done() called too many times" or undercounted the new loop's queue. It now settles the queue the task came from. The rest are test isolation fixes for failures that only appeared when another file ran first on the same xdist worker: a replaced user_api_key_cache, breaker metrics unregistered by prometheus tests, semantic_router's health-check filter on uvicorn.access, logging tasks carried over from bedrock tests, a Router-written model_cost entry, and a stray post captured by the langflow test. The token counter check now asserts bounded chunking instead of wall-clock time. * test(e2e/ui): wait for the logout redirect before visiting a protected page Logout revokes the session server-side before clearing cookies and navigating, so an immediate page.goto either ran with the cookie still set or was aborted by the logout redirect (net::ERR_ABORTED). * test(unit): restore the prometheus metrics config per test and settle logs carried from earlier tests in the a2a cost tests * test(router): pin the router clock in the usage counter tests so a minute rollover cannot empty the read * test(e2e/ui): wait for logout to clear the token cookie instead of for a login redirect * test(integration/mcp): answer the model-info probe another test's proxy sends to the model double
451 lines
14 KiB
Python
451 lines
14 KiB
Python
"""
|
|
Test A2A cost calculator with cost_per_query parameter.
|
|
"""
|
|
|
|
import asyncio
|
|
from typing import Any, AsyncIterator, Optional
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import pytest
|
|
|
|
import litellm
|
|
from litellm.integrations.custom_logger import CustomLogger
|
|
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
|
|
|
|
|
async def _reset_callbacks_and_settle_pending_logs() -> None:
|
|
litellm.logging_callback_manager._reset_all_callbacks()
|
|
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0)
|
|
|
|
|
|
def _make_send_message_request(request_id: str, user_text: str = "Hello"):
|
|
from a2a.compat.v0_3.types import MessageSendParams, SendMessageRequest
|
|
|
|
return SendMessageRequest(
|
|
id=request_id,
|
|
params=MessageSendParams(
|
|
message={
|
|
"role": "user",
|
|
"parts": [{"kind": "text", "text": user_text}],
|
|
"messageId": "msg-1",
|
|
}
|
|
),
|
|
)
|
|
|
|
|
|
async def _mock_execute_a2a_send(
|
|
a2a_client: Any,
|
|
request: Any,
|
|
**kwargs: Any,
|
|
) -> Any:
|
|
mock_response = MagicMock()
|
|
mock_response.model_dump = MagicMock(
|
|
return_value={
|
|
"id": request.id,
|
|
"jsonrpc": "2.0",
|
|
"result": {"status": "completed"},
|
|
}
|
|
)
|
|
return mock_response
|
|
|
|
|
|
async def _mock_execute_a2a_send_with_assistant_reply(
|
|
a2a_client: Any,
|
|
request: Any,
|
|
**kwargs: Any,
|
|
) -> Any:
|
|
mock_response = MagicMock()
|
|
mock_response.model_dump = MagicMock(
|
|
return_value={
|
|
"id": request.id,
|
|
"jsonrpc": "2.0",
|
|
"result": {
|
|
"status": {"state": "completed"},
|
|
"message": {
|
|
"role": "assistant",
|
|
"parts": [
|
|
{
|
|
"kind": "text",
|
|
"text": "Hello! I am your assistant. How can I help you today?",
|
|
}
|
|
],
|
|
"messageId": "msg-456",
|
|
},
|
|
},
|
|
}
|
|
)
|
|
return mock_response
|
|
|
|
|
|
def _make_streaming_request(request_id: str):
|
|
from a2a.compat.v0_3.types import MessageSendParams, SendStreamingMessageRequest
|
|
|
|
return SendStreamingMessageRequest(
|
|
id=request_id,
|
|
params=MessageSendParams(
|
|
message={
|
|
"role": "user",
|
|
"parts": [{"kind": "text", "text": "Hello"}],
|
|
"messageId": "msg-1",
|
|
}
|
|
),
|
|
)
|
|
|
|
|
|
async def _mock_stream_messages(a2a_client: Any, request: Any) -> AsyncIterator[Any]:
|
|
from a2a.compat.v0_3.types import (
|
|
Message,
|
|
Part,
|
|
Role,
|
|
SendStreamingMessageResponse,
|
|
SendStreamingMessageSuccessResponse,
|
|
TextPart,
|
|
)
|
|
|
|
msg = Message(
|
|
message_id="msg-agent",
|
|
role=Role.agent,
|
|
parts=[Part(root=TextPart(kind="text", text="hello"))],
|
|
kind="message",
|
|
)
|
|
for _ in range(2):
|
|
yield SendStreamingMessageResponse(root=SendStreamingMessageSuccessResponse(id=request.id, result=msg))
|
|
|
|
|
|
class CostLogger(CustomLogger):
|
|
"""Custom logger to capture response_cost."""
|
|
|
|
def __init__(self):
|
|
self.response_cost: Optional[float] = None
|
|
super().__init__()
|
|
|
|
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
|
slp = kwargs.get("standard_logging_object")
|
|
if slp:
|
|
self.response_cost = (
|
|
slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None)
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_asend_message_uses_cost_per_query(monkeypatch):
|
|
"""
|
|
Test that asend_message uses cost_per_query param for response_cost.
|
|
"""
|
|
from litellm.a2a_protocol import asend_message
|
|
|
|
# Setup logger
|
|
await _reset_callbacks_and_settle_pending_logs()
|
|
cost_logger = CostLogger()
|
|
monkeypatch.setattr(litellm, "callbacks", [cost_logger])
|
|
|
|
# Mock A2A client
|
|
mock_client = MagicMock()
|
|
mock_client._litellm_agent_card = MagicMock()
|
|
mock_client._litellm_agent_card.name = "test-agent"
|
|
|
|
mock_request = _make_send_message_request("test-123")
|
|
|
|
# Call asend_message with cost_per_query
|
|
with patch(
|
|
"litellm.a2a_protocol.main._execute_a2a_send_with_retry",
|
|
new=_mock_execute_a2a_send,
|
|
):
|
|
await asend_message(
|
|
a2a_client=mock_client,
|
|
request=mock_request,
|
|
cost_per_query=0.05,
|
|
)
|
|
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert cost_logger.response_cost == 0.05
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_asend_message_uses_cost_per_query_from_litellm_params_dict(monkeypatch):
|
|
"""
|
|
Proxy passes agent pricing as the litellm_params dict param (not top-level
|
|
kwargs). Regression for cost_per_query landing at $0 on the native path.
|
|
"""
|
|
from litellm.a2a_protocol import asend_message
|
|
|
|
await _reset_callbacks_and_settle_pending_logs()
|
|
cost_logger = CostLogger()
|
|
monkeypatch.setattr(litellm, "callbacks", [cost_logger])
|
|
|
|
mock_client = MagicMock()
|
|
mock_client._litellm_agent_card = MagicMock()
|
|
mock_client._litellm_agent_card.name = "test-agent"
|
|
|
|
mock_request = _make_send_message_request("test-123")
|
|
|
|
with patch(
|
|
"litellm.a2a_protocol.main._execute_a2a_send_with_retry",
|
|
new=_mock_execute_a2a_send,
|
|
):
|
|
await asend_message(
|
|
a2a_client=mock_client,
|
|
request=mock_request,
|
|
litellm_params={
|
|
"cost_per_query": 0.5,
|
|
"input_cost_per_token": 0.099999,
|
|
"output_cost_per_token": 0.1,
|
|
},
|
|
)
|
|
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert cost_logger.response_cost == 0.5
|
|
|
|
|
|
class TokenAndCostLogger(CustomLogger):
|
|
"""Custom logger to capture both token counts and cost."""
|
|
|
|
def __init__(self):
|
|
self.response_cost: Optional[float] = None
|
|
self.prompt_tokens: Optional[int] = None
|
|
self.completion_tokens: Optional[int] = None
|
|
super().__init__()
|
|
|
|
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
|
slp = kwargs.get("standard_logging_object")
|
|
if slp:
|
|
self.response_cost = (
|
|
slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None)
|
|
)
|
|
self.prompt_tokens = (
|
|
slp.get("prompt_tokens") if isinstance(slp, dict) else getattr(slp, "prompt_tokens", None)
|
|
)
|
|
self.completion_tokens = (
|
|
slp.get("completion_tokens") if isinstance(slp, dict) else getattr(slp, "completion_tokens", None)
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_asend_message_uses_input_output_cost_per_token(monkeypatch):
|
|
"""
|
|
Test that asend_message calculates cost using input_cost_per_token and output_cost_per_token.
|
|
Validates exact cost calculation: cost = (prompt_tokens * input_cost) + (completion_tokens * output_cost)
|
|
"""
|
|
from litellm.a2a_protocol import asend_message
|
|
|
|
# Setup logger
|
|
await _reset_callbacks_and_settle_pending_logs()
|
|
token_cost_logger = TokenAndCostLogger()
|
|
monkeypatch.setattr(litellm, "callbacks", [token_cost_logger])
|
|
|
|
# Mock A2A client
|
|
mock_client = MagicMock()
|
|
mock_client._litellm_agent_card = MagicMock()
|
|
mock_client._litellm_agent_card.name = "test-agent"
|
|
|
|
mock_request = _make_send_message_request("test-123", user_text="Hello, what can you do?")
|
|
|
|
# Define specific cost per token values
|
|
input_cost_per_token = 0.00001 # $0.01 per 1000 tokens
|
|
output_cost_per_token = 0.00002 # $0.02 per 1000 tokens
|
|
|
|
with patch(
|
|
"litellm.a2a_protocol.main._execute_a2a_send_with_retry",
|
|
new=_mock_execute_a2a_send_with_assistant_reply,
|
|
):
|
|
await asend_message(
|
|
a2a_client=mock_client,
|
|
request=mock_request,
|
|
input_cost_per_token=input_cost_per_token,
|
|
output_cost_per_token=output_cost_per_token,
|
|
)
|
|
|
|
await asyncio.sleep(0.1)
|
|
|
|
# Get actual token counts from logger
|
|
prompt_tokens = token_cost_logger.prompt_tokens
|
|
completion_tokens = token_cost_logger.completion_tokens
|
|
response_cost = token_cost_logger.response_cost
|
|
|
|
print(f"\n=== Token-Based Cost Results ===")
|
|
print(f"prompt_tokens: {prompt_tokens}")
|
|
print(f"completion_tokens: {completion_tokens}")
|
|
print(f"input_cost_per_token: {input_cost_per_token}")
|
|
print(f"output_cost_per_token: {output_cost_per_token}")
|
|
print(f"response_cost: {response_cost}")
|
|
|
|
# Verify tokens were captured
|
|
assert prompt_tokens is not None, "prompt_tokens should be captured"
|
|
assert completion_tokens is not None, "completion_tokens should be captured"
|
|
assert response_cost is not None, "response_cost should be captured"
|
|
|
|
# Calculate expected cost
|
|
expected_cost = (prompt_tokens * input_cost_per_token) + (completion_tokens * output_cost_per_token)
|
|
print(f"expected_cost: {expected_cost}")
|
|
|
|
# Verify exact cost calculation
|
|
assert response_cost == expected_cost, f"response_cost {response_cost} should equal expected {expected_cost}"
|
|
|
|
|
|
class AgentIdLogger(CustomLogger):
|
|
"""Custom logger to capture agent_id from kwargs."""
|
|
|
|
def __init__(self):
|
|
self.agent_id: Optional[str] = None
|
|
self.kwargs: Optional[dict] = None
|
|
super().__init__()
|
|
|
|
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
|
self.kwargs = kwargs
|
|
self.agent_id = kwargs.get("agent_id")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_asend_message_passes_agent_id_to_callback(monkeypatch):
|
|
"""
|
|
Test that asend_message passes agent_id to callbacks via kwargs.
|
|
"""
|
|
from litellm.a2a_protocol import asend_message
|
|
|
|
# Setup logger
|
|
await _reset_callbacks_and_settle_pending_logs()
|
|
agent_id_logger = AgentIdLogger()
|
|
monkeypatch.setattr(litellm, "callbacks", [agent_id_logger])
|
|
|
|
# Mock A2A client
|
|
mock_client = MagicMock()
|
|
mock_client._litellm_agent_card = MagicMock()
|
|
mock_client._litellm_agent_card.name = "test-agent"
|
|
|
|
mock_request = _make_send_message_request("test-123")
|
|
|
|
test_agent_id = "agent-uuid-12345"
|
|
|
|
# Call asend_message with agent_id
|
|
with patch(
|
|
"litellm.a2a_protocol.main._execute_a2a_send_with_retry",
|
|
new=_mock_execute_a2a_send,
|
|
):
|
|
await asend_message(
|
|
a2a_client=mock_client,
|
|
request=mock_request,
|
|
agent_id=test_agent_id,
|
|
)
|
|
|
|
await asyncio.sleep(0.1)
|
|
|
|
# Verify agent_id was passed to callback
|
|
assert agent_id_logger.agent_id == test_agent_id, (
|
|
f"Expected agent_id '{test_agent_id}', got '{agent_id_logger.agent_id}'"
|
|
)
|
|
|
|
|
|
class MetadataLogger(CustomLogger):
|
|
"""Custom logger to capture metadata from kwargs for proxy spend tracking."""
|
|
|
|
def __init__(self):
|
|
self.metadata: Optional[dict] = None
|
|
self.litellm_params: Optional[dict] = None
|
|
self.user_api_key: Optional[str] = None
|
|
self.user_id: Optional[str] = None
|
|
self.team_id: Optional[str] = None
|
|
super().__init__()
|
|
|
|
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
|
self.litellm_params = kwargs.get("litellm_params", {})
|
|
self.metadata = self.litellm_params.get("metadata", {})
|
|
self.user_api_key = self.metadata.get("user_api_key")
|
|
self.user_id = self.metadata.get("user_api_key_user_id")
|
|
self.team_id = self.metadata.get("user_api_key_team_id")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_asend_message_streaming_propagates_metadata():
|
|
"""
|
|
Test that asend_message_streaming propagates metadata to logging object.
|
|
This ensures user_api_key, user_id, team_id are available for SpendLogs.
|
|
"""
|
|
from litellm.a2a_protocol import asend_message_streaming
|
|
|
|
# Setup logger
|
|
await _reset_callbacks_and_settle_pending_logs()
|
|
metadata_logger = MetadataLogger()
|
|
litellm.logging_callback_manager.add_litellm_async_success_callback(metadata_logger)
|
|
|
|
# Mock A2A client
|
|
mock_client = MagicMock()
|
|
mock_client._litellm_agent_card = MagicMock()
|
|
mock_client._litellm_agent_card.name = "test-agent"
|
|
|
|
mock_request = _make_streaming_request("test-stream-metadata")
|
|
|
|
# Metadata from proxy (contains user_api_key, user_id, team_id for SpendLogs)
|
|
test_metadata = {
|
|
"user_api_key": "sk-test-key-hash-12345",
|
|
"user_api_key_user_id": "user-uuid-123",
|
|
"user_api_key_team_id": "team-uuid-456",
|
|
}
|
|
|
|
# Consume streaming response with metadata
|
|
chunks = []
|
|
with patch(
|
|
"litellm.a2a_protocol.main._stream_messages",
|
|
new=_mock_stream_messages,
|
|
):
|
|
async for chunk in asend_message_streaming(
|
|
a2a_client=mock_client,
|
|
request=mock_request,
|
|
metadata=test_metadata,
|
|
):
|
|
chunks.append(chunk)
|
|
|
|
await asyncio.sleep(0.2)
|
|
|
|
# Verify metadata was propagated to callback
|
|
assert metadata_logger.user_api_key == "sk-test-key-hash-12345"
|
|
assert metadata_logger.user_id == "user-uuid-123"
|
|
assert metadata_logger.team_id == "team-uuid-456"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_asend_message_streaming_triggers_callbacks():
|
|
"""
|
|
Test that asend_message_streaming triggers callbacks after stream completes.
|
|
"""
|
|
from litellm.a2a_protocol import asend_message_streaming
|
|
|
|
# Setup logger - must use logging_callback_manager to properly register
|
|
await _reset_callbacks_and_settle_pending_logs()
|
|
callback_logger = AgentIdLogger()
|
|
litellm.logging_callback_manager.add_litellm_async_success_callback(callback_logger)
|
|
litellm.logging_callback_manager.add_litellm_success_callback(callback_logger)
|
|
|
|
# Mock A2A client
|
|
mock_client = MagicMock()
|
|
mock_client._litellm_agent_card = MagicMock()
|
|
mock_client._litellm_agent_card.name = "test-agent"
|
|
|
|
mock_request = _make_streaming_request("test-stream-123")
|
|
|
|
test_agent_id = "test-agent-id-streaming"
|
|
|
|
# Consume streaming response
|
|
chunks = []
|
|
with patch(
|
|
"litellm.a2a_protocol.main._stream_messages",
|
|
new=_mock_stream_messages,
|
|
):
|
|
async for chunk in asend_message_streaming(
|
|
a2a_client=mock_client,
|
|
request=mock_request,
|
|
agent_id=test_agent_id,
|
|
):
|
|
chunks.append(chunk)
|
|
|
|
await asyncio.sleep(0.2)
|
|
|
|
# Verify chunks were received
|
|
assert len(chunks) == 2
|
|
|
|
# Verify callbacks WERE triggered after stream completed
|
|
assert callback_logger.kwargs is not None, "Streaming should trigger callbacks after completion"
|
|
assert callback_logger.agent_id == test_agent_id, (
|
|
f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'"
|
|
)
|