litellm/tests/unit/a2a_protocol/test_cost_calculator.py
yuneng-jiang 14f4c34c61
fix(ci): stop stale CI reds, keep unit tests off the host env, retry CyberArk policy conflicts (#43294)
* 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
2026-09-26 09:25:13 -07:00

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}'"
)