From 29c215579699d532df8e85f282469a89fabcd2b4 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 21 Jun 2024 16:49:57 -0700 Subject: [PATCH 1/2] fix cost tracking by tags --- litellm/proxy/proxy_server.py | 4 +- .../spend_management_endpoints.py | 24 ++-- .../spend_tracking/spend_tracking_utils.py | 125 ++++++++++++++++++ litellm/proxy/utils.py | 115 ---------------- 4 files changed, 139 insertions(+), 129 deletions(-) rename litellm/proxy/{spend_reporting_endpoints => spend_tracking}/spend_management_endpoints.py (99%) create mode 100644 litellm/proxy/spend_tracking/spend_tracking_utils.py diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8eac72629ad..021b59e295f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -165,9 +165,10 @@ from litellm.proxy.secret_managers.aws_secret_manager import ( load_aws_secret_manager, ) from litellm.proxy.secret_managers.google_kms import load_google_kms -from litellm.proxy.spend_reporting_endpoints.spend_management_endpoints import ( +from litellm.proxy.spend_tracking.spend_management_endpoints import ( router as spend_management_router, ) +from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload from litellm.proxy.utils import ( DBClient, PrismaClient, @@ -180,7 +181,6 @@ from litellm.proxy.utils import ( encrypt_value, get_error_message_str, get_instance_fn, - get_logging_payload, hash_token, html_form, missing_keys_html_form, diff --git a/litellm/proxy/spend_reporting_endpoints/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py similarity index 99% rename from litellm/proxy/spend_reporting_endpoints/spend_management_endpoints.py rename to litellm/proxy/spend_tracking/spend_management_endpoints.py index 901a926456d..11edd188733 100644 --- a/litellm/proxy/spend_reporting_endpoints/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1,13 +1,14 @@ #### SPEND MANAGEMENT ##### -from typing import Optional, List +from datetime import datetime, timedelta, timezone +from typing import List, Optional + +import fastapi +from fastapi import APIRouter, Depends, Header, HTTPException, Request, status + import litellm from litellm._logging import verbose_proxy_logger -from datetime import datetime, timedelta, timezone -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -import fastapi -from fastapi import Depends, Request, APIRouter, Header, status -from fastapi import HTTPException from litellm.proxy._types import * +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth router = APIRouter() @@ -227,7 +228,7 @@ async def get_global_activity( start_date_obj = datetime.strptime(start_date, "%Y-%m-%d") end_date_obj = datetime.strptime(end_date, "%Y-%m-%d") - from litellm.proxy.proxy_server import prisma_client, llm_router + from litellm.proxy.proxy_server import llm_router, prisma_client try: if prisma_client is None: @@ -355,7 +356,7 @@ async def get_global_activity_model( start_date_obj = datetime.strptime(start_date, "%Y-%m-%d") end_date_obj = datetime.strptime(end_date, "%Y-%m-%d") - from litellm.proxy.proxy_server import prisma_client, llm_router, premium_user + from litellm.proxy.proxy_server import llm_router, premium_user, prisma_client try: if prisma_client is None: @@ -500,7 +501,7 @@ async def get_global_activity_exceptions_per_deployment( start_date_obj = datetime.strptime(start_date, "%Y-%m-%d") end_date_obj = datetime.strptime(end_date, "%Y-%m-%d") - from litellm.proxy.proxy_server import prisma_client, llm_router, premium_user + from litellm.proxy.proxy_server import llm_router, premium_user, prisma_client try: if prisma_client is None: @@ -634,7 +635,7 @@ async def get_global_activity_exceptions( start_date_obj = datetime.strptime(start_date, "%Y-%m-%d") end_date_obj = datetime.strptime(end_date, "%Y-%m-%d") - from litellm.proxy.proxy_server import prisma_client, llm_router + from litellm.proxy.proxy_server import llm_router, prisma_client try: if prisma_client is None: @@ -739,7 +740,7 @@ async def get_global_spend_provider( start_date_obj = datetime.strptime(start_date, "%Y-%m-%d") end_date_obj = datetime.strptime(end_date, "%Y-%m-%d") - from litellm.proxy.proxy_server import prisma_client, llm_router + from litellm.proxy.proxy_server import llm_router, prisma_client try: if prisma_client is None: @@ -1091,7 +1092,6 @@ async def global_view_spend_tags( """ from enterprise.utils import ui_get_spend_by_tags - from litellm.proxy.proxy_server import prisma_client try: diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py new file mode 100644 index 00000000000..e7bdac9aeea --- /dev/null +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -0,0 +1,125 @@ +import json +import traceback +from typing import Optional + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload +from litellm.proxy.utils import hash_token + + +def get_logging_payload( + kwargs, response_obj, start_time, end_time, end_user_id: Optional[str] +) -> SpendLogsPayload: + from pydantic import Json + + from litellm.proxy._types import LiteLLM_SpendLogs + + verbose_proxy_logger.debug( + f"SpendTable: get_logging_payload - kwargs: {kwargs}\n\n" + ) + + if kwargs is None: + kwargs = {} + # standardize this function to be used across, s3, dynamoDB, langfuse logging + litellm_params = kwargs.get("litellm_params", {}) + metadata = ( + litellm_params.get("metadata", {}) or {} + ) # if litellm_params['metadata'] == None + completion_start_time = kwargs.get("completion_start_time", end_time) + call_type = kwargs.get("call_type") + cache_hit = kwargs.get("cache_hit", False) + usage = response_obj["usage"] + if type(usage) == litellm.Usage: + usage = dict(usage) + id = response_obj.get("id", kwargs.get("litellm_call_id")) + api_key = metadata.get("user_api_key", "") + if api_key is not None and isinstance(api_key, str) and api_key.startswith("sk-"): + # hash the api_key + api_key = hash_token(api_key) + + _model_id = metadata.get("model_info", {}).get("id", "") + _model_group = metadata.get("model_group", "") + + request_tags = ( + json.dumps(metadata.get("tags", [])) + if isinstance(metadata.get("tags", []), dict) + else "[]" + ) + + # clean up litellm metadata + clean_metadata = SpendLogsMetadata( + user_api_key=None, + user_api_key_alias=None, + user_api_key_team_id=None, + user_api_key_user_id=None, + user_api_key_team_alias=None, + spend_logs_metadata=None, + ) + if isinstance(metadata, dict): + verbose_proxy_logger.debug( + "getting payload for SpendLogs, available keys in metadata: " + + str(list(metadata.keys())) + ) + + # Filter the metadata dictionary to include only the specified keys + clean_metadata = SpendLogsMetadata( + **{ # type: ignore + key: metadata[key] + for key in SpendLogsMetadata.__annotations__.keys() + if key in metadata + } + ) + + if litellm.cache is not None: + cache_key = litellm.cache.get_cache_key(**kwargs) + else: + cache_key = "Cache OFF" + if cache_hit is True: + import time + + id = f"{id}_cache_hit{time.time()}" # SpendLogs does not allow duplicate request_id + + try: + payload: SpendLogsPayload = SpendLogsPayload( + request_id=str(id), + call_type=call_type or "", + api_key=str(api_key), + cache_hit=str(cache_hit), + startTime=start_time, + endTime=end_time, + completionStartTime=completion_start_time, + model=kwargs.get("model", "") or "", + user=kwargs.get("litellm_params", {}) + .get("metadata", {}) + .get("user_api_key_user_id", "") + or "", + team_id=kwargs.get("litellm_params", {}) + .get("metadata", {}) + .get("user_api_key_team_id", "") + or "", + metadata=json.dumps(clean_metadata), + cache_key=cache_key, + spend=kwargs.get("response_cost", 0), + total_tokens=usage.get("total_tokens", 0), + prompt_tokens=usage.get("prompt_tokens", 0), + completion_tokens=usage.get("completion_tokens", 0), + request_tags=request_tags, + end_user=end_user_id or "", + api_base=litellm_params.get("api_base", ""), + model_group=_model_group, + model_id=_model_id, + ) + + verbose_proxy_logger.debug( + "SpendTable: created payload - payload: %s\n\n", payload + ) + + return payload + except Exception as e: + verbose_proxy_logger.error( + "Error creating spendlogs object - {}\n{}".format( + str(e), traceback.format_exc() + ) + ) + raise e diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d1e1d3576f4..1e75db213c0 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -2005,121 +2005,6 @@ def hash_token(token: str): return hashed_token -def get_logging_payload( - kwargs, response_obj, start_time, end_time, end_user_id: Optional[str] -) -> SpendLogsPayload: - from pydantic import Json - - from litellm.proxy._types import LiteLLM_SpendLogs - - verbose_proxy_logger.debug( - f"SpendTable: get_logging_payload - kwargs: {kwargs}\n\n" - ) - - if kwargs is None: - kwargs = {} - # standardize this function to be used across, s3, dynamoDB, langfuse logging - litellm_params = kwargs.get("litellm_params", {}) - metadata = ( - litellm_params.get("metadata", {}) or {} - ) # if litellm_params['metadata'] == None - completion_start_time = kwargs.get("completion_start_time", end_time) - call_type = kwargs.get("call_type") - cache_hit = kwargs.get("cache_hit", False) - usage = response_obj["usage"] - if type(usage) == litellm.Usage: - usage = dict(usage) - id = response_obj.get("id", kwargs.get("litellm_call_id")) - api_key = metadata.get("user_api_key", "") - if api_key is not None and isinstance(api_key, str) and api_key.startswith("sk-"): - # hash the api_key - api_key = hash_token(api_key) - - _model_id = metadata.get("model_info", {}).get("id", "") - _model_group = metadata.get("model_group", "") - - # clean up litellm metadata - clean_metadata = SpendLogsMetadata( - user_api_key=None, - user_api_key_alias=None, - user_api_key_team_id=None, - user_api_key_user_id=None, - user_api_key_team_alias=None, - spend_logs_metadata=None, - ) - if isinstance(metadata, dict): - verbose_proxy_logger.debug( - "getting payload for SpendLogs, available keys in metadata: " - + str(list(metadata.keys())) - ) - - # Filter the metadata dictionary to include only the specified keys - clean_metadata = SpendLogsMetadata( - **{ # type: ignore - key: metadata[key] - for key in SpendLogsMetadata.__annotations__.keys() - if key in metadata - } - ) - - if litellm.cache is not None: - cache_key = litellm.cache.get_cache_key(**kwargs) - else: - cache_key = "Cache OFF" - if cache_hit is True: - import time - - id = f"{id}_cache_hit{time.time()}" # SpendLogs does not allow duplicate request_id - - try: - payload: SpendLogsPayload = SpendLogsPayload( - request_id=str(id), - call_type=call_type or "", - api_key=str(api_key), - cache_hit=str(cache_hit), - startTime=start_time, - endTime=end_time, - completionStartTime=completion_start_time, - model=kwargs.get("model", "") or "", - user=kwargs.get("litellm_params", {}) - .get("metadata", {}) - .get("user_api_key_user_id", "") - or "", - team_id=kwargs.get("litellm_params", {}) - .get("metadata", {}) - .get("user_api_key_team_id", "") - or "", - metadata=json.dumps(clean_metadata), - cache_key=cache_key, - spend=kwargs.get("response_cost", 0), - total_tokens=usage.get("total_tokens", 0), - prompt_tokens=usage.get("prompt_tokens", 0), - completion_tokens=usage.get("completion_tokens", 0), - request_tags=( - json.dumps(metadata.get("tags", [])) - if isinstance(metadata.get("tags", []), dict) - else "[]" - ), - end_user=end_user_id or "", - api_base=litellm_params.get("api_base", ""), - model_group=_model_group, - model_id=_model_id, - ) - - verbose_proxy_logger.debug( - "SpendTable: created payload - payload: %s\n\n", payload - ) - - return payload - except Exception as e: - verbose_proxy_logger.error( - "Error creating spendlogs object - {}\n{}".format( - str(e), traceback.format_exc() - ) - ) - raise e - - def _extract_from_regex(duration: str) -> Tuple[int, str]: match = re.match(r"(\d+)(mo|[smhd]?)", duration) From fff928b10bb09b686f565058ee1012b00f60718f Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 21 Jun 2024 16:52:42 -0700 Subject: [PATCH 2/2] fix testing spend_tracking --- litellm/tests/test_blocked_user_list.py | 63 ++++++++++++----------- litellm/tests/test_key_generate_prisma.py | 2 +- litellm/tests/test_spend_logs.py | 34 ++++++++---- litellm/tests/test_update_spend.py | 59 ++++++++++----------- 4 files changed, 86 insertions(+), 72 deletions(-) diff --git a/litellm/tests/test_blocked_user_list.py b/litellm/tests/test_blocked_user_list.py index 3af1b246b27..12b4ab2f722 100644 --- a/litellm/tests/test_blocked_user_list.py +++ b/litellm/tests/test_blocked_user_list.py @@ -2,9 +2,14 @@ ## This tests the blocked user pre call hook for the proxy server -import sys, os, asyncio, time, random -from datetime import datetime +import asyncio +import os +import random +import sys +import time import traceback +from datetime import datetime + from dotenv import load_dotenv from fastapi import Request @@ -14,57 +19,53 @@ import os sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path +import asyncio +import logging + import pytest + import litellm +from litellm import Router, mock_completion +from litellm._logging import verbose_proxy_logger +from litellm.caching import DualCache +from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.enterprise.enterprise_hooks.blocked_user_list import ( _ENTERPRISE_BlockedUserList, ) -from litellm import Router, mock_completion -from litellm.proxy.utils import ProxyLogging, hash_token -from litellm.proxy._types import UserAPIKeyAuth -from litellm.caching import DualCache -from litellm.proxy.utils import PrismaClient, ProxyLogging, hash_token - -import pytest, logging, asyncio -import litellm, asyncio -from litellm.proxy.proxy_server import ( - user_api_key_auth, - block_user, +from litellm.proxy.management_endpoints.internal_user_endpoints import ( + new_user, + user_info, + user_update, ) from litellm.proxy.management_endpoints.key_management_endpoints import ( delete_key_fn, - info_key_fn, - update_key_fn, generate_key_fn, generate_key_helper_fn, + info_key_fn, + update_key_fn, ) -from litellm.proxy.management_endpoints.internal_user_endpoints import ( - new_user, - user_update, - user_info, -) -from litellm.proxy.spend_reporting_endpoints.spend_management_endpoints import ( - spend_user_fn, +from litellm.proxy.proxy_server import block_user, user_api_key_auth +from litellm.proxy.spend_tracking.spend_management_endpoints import ( spend_key_fn, + spend_user_fn, view_spend_logs, ) from litellm.proxy.utils import PrismaClient, ProxyLogging, hash_token -from litellm._logging import verbose_proxy_logger verbose_proxy_logger.setLevel(level=logging.DEBUG) +from starlette.datastructures import URL + +from litellm.caching import DualCache from litellm.proxy._types import ( - NewUserRequest, - GenerateKeyRequest, - DynamoDBArgs, - KeyRequest, - UpdateKeyRequest, - GenerateKeyRequest, BlockUsers, + DynamoDBArgs, + GenerateKeyRequest, + KeyRequest, + NewUserRequest, + UpdateKeyRequest, ) from litellm.proxy.utils import DBClient -from starlette.datastructures import URL -from litellm.caching import DualCache proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) diff --git a/litellm/tests/test_key_generate_prisma.py b/litellm/tests/test_key_generate_prisma.py index 594b4d77c55..5607d6c5b2c 100644 --- a/litellm/tests/test_key_generate_prisma.py +++ b/litellm/tests/test_key_generate_prisma.py @@ -75,7 +75,7 @@ from litellm.proxy.proxy_server import ( new_end_user, user_api_key_auth, ) -from litellm.proxy.spend_reporting_endpoints.spend_management_endpoints import ( +from litellm.proxy.spend_tracking.spend_management_endpoints import ( spend_key_fn, spend_user_fn, view_spend_logs, diff --git a/litellm/tests/test_spend_logs.py b/litellm/tests/test_spend_logs.py index b56bb5e2e40..3e8301e1e44 100644 --- a/litellm/tests/test_spend_logs.py +++ b/litellm/tests/test_spend_logs.py @@ -1,26 +1,32 @@ -import sys, os -import traceback, uuid +import os +import sys +import traceback +import uuid + from dotenv import load_dotenv from fastapi import Request from fastapi.routing import APIRoute load_dotenv() -import os, io, time +import io +import os +import time # this file is to test litellm/proxy sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path -import pytest, logging, asyncio -import litellm, asyncio -import json +import asyncio import datetime -from litellm.proxy.utils import ( - get_logging_payload, - SpendLogsPayload, - SpendLogsMetadata, -) # noqa: E402 +import json +import logging + +import pytest + +import litellm +from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload +from litellm.proxy.utils import SpendLogsMetadata, SpendLogsPayload # noqa: E402 def test_spend_logs_payload(): @@ -53,6 +59,7 @@ def test_spend_logs_payload(): "model_alias_map": {}, "completion_call_id": None, "metadata": { + "tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], "user_api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", "user_api_key_alias": None, "user_api_end_user_max_budget": None, @@ -193,3 +200,8 @@ def test_spend_logs_payload(): assert isinstance(payload["metadata"], str) payload["metadata"] = json.loads(payload["metadata"]) assert set(payload["metadata"].keys()) == set(expected_metadata_keys) + + # This is crucial - used in PROD, it should pass, related issue: https://github.com/BerriAI/litellm/issues/4334 + assert ( + payload["request_tags"] == '["model-anthropic-claude-v2.1", "app-ishaan-prod"]' + ) diff --git a/litellm/tests/test_update_spend.py b/litellm/tests/test_update_spend.py index c0bdd5cf910..fe06229ca26 100644 --- a/litellm/tests/test_update_spend.py +++ b/litellm/tests/test_update_spend.py @@ -2,9 +2,14 @@ ## This tests the batch update spend logic on the proxy server -import sys, os, asyncio, time, random -from datetime import datetime +import asyncio +import os +import random +import sys +import time import traceback +from datetime import datetime + from dotenv import load_dotenv from fastapi import Request @@ -14,54 +19,50 @@ import os sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path +import asyncio +import logging + import pytest + import litellm from litellm import Router, mock_completion -from litellm.proxy.utils import ProxyLogging -from litellm.proxy._types import UserAPIKeyAuth +from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache -from litellm.proxy.utils import PrismaClient, ProxyLogging, hash_token - -import pytest, logging, asyncio -import litellm, asyncio -from litellm.proxy.proxy_server import ( - user_api_key_auth, - block_user, -) -from litellm.proxy.spend_reporting_endpoints.spend_management_endpoints import ( - spend_user_fn, - spend_key_fn, - view_spend_logs, -) +from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.management_endpoints.internal_user_endpoints import ( new_user, - user_update, user_info, + user_update, ) from litellm.proxy.management_endpoints.key_management_endpoints import ( delete_key_fn, - info_key_fn, - update_key_fn, generate_key_fn, generate_key_helper_fn, + info_key_fn, + update_key_fn, +) +from litellm.proxy.proxy_server import block_user, user_api_key_auth +from litellm.proxy.spend_tracking.spend_management_endpoints import ( + spend_key_fn, + spend_user_fn, + view_spend_logs, ) from litellm.proxy.utils import PrismaClient, ProxyLogging, hash_token, update_spend -from litellm._logging import verbose_proxy_logger verbose_proxy_logger.setLevel(level=logging.DEBUG) +from starlette.datastructures import URL + +from litellm.caching import DualCache from litellm.proxy._types import ( - NewUserRequest, - GenerateKeyRequest, - DynamoDBArgs, - KeyRequest, - UpdateKeyRequest, - GenerateKeyRequest, BlockUsers, + DynamoDBArgs, + GenerateKeyRequest, + KeyRequest, + NewUserRequest, + UpdateKeyRequest, ) from litellm.proxy.utils import DBClient -from starlette.datastructures import URL -from litellm.caching import DualCache proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())