From 98365205acd08eaeaa870f389c490293424dabaa Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 17 Nov 2025 23:00:53 -0800 Subject: [PATCH] Deduplicate /tag/daily/activity metadata --- litellm/proxy/db/db_spend_update_writer.py | 3 + .../common_daily_activity.py | 59 ++++++++++--- .../tag_management_endpoints.py | 2 + .../proxy/db/test_db_spend_update_writer.py | 55 +++++++++++++ .../test_common_daily_activity.py | 82 +++++++++++++++++-- 5 files changed, 184 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 06b5301424..20915983d4 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1107,6 +1107,9 @@ class DBSpendUpdateWriter: ) } + if entity_type == "tag" and "request_id" in transaction: + update_data["request_id"] = transaction.get("request_id") + table.upsert( where=where_clause, data={ diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index dbf1cdf514..cd28cbb714 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -1,5 +1,5 @@ from datetime import datetime -from typing import Any, Dict, List, Optional, Set, Union +from typing import Any, Callable, Dict, List, Optional, Set, Union from fastapi import HTTPException, status @@ -32,6 +32,40 @@ def update_metrics(existing_metrics: SpendMetrics, record: Any) -> SpendMetrics: return existing_metrics +def _is_user_agent_tag(tag: Optional[str]) -> bool: + """Determine whether a tag should be treated as a User-Agent tag.""" + if not tag: + return False + normalized_tag = tag.strip().lower() + return normalized_tag.startswith("user-agent:") or normalized_tag.startswith("user agent:") + + +def compute_tag_metadata_totals(records: List[Any]) -> SpendMetrics: + """ + Deduplicate spend metrics for tags using request_id, ignoring User-Agent prefixed tags. + + Each unique request_id contributes at most one record (the tag with max spend) to metadata. + """ + deduped_records: Dict[str, Any] = {} + for record in records: + request_id = getattr(record, "request_id", None) + if not request_id: + continue + + tag_value = getattr(record, "tag", None) + if _is_user_agent_tag(tag_value): + continue + + current_best = deduped_records.get(request_id) + if current_best is None or record.spend > current_best.spend: + deduped_records[request_id] = record + + metadata_metrics = SpendMetrics() + for record in deduped_records.values(): + update_metrics(metadata_metrics, record) + return metadata_metrics + + def update_breakdown_metrics( breakdown: BreakdownMetrics, record: Any, @@ -380,6 +414,7 @@ async def get_daily_activity( page: int, page_size: int, exclude_entity_ids: Optional[List[str]] = None, + metadata_metrics_func: Optional[Callable[[List[Any]], SpendMetrics]] = None, ) -> SpendAnalyticsPaginatedResponse: """Common function to get daily activity for any entity type.""" @@ -428,18 +463,22 @@ async def get_daily_activity( entity_metadata_field=entity_metadata_field, ) + metadata_metrics = aggregated["totals"] + if metadata_metrics_func: + metadata_metrics = metadata_metrics_func(daily_spend_data) + return SpendAnalyticsPaginatedResponse( results=aggregated["results"], metadata=DailySpendMetadata( - total_spend=aggregated["totals"].spend, - total_prompt_tokens=aggregated["totals"].prompt_tokens, - total_completion_tokens=aggregated["totals"].completion_tokens, - total_tokens=aggregated["totals"].total_tokens, - total_api_requests=aggregated["totals"].api_requests, - total_successful_requests=aggregated["totals"].successful_requests, - total_failed_requests=aggregated["totals"].failed_requests, - total_cache_read_input_tokens=aggregated["totals"].cache_read_input_tokens, - total_cache_creation_input_tokens=aggregated["totals"].cache_creation_input_tokens, + total_spend=metadata_metrics.spend, + total_prompt_tokens=metadata_metrics.prompt_tokens, + total_completion_tokens=metadata_metrics.completion_tokens, + total_tokens=metadata_metrics.total_tokens, + total_api_requests=metadata_metrics.api_requests, + total_successful_requests=metadata_metrics.successful_requests, + total_failed_requests=metadata_metrics.failed_requests, + total_cache_read_input_tokens=metadata_metrics.cache_read_input_tokens, + total_cache_creation_input_tokens=metadata_metrics.cache_creation_input_tokens, page=page, total_pages=-(-total_count // page_size), # Ceiling division has_more=(page * page_size) < total_count, diff --git a/litellm/proxy/management_endpoints/tag_management_endpoints.py b/litellm/proxy/management_endpoints/tag_management_endpoints.py index 1366c2ef4e..f292ffd52b 100644 --- a/litellm/proxy/management_endpoints/tag_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tag_management_endpoints.py @@ -22,6 +22,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.management_endpoints.common_daily_activity import ( SpendAnalyticsPaginatedResponse, + compute_tag_metadata_totals, get_daily_activity, ) from litellm.proxy.management_helpers.utils import handle_budget_for_entity @@ -533,4 +534,5 @@ async def get_tag_daily_activity( api_key=api_key, page=page, page_size=page_size, + metadata_metrics_func=compute_tag_metadata_totals, ) diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 6dbbbdd744..8a5860165a 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -221,6 +221,61 @@ async def test_update_daily_spend_sorting(): # Verify that table.upsert was called mock_table.upsert.assert_has_calls(upsert_calls) + + +@pytest.mark.asyncio +async def test_update_daily_spend_tag_with_request_id(): + """ + Test that request_id is included in update_data when updating tag transactions. + """ + # Setup + mock_prisma_client = MagicMock() + mock_batcher = MagicMock() + mock_table = MagicMock() + mock_prisma_client.db.batch_.return_value.__aenter__.return_value = mock_batcher + mock_batcher.litellm_dailytagspend = mock_table + + # Create a transaction with request_id + daily_spend_transactions = { + "test_key": { + "tag": "prod-tag", + "date": "2024-01-01", + "api_key": "test-api-key", + "model": "gpt-4", + "custom_llm_provider": "openai", + "mcp_namespaced_tool_name": "", + "prompt_tokens": 10, + "completion_tokens": 20, + "spend": 0.1, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + "request_id": "test-request-id-123", + } + } + + # Call the method + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=1, + prisma_client=mock_prisma_client, + proxy_logging_obj=MagicMock(), + daily_spend_transactions=daily_spend_transactions, + entity_type="tag", + entity_id_field="tag", + table_name="litellm_dailytagspend", + unique_constraint_name="tag_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name", + ) + + # Verify that table.upsert was called + mock_table.upsert.assert_called_once() + + # Verify request_id is in update_data + call_args = mock_table.upsert.call_args[1] + update_data = call_args["data"]["update"] + assert "request_id" in update_data + assert update_data["request_id"] == "test-request-id-123" + + # Tag Spend Tracking Tests diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index ffaed2d88f..bbdc4b1edf 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -1,20 +1,18 @@ -import json import os import sys -from datetime import datetime, timezone from unittest.mock import AsyncMock, MagicMock import pytest -from fastapi.testclient import TestClient sys.path.insert( 0, os.path.abspath("../../../..") ) # Adds the parent directory to the system path -from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity -from litellm.proxy.proxy_server import app - -client = TestClient(app) +from litellm.proxy.management_endpoints.common_daily_activity import ( + _is_user_agent_tag, + compute_tag_metadata_totals, + get_daily_activity, +) @pytest.mark.asyncio @@ -56,3 +54,73 @@ async def test_get_daily_activity_empty_entity_id_list(): # Check that team_id is set to empty list assert "team_id" in where_conditions assert where_conditions["team_id"] == {"in": []} + + +def test_is_user_agent_tag(): + """Test _is_user_agent_tag function.""" + # Test None and empty string + assert _is_user_agent_tag(None) is False + assert _is_user_agent_tag("") is False + + # Test user-agent variations (should return True) + assert _is_user_agent_tag("user-agent:chrome") is True + assert _is_user_agent_tag("user agent:firefox") is True + assert _is_user_agent_tag("USER-AGENT:safari") is True + assert _is_user_agent_tag("User Agent:edge") is True + assert _is_user_agent_tag(" user-agent:opera ") is True # with whitespace + + # Test regular tags (should return False) + assert _is_user_agent_tag("production") is False + assert _is_user_agent_tag("tag:value") is False + assert _is_user_agent_tag("user-agent-tag") is False # no colon + + +def test_compute_tag_metadata_totals(): + """Test compute_tag_metadata_totals function.""" + # Create mock records + class MockRecord: + def __init__(self, request_id, tag, spend, prompt_tokens=10, completion_tokens=5): + self.request_id = request_id + self.tag = tag + self.spend = spend + self.prompt_tokens = prompt_tokens + self.completion_tokens = completion_tokens + self.total_tokens = prompt_tokens + completion_tokens + self.cache_read_input_tokens = 0 + self.cache_creation_input_tokens = 0 + self.api_requests = 1 + self.successful_requests = 1 + self.failed_requests = 0 + + # Test deduplication by request_id (keeps max spend) + records = [ + MockRecord("req-1", "production", spend=10.0), + MockRecord("req-1", "staging", spend=20.0), # Higher spend, should be kept + MockRecord("req-2", "production", spend=15.0), + ] + result = compute_tag_metadata_totals(records) + assert result.spend == 35.0 # 20.0 + 15.0 (deduplicated req-1) + assert result.prompt_tokens == 20 # 10 + 10 (only deduplicated records) + assert result.completion_tokens == 10 # 5 + 5 (only deduplicated records) + + # Test ignoring user-agent tags + records_with_ua = [ + MockRecord("req-1", "production", spend=10.0), + MockRecord("req-1", "user-agent:chrome", spend=50.0), # Should be ignored + MockRecord("req-2", "staging", spend=15.0), + ] + result = compute_tag_metadata_totals(records_with_ua) + assert result.spend == 25.0 # 10.0 + 15.0 (user-agent ignored) + + # Test ignoring records without request_id + records_no_req_id = [ + MockRecord("req-1", "production", spend=10.0), + MockRecord(None, "staging", spend=20.0), # Should be ignored + ] + result = compute_tag_metadata_totals(records_no_req_id) + assert result.spend == 10.0 + + # Test empty records + result = compute_tag_metadata_totals([]) + assert result.spend == 0.0 + assert result.prompt_tokens == 0