From b3bd104f2446b8087fb003902cecfa28e58b3806 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 21 Dec 2024 14:17:12 -0800 Subject: [PATCH] (refactor) - fix from enterprise.utils import ui_get_spend_by_tags (#7352) * ui - refactor ui_get_spend_by_tags * fix typing --- enterprise/utils.py | 214 ------------------ .../spend_management_endpoints.py | 132 +++++++++-- 2 files changed, 115 insertions(+), 231 deletions(-) delete mode 100644 enterprise/utils.py diff --git a/enterprise/utils.py b/enterprise/utils.py deleted file mode 100644 index b252a064bb..0000000000 --- a/enterprise/utils.py +++ /dev/null @@ -1,214 +0,0 @@ -# Enterprise Proxy Util Endpoints -from typing import Optional, List -from litellm.proxy.proxy_server import PrismaClient, HTTPException -from litellm.llms.custom_httpx.http_handler import HTTPHandler -import collections -import httpx -from datetime import datetime - - -async def get_spend_by_tags( - prisma_client: PrismaClient, start_date=None, end_date=None -): - response = await prisma_client.db.query_raw( - """ - SELECT - jsonb_array_elements_text(request_tags) AS individual_request_tag, - COUNT(*) AS log_count, - SUM(spend) AS total_spend - FROM "LiteLLM_SpendLogs" - GROUP BY individual_request_tag; - """ - ) - - return response - - -async def ui_get_spend_by_tags( - start_date: str, - end_date: str, - prisma_client: Optional[PrismaClient] = None, - tags_str: Optional[str] = None, -): - """ - Should cover 2 cases: - 1. When user is getting spend for all_tags. "all_tags" in tags_list - 2. When user is getting spend for specific tags. - """ - - # tags_str is a list of strings csv of tags - # tags_str = tag1,tag2,tag3 - # convert to list if it's not None - tags_list: Optional[List[str]] = None - if tags_str is not None and len(tags_str) > 0: - tags_list = tags_str.split(",") - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No db connected"}) - - response = None - if tags_list is None or (isinstance(tags_list, list) and "all-tags" in tags_list): - # Get spend for all tags - sql_query = """ - SELECT - individual_request_tag, - spend_date, - log_count, - total_spend - FROM DailyTagSpend - WHERE spend_date >= $1::date AND spend_date <= $2::date - ORDER BY total_spend DESC; - """ - response = await prisma_client.db.query_raw( - sql_query, - start_date, - end_date, - ) - else: - # filter by tags list - sql_query = """ - SELECT - individual_request_tag, - SUM(log_count) AS log_count, - SUM(total_spend) AS total_spend - FROM DailyTagSpend - WHERE spend_date >= $1::date AND spend_date <= $2::date - AND individual_request_tag = ANY($3::text[]) - GROUP BY individual_request_tag - ORDER BY total_spend DESC; - """ - response = await prisma_client.db.query_raw( - sql_query, - start_date, - end_date, - tags_list, - ) - - # print("tags - spend") - # print(response) - # Bar Chart 1 - Spend per tag - Top 10 tags by spend - total_spend_per_tag: collections.defaultdict = collections.defaultdict(float) - total_requests_per_tag: collections.defaultdict = collections.defaultdict(int) - for row in response: - tag_name = row["individual_request_tag"] - tag_spend = row["total_spend"] - - total_spend_per_tag[tag_name] += tag_spend - total_requests_per_tag[tag_name] += row["log_count"] - - sorted_tags = sorted(total_spend_per_tag.items(), key=lambda x: x[1], reverse=True) - # convert to ui format - ui_tags = [] - for tag in sorted_tags: - current_spend = tag[1] - if current_spend is not None and isinstance(current_spend, float): - current_spend = round(current_spend, 4) - ui_tags.append( - { - "name": tag[0], - "spend": current_spend, - "log_count": total_requests_per_tag[tag[0]], - } - ) - - return {"spend_per_tag": ui_tags} - - -def _forecast_daily_cost(data: list): - from datetime import timedelta - - if len(data) == 0: - return { - "response": [], - "predicted_spend": "Current Spend = $0, Predicted = $0", - } - first_entry = data[0] - last_entry = data[-1] - - # get the date today - today_date = datetime.today().date() - - today_day_month = today_date.month - - # Parse the date from the first entry - first_entry_date = datetime.strptime(first_entry["date"], "%Y-%m-%d").date() - last_entry_date = datetime.strptime(last_entry["date"], "%Y-%m-%d") - - print("last entry date", last_entry_date) - - # Calculate the last day of the month - last_day_of_todays_month = datetime( - today_date.year, today_date.month % 12 + 1, 1 - ) - timedelta(days=1) - - print("last day of todays month", last_day_of_todays_month) - # Calculate the remaining days in the month - remaining_days = (last_day_of_todays_month - last_entry_date).days - - print("remaining days", remaining_days) - - current_spend_this_month = 0 - series = {} - for entry in data: - date = entry["date"] - spend = entry["spend"] - series[date] = spend - - # check if the date is in this month - if datetime.strptime(date, "%Y-%m-%d").month == today_day_month: - current_spend_this_month += spend - - if len(series) < 10: - num_items_to_fill = 11 - len(series) - - # avg spend for all days in series - avg_spend = sum(series.values()) / len(series) - for i in range(num_items_to_fill): - # go backwards from the first entry - date = first_entry_date - timedelta(days=i) - series[date.strftime("%Y-%m-%d")] = avg_spend - series[date.strftime("%Y-%m-%d")] = avg_spend - - payload = {"series": series, "count": remaining_days} - print("Prediction Data:", payload) - - headers = { - "Content-Type": "application/json", - } - - client = HTTPHandler() - - try: - response = client.post( - url="https://trend-api-production.up.railway.app/forecast", - json=payload, - headers=headers, - ) - except httpx.HTTPStatusError as e: - raise HTTPException( - status_code=500, - detail={"error": f"Error getting forecast: {e.response.text}"}, - ) - - json_response = response.json() - forecast_data = json_response["forecast"] - - # print("Forecast Data:", forecast_data) - - response_data = [] - total_predicted_spend = current_spend_this_month - for date in forecast_data: - spend = forecast_data[date] - entry = { - "date": date, - "predicted_spend": spend, - } - total_predicted_spend += spend - response_data.append(entry) - - # get month as a string, Jan, Feb, etc. - today_month = today_date.strftime("%B") - predicted_spend = ( - f"Predicted Spend for { today_month } 2024, ${total_predicted_spend}" - ) - return {"response": response_data, "predicted_spend": predicted_spend} diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 4eb78f4261..6af8593bd7 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1,10 +1,11 @@ #### SPEND MANAGEMENT ##### +import collections import os from datetime import datetime, timedelta -from typing import List, Optional +from typing import TYPE_CHECKING, Any, List, Optional import fastapi -from fastapi import APIRouter, Depends, HTTPException, Request, status +from fastapi import APIRouter, Depends, HTTPException, status import litellm from litellm._logging import verbose_proxy_logger @@ -16,6 +17,11 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import ( ) from litellm.proxy.utils import handle_exception_on_proxy +if TYPE_CHECKING: + from litellm.proxy.proxy_server import PrismaClient +else: + PrismaClient = Any + router = APIRouter() @@ -1354,7 +1360,6 @@ async def global_view_spend_tags( """ import traceback - from enterprise.utils import ui_get_spend_by_tags from litellm.proxy.proxy_server import prisma_client try: @@ -2451,20 +2456,6 @@ async def global_spend_models( return response -@router.post( - "/global/predict/spend/logs", - tags=["Budget & Spend Tracking"], - dependencies=[Depends(user_api_key_auth)], - include_in_schema=False, -) -async def global_predict_spend_logs(request: Request): - from enterprise.utils import _forecast_daily_cost - - data = await request.json() - data = data.get("data") - return _forecast_daily_cost(data) - - @router.get("/provider/budgets", response_model=ProviderBudgetResponse) async def provider_budgets() -> ProviderBudgetResponse: """ @@ -2554,3 +2545,110 @@ async def provider_budgets() -> ProviderBudgetResponse: "/provider/budgets: Exception occured - {}".format(str(e)) ) raise handle_exception_on_proxy(e) + + +async def get_spend_by_tags( + prisma_client: PrismaClient, start_date=None, end_date=None +): + response = await prisma_client.db.query_raw( + """ + SELECT + jsonb_array_elements_text(request_tags) AS individual_request_tag, + COUNT(*) AS log_count, + SUM(spend) AS total_spend + FROM "LiteLLM_SpendLogs" + GROUP BY individual_request_tag; + """ + ) + + return response + + +async def ui_get_spend_by_tags( + start_date: str, + end_date: str, + prisma_client: Optional[PrismaClient] = None, + tags_str: Optional[str] = None, +): + """ + Should cover 2 cases: + 1. When user is getting spend for all_tags. "all_tags" in tags_list + 2. When user is getting spend for specific tags. + """ + + # tags_str is a list of strings csv of tags + # tags_str = tag1,tag2,tag3 + # convert to list if it's not None + tags_list: Optional[List[str]] = None + if tags_str is not None and len(tags_str) > 0: + tags_list = tags_str.split(",") + + if prisma_client is None: + raise HTTPException(status_code=500, detail={"error": "No db connected"}) + + response = None + if tags_list is None or (isinstance(tags_list, list) and "all-tags" in tags_list): + # Get spend for all tags + sql_query = """ + SELECT + individual_request_tag, + spend_date, + log_count, + total_spend + FROM DailyTagSpend + WHERE spend_date >= $1::date AND spend_date <= $2::date + ORDER BY total_spend DESC; + """ + response = await prisma_client.db.query_raw( + sql_query, + start_date, + end_date, + ) + else: + # filter by tags list + sql_query = """ + SELECT + individual_request_tag, + SUM(log_count) AS log_count, + SUM(total_spend) AS total_spend + FROM DailyTagSpend + WHERE spend_date >= $1::date AND spend_date <= $2::date + AND individual_request_tag = ANY($3::text[]) + GROUP BY individual_request_tag + ORDER BY total_spend DESC; + """ + response = await prisma_client.db.query_raw( + sql_query, + start_date, + end_date, + tags_list, + ) + + # print("tags - spend") + # print(response) + # Bar Chart 1 - Spend per tag - Top 10 tags by spend + total_spend_per_tag: collections.defaultdict = collections.defaultdict(float) + total_requests_per_tag: collections.defaultdict = collections.defaultdict(int) + for row in response: + tag_name = row["individual_request_tag"] + tag_spend = row["total_spend"] + + total_spend_per_tag[tag_name] += tag_spend + total_requests_per_tag[tag_name] += row["log_count"] + + sorted_tags = sorted(total_spend_per_tag.items(), key=lambda x: x[1], reverse=True) + # convert to ui format + ui_tags = [] + for tag in sorted_tags: + current_spend = tag[1] + if current_spend is not None and isinstance(current_spend, float): + current_spend = round(current_spend, 4) + ui_tags.append( + { + "name": tag[0], + "spend": current_spend, + "log_count": total_requests_per_tag[tag[0]], + } + ) + + return {"spend_per_tag": ui_tags}