From e1e1e2e566ca4075167b35b38f48b38c2aa99d3a Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 30 Aug 2024 15:46:45 -0700 Subject: [PATCH 01/14] add example custom --- .../proxy/example_config_yaml/custom_auth.py | 50 +++---------------- litellm/proxy/proxy_config.yaml | 3 +- 2 files changed, 9 insertions(+), 44 deletions(-) diff --git a/litellm/proxy/example_config_yaml/custom_auth.py b/litellm/proxy/example_config_yaml/custom_auth.py index 6cecf466c10..2726b1c3d16 100644 --- a/litellm/proxy/example_config_yaml/custom_auth.py +++ b/litellm/proxy/example_config_yaml/custom_auth.py @@ -1,50 +1,14 @@ -from litellm.proxy._types import UserAPIKeyAuth, GenerateKeyRequest from fastapi import Request -import os + +from litellm.proxy._types import UserAPIKeyAuth async def user_api_key_auth(request: Request, api_key: str) -> UserAPIKeyAuth: try: - modified_master_key = f"{os.getenv('PROXY_MASTER_KEY')}-1234" - if api_key == modified_master_key: - return UserAPIKeyAuth(api_key=api_key) - raise Exception + return UserAPIKeyAuth( + api_key="best-api-key-ever", + user_id="best-user-id-ever", + team_id="best-team-id-ever", + ) except: raise Exception - - -async def generate_key_fn(data: GenerateKeyRequest): - """ - Asynchronously decides if a key should be generated or not based on the provided data. - - Args: - data (GenerateKeyRequest): The data to be used for decision making. - - Returns: - bool: True if a key should be generated, False otherwise. - """ - # decide if a key should be generated or not - data_json = data.json() # type: ignore - - # Unpacking variables - team_id = data_json.get("team_id") - duration = data_json.get("duration") - models = data_json.get("models") - aliases = data_json.get("aliases") - config = data_json.get("config") - spend = data_json.get("spend") - user_id = data_json.get("user_id") - max_parallel_requests = data_json.get("max_parallel_requests") - metadata = data_json.get("metadata") - tpm_limit = data_json.get("tpm_limit") - rpm_limit = data_json.get("rpm_limit") - - if team_id is not None and len(team_id) > 0: - return { - "decision": True, - } - else: - return { - "decision": True, - "message": "This violates LiteLLM Proxy Rules. No team id provided.", - } diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index c5f736bacce..708f0e27ca3 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -20,4 +20,5 @@ router_settings: enable_tag_filtering: True # 👈 Key Change general_settings: - master_key: sk-1234 \ No newline at end of file + master_key: sk-1234 + custom_auth: example_config_yaml.custom_auth.user_api_key_auth \ No newline at end of file From f50374e81dbe3b1ff3b3269df25ac1811cc56ef0 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 30 Aug 2024 15:52:47 -0700 Subject: [PATCH 02/14] use helper class for pass through success handler --- .../pass_through_endpoints.py | 15 ++- .../pass_through_endpoints/success_handler.py | 105 ++++++++++++++++++ 2 files changed, 117 insertions(+), 3 deletions(-) create mode 100644 litellm/proxy/pass_through_endpoints/success_handler.py diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index b9ab7526a3c..f34efdcf390 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -3,6 +3,7 @@ import asyncio import json import traceback from base64 import b64encode +from datetime import datetime from typing import AsyncIterable, List, Optional import httpx @@ -20,6 +21,7 @@ from fastapi.responses import StreamingResponse import litellm from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._types import ( ConfigFieldInfo, ConfigFieldUpdate, @@ -30,8 +32,12 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from .success_handler import PassThroughEndpointLogging + router = APIRouter() +pass_through_endpoint_logging = PassThroughEndpointLogging() + async def set_env_variables_in_header(custom_headers: dict): """ @@ -330,7 +336,7 @@ async def pass_through_request( async_client = httpx.AsyncClient(timeout=600) # create logging object - start_time = time.time() + start_time = datetime.now() logging_obj = Logging( model="unknown", messages=[{"role": "user", "content": "no-message-pass-through-endpoint"}], @@ -473,12 +479,15 @@ async def pass_through_request( content = await response.aread() ## LOG SUCCESS - end_time = time.time() + end_time = datetime.now() - await logging_obj.async_success_handler( + await pass_through_endpoint_logging.pass_through_async_success_handler( + httpx_response=response, + url_route=str(url), result="", start_time=start_time, end_time=end_time, + logging_obj=logging_obj, cache_hit=False, ) diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py new file mode 100644 index 00000000000..618f68659ea --- /dev/null +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -0,0 +1,105 @@ +import re +from datetime import datetime + +import httpx + +import litellm +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.vertex_ai_and_google_ai_studio.gemini.vertex_and_google_ai_studio_gemini import ( + VertexLLM, +) +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + +class PassThroughEndpointLogging: + def __init__(self): + self.TRACKED_VERTEX_ROUTES = [ + "generateContent", + "streamGenerateContent", + "predict", + ] + + async def pass_through_async_success_handler( + self, + httpx_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + url_route: str, + result: str, + start_time: datetime, + end_time: datetime, + cache_hit: bool, + **kwargs, + ): + if self.is_vertex_route(url_route): + await self.vertex_passthrough_handler( + httpx_response=httpx_response, + logging_obj=logging_obj, + url_route=url_route, + result=result, + start_time=start_time, + end_time=end_time, + cache_hit=cache_hit, + **kwargs, + ) + else: + await logging_obj.async_success_handler( + result="", + start_time=start_time, + end_time=end_time, + cache_hit=False, + ) + + def is_vertex_route(self, url_route: str): + for route in self.TRACKED_VERTEX_ROUTES: + if route in url_route: + return True + return False + + def extract_model_from_url(self, url: str) -> str: + pattern = r"/models/([^:]+)" + match = re.search(pattern, url) + if match: + return match.group(1) + return "unknown" + + async def vertex_passthrough_handler( + self, + httpx_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + url_route: str, + result: str, + start_time: datetime, + end_time: datetime, + cache_hit: bool, + **kwargs, + ): + if "generateContent" in url_route: + model = self.extract_model_from_url(url_route) + + instance_of_vertex_llm = VertexLLM() + litellm_model_response: litellm.ModelResponse = ( + instance_of_vertex_llm._process_response( + model=model, + messages=[ + {"role": "user", "content": "no-message-pass-through-endpoint"} + ], + response=httpx_response, + model_response=litellm.ModelResponse(), + logging_obj=logging_obj, + optional_params={}, + litellm_params={}, + api_key="", + data={}, + print_verbose=litellm.print_verbose, + encoding=None, + ) + ) + logging_obj.model = litellm_model_response.model + logging_obj.model_call_details["model"] = logging_obj.model + + await logging_obj.async_success_handler( + result=litellm_model_response, + start_time=start_time, + end_time=end_time, + cache_hit=cache_hit, + ) From 8ed0ffea541374886f5b85a1619d4572ffff7b49 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 30 Aug 2024 16:22:28 -0700 Subject: [PATCH 03/14] fix use existing custom_auth.py --- .../proxy/example_config_yaml/custom_auth.py | 50 ++++++++++++++++--- 1 file changed, 44 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/example_config_yaml/custom_auth.py b/litellm/proxy/example_config_yaml/custom_auth.py index 2726b1c3d16..66c213dcfd6 100644 --- a/litellm/proxy/example_config_yaml/custom_auth.py +++ b/litellm/proxy/example_config_yaml/custom_auth.py @@ -1,14 +1,52 @@ +import os + from fastapi import Request -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import GenerateKeyRequest, UserAPIKeyAuth async def user_api_key_auth(request: Request, api_key: str) -> UserAPIKeyAuth: try: - return UserAPIKeyAuth( - api_key="best-api-key-ever", - user_id="best-user-id-ever", - team_id="best-team-id-ever", - ) + modified_master_key = f"{os.getenv('PROXY_MASTER_KEY')}-1234" + if api_key == modified_master_key: + return UserAPIKeyAuth(api_key=api_key) + raise Exception except: raise Exception + + +async def generate_key_fn(data: GenerateKeyRequest): + """ + Asynchronously decides if a key should be generated or not based on the provided data. + + Args: + data (GenerateKeyRequest): The data to be used for decision making. + + Returns: + bool: True if a key should be generated, False otherwise. + """ + # decide if a key should be generated or not + data_json = data.json() # type: ignore + + # Unpacking variables + team_id = data_json.get("team_id") + duration = data_json.get("duration") + models = data_json.get("models") + aliases = data_json.get("aliases") + config = data_json.get("config") + spend = data_json.get("spend") + user_id = data_json.get("user_id") + max_parallel_requests = data_json.get("max_parallel_requests") + metadata = data_json.get("metadata") + tpm_limit = data_json.get("tpm_limit") + rpm_limit = data_json.get("rpm_limit") + + if team_id is not None and len(team_id) > 0: + return { + "decision": True, + } + else: + return { + "decision": True, + "message": "This violates LiteLLM Proxy Rules. No team id provided.", + } From f3f85f61416efb868f8fba201311078294499c42 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 30 Aug 2024 16:26:00 -0700 Subject: [PATCH 04/14] add test for vertex basic pass throgh --- .circleci/config.yml | 101 ++++++++++++++++++ .../example_config_yaml/custom_auth_basic.py | 14 +++ .../pass_through_config.yaml | 9 ++ tests/pass_through_tests/test_vertex_ai.py | 77 +++++++++++++ tests/pass_through_tests/vertex_key.json | 13 +++ 5 files changed, 214 insertions(+) create mode 100644 litellm/proxy/example_config_yaml/custom_auth_basic.py create mode 100644 litellm/proxy/example_config_yaml/pass_through_config.yaml create mode 100644 tests/pass_through_tests/test_vertex_ai.py create mode 100644 tests/pass_through_tests/vertex_key.json diff --git a/.circleci/config.yml b/.circleci/config.yml index 585502710f4..efc0c720cbc 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -363,6 +363,100 @@ jobs: - store_test_results: path: test-results + proxy_pass_through_endpoint_tests: + machine: + image: ubuntu-2204:2023.10.1 + resource_class: xlarge + working_directory: ~/project + steps: + - checkout + - run: + name: Install Docker CLI (In case it's not already installed) + command: | + sudo apt-get update + sudo apt-get install -y docker-ce docker-ce-cli containerd.io + - run: + name: Install Python 3.9 + command: | + curl https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh --output miniconda.sh + bash miniconda.sh -b -p $HOME/miniconda + export PATH="$HOME/miniconda/bin:$PATH" + conda init bash + source ~/.bashrc + conda create -n myenv python=3.9 -y + conda activate myenv + python --version + - run: + name: Install Dependencies + command: | + pip install "pytest==7.3.1" + pip install "pytest-retry==1.6.3" + pip install "pytest-asyncio==0.21.1" + pip install "google-cloud-aiplatform==1.43.0" + pip install aiohttp + pip install "openai==1.40.0" + python -m pip install --upgrade pip + pip install "pydantic==2.7.1" + pip install "pytest==7.3.1" + pip install "pytest-mock==3.12.0" + pip install "pytest-asyncio==0.21.1" + pip install "boto3==1.34.34" + pip install mypy + pip install pyarrow + pip install numpydoc + pip install prisma + pip install fastapi + pip install jsonschema + pip install "httpx==0.24.1" + pip install "anyio==3.7.1" + pip install "asyncio==3.4.3" + pip install "PyGithub==1.59.1" + - run: + name: Build Docker image + command: docker build -t my-app:latest -f Dockerfile.database . + - run: + name: Run Docker container + command: | + docker run -d \ + -p 4000:4000 \ + -e DATABASE_URL=$PROXY_DATABASE_URL \ + -e LITELLM_MASTER_KEY="sk-1234" \ + -e OPENAI_API_KEY=$OPENAI_API_KEY \ + -e LITELLM_LICENSE=$LITELLM_LICENSE \ + --name my-app \ + -v $(pwd)/litellm/proxy/example_config_yaml/pass_through_config.yaml:/app/config.yaml \ + -v $(pwd)/litellm/proxy/example_config_yaml/custom_auth_basic.py:/app/custom_auth_basic.py \ + my-app:latest \ + --config /app/config.yaml \ + --port 4000 \ + --detailed_debug \ + - run: + name: Install curl and dockerize + command: | + sudo apt-get update + sudo apt-get install -y curl + sudo wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz + sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz + sudo rm dockerize-linux-amd64-v0.6.1.tar.gz + - run: + name: Start outputting logs + command: docker logs -f my-app + background: true + - run: + name: Wait for app to be ready + command: dockerize -wait http://localhost:4000 -timeout 5m + - run: + name: Run tests + command: | + pwd + ls + python -m pytest -vv tests/pass_through_tests/test_vertex_ai.py -x --junitxml=test-results/junit.xml --durations=5 + no_output_timeout: 120m + + # Store test results + - store_test_results: + path: test-results + publish_to_pypi: docker: - image: cimg/python:3.8 @@ -457,6 +551,12 @@ workflows: only: - main - /litellm_.*/ + - proxy_pass_through_endpoint_tests: + filters: + branches: + only: + - main + - /litellm_.*/ - installing_litellm_on_python: filters: branches: @@ -468,6 +568,7 @@ workflows: - local_testing - build_and_test - proxy_log_to_otel_tests + - proxy_pass_through_endpoint_tests filters: branches: only: diff --git a/litellm/proxy/example_config_yaml/custom_auth_basic.py b/litellm/proxy/example_config_yaml/custom_auth_basic.py new file mode 100644 index 00000000000..2726b1c3d16 --- /dev/null +++ b/litellm/proxy/example_config_yaml/custom_auth_basic.py @@ -0,0 +1,14 @@ +from fastapi import Request + +from litellm.proxy._types import UserAPIKeyAuth + + +async def user_api_key_auth(request: Request, api_key: str) -> UserAPIKeyAuth: + try: + return UserAPIKeyAuth( + api_key="best-api-key-ever", + user_id="best-user-id-ever", + team_id="best-team-id-ever", + ) + except: + raise Exception diff --git a/litellm/proxy/example_config_yaml/pass_through_config.yaml b/litellm/proxy/example_config_yaml/pass_through_config.yaml new file mode 100644 index 00000000000..db9558b3cf3 --- /dev/null +++ b/litellm/proxy/example_config_yaml/pass_through_config.yaml @@ -0,0 +1,9 @@ +model_list: + - model_name: fake-openai-endpoint + litellm_params: + model: openai/fake + api_key: fake-key + api_base: https://exampleopenaiendpoint-production.up.railway.app/ +general_settings: + master_key: sk-1234 + custom_auth: example_config_yaml.custom_auth.user_api_key_auth \ No newline at end of file diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py new file mode 100644 index 00000000000..10660ddb95e --- /dev/null +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -0,0 +1,77 @@ +""" +Test Vertex AI Pass Through + +1. use Credentials client side, Assert SpendLog was created +""" + +import vertexai +from vertexai.preview.generative_models import GenerativeModel +import tempfile +import json +import os +from google.oauth2 import service_account +import google.auth.transport.requests + + +# Path to your service account JSON file +SERVICE_ACCOUNT_FILE = "path/to/your/service-account.json" + + +def load_vertex_ai_credentials(): + # Define the path to the vertex_key.json file + print("loading vertex ai credentials") + filepath = os.path.dirname(os.path.abspath(__file__)) + vertex_key_path = filepath + "/vertex_key.json" + + # Read the existing content of the file or create an empty dictionary + try: + with open(vertex_key_path, "r") as file: + # Read the file content + print("Read vertexai file path") + content = file.read() + + # If the file is empty or not valid JSON, create an empty dictionary + if not content or not content.strip(): + service_account_key_data = {} + else: + # Attempt to load the existing JSON content + file.seek(0) + service_account_key_data = json.load(file) + except FileNotFoundError: + # If the file doesn't exist, create an empty dictionary + service_account_key_data = {} + + # Update the service_account_key_data with environment variables + private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "") + private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "") + private_key = private_key.replace("\\n", "\n") + service_account_key_data["private_key_id"] = private_key_id + service_account_key_data["private_key"] = private_key + + # Create a temporary file + with tempfile.NamedTemporaryFile(mode="w+", delete=False) as temp_file: + # Write the updated content to the temporary files + json.dump(service_account_key_data, temp_file, indent=2) + + # Export the temporary file as GOOGLE_APPLICATION_CREDENTIALS + os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name) + + +LITE_LLM_ENDPOINT = "http://localhost:4000" + + +async def test_basic_vertex_ai_pass_through_with_spendlog(): + + vertexai.init( + project="adroit-crow-413218", + location="us-central1", + api_endpoint=f"{LITE_LLM_ENDPOINT}/vertex-ai", + api_transport="rest", + ) + + model = GenerativeModel(model_name="gemini-1.0-pro") + response = model.generate_content("hi") + + print("response", response) + + pass diff --git a/tests/pass_through_tests/vertex_key.json b/tests/pass_through_tests/vertex_key.json new file mode 100644 index 00000000000..e2fd8512b19 --- /dev/null +++ b/tests/pass_through_tests/vertex_key.json @@ -0,0 +1,13 @@ +{ + "type": "service_account", + "project_id": "adroit-crow-413218", + "private_key_id": "", + "private_key": "", + "client_email": "test-adroit-crow@adroit-crow-413218.iam.gserviceaccount.com", + "client_id": "104886546564708740969", + "auth_uri": "https://accounts.google.com/o/oauth2/auth", + "token_uri": "https://oauth2.googleapis.com/token", + "auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs", + "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-adroit-crow%40adroit-crow-413218.iam.gserviceaccount.com", + "universe_domain": "googleapis.com" +} From c86d1cb391e7247d2f5480aa40126f44a325ce87 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 30 Aug 2024 16:29:32 -0700 Subject: [PATCH 05/14] fix tests --- .circleci/config.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.circleci/config.yml b/.circleci/config.yml index efc0c720cbc..6a99d2a4493 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -251,7 +251,7 @@ jobs: command: | pwd ls - python -m pytest -vv tests/ -x --junitxml=test-results/junit.xml --durations=5 --ignore=tests/otel_tests + python -m pytest -vv tests/ -x --junitxml=test-results/junit.xml --durations=5 --ignore=tests/otel_tests --ignore=tests/pass_through_tests no_output_timeout: 120m # Store test results @@ -450,7 +450,7 @@ jobs: command: | pwd ls - python -m pytest -vv tests/pass_through_tests/test_vertex_ai.py -x --junitxml=test-results/junit.xml --durations=5 + python -m pytest -vv tests/pass_through_tests/ -x --junitxml=test-results/junit.xml --durations=5 no_output_timeout: 120m # Store test results From 414d2dcb523daa3a4aecd4e73e1a771c54c4ea3a Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 30 Aug 2024 16:35:07 -0700 Subject: [PATCH 06/14] call spend logs endpoint --- .../pass_through_config.yaml | 2 +- tests/pass_through_tests/test_vertex_ai.py | 20 +++++++++++++++++++ 2 files changed, 21 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/example_config_yaml/pass_through_config.yaml b/litellm/proxy/example_config_yaml/pass_through_config.yaml index db9558b3cf3..41d581249f1 100644 --- a/litellm/proxy/example_config_yaml/pass_through_config.yaml +++ b/litellm/proxy/example_config_yaml/pass_through_config.yaml @@ -6,4 +6,4 @@ model_list: api_base: https://exampleopenaiendpoint-production.up.railway.app/ general_settings: master_key: sk-1234 - custom_auth: example_config_yaml.custom_auth.user_api_key_auth \ No newline at end of file + custom_auth: custom_auth_basic.user_api_key_auth \ No newline at end of file diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index 10660ddb95e..30600c97d6c 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -57,6 +57,23 @@ def load_vertex_ai_credentials(): os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name) +async def call_spend_logs_endpoint(): + """ + Call this + curl -X GET "http://0.0.0.0:4000/spend/logs?start_date={}" -H "Authorization: Bearer sk-1234" + """ + import datetime + import requests + + todays_date = datetime.datetime.now().strftime("%Y-%m-%d") + url = f"http://0.0.0.0:4000/spend/logs?start_date={todays_date}" + headers = {"Authorization": f"Bearer sk-1234"} + response = requests.get(url, headers=headers) + print("response from call_spend_logs_endpoint", response) + + return response + + LITE_LLM_ENDPOINT = "http://localhost:4000" @@ -74,4 +91,7 @@ async def test_basic_vertex_ai_pass_through_with_spendlog(): print("response", response) + _spend_logs_response = await call_spend_logs_endpoint() + print("spend logs response", _spend_logs_response) + pass From f43060e8df33783f0e729530c04d9b489032fc9d Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 30 Aug 2024 16:40:41 -0700 Subject: [PATCH 07/14] mark as async --- tests/pass_through_tests/test_vertex_ai.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index 30600c97d6c..d2c047b2771 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -9,8 +9,7 @@ from vertexai.preview.generative_models import GenerativeModel import tempfile import json import os -from google.oauth2 import service_account -import google.auth.transport.requests +import pytest # Path to your service account JSON file @@ -77,6 +76,7 @@ async def call_spend_logs_endpoint(): LITE_LLM_ENDPOINT = "http://localhost:4000" +@pytest.mark.asyncio() async def test_basic_vertex_ai_pass_through_with_spendlog(): vertexai.init( From 2c86a624746c80cd09c52f0d6ac18f4fd6af79f5 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 30 Aug 2024 16:50:23 -0700 Subject: [PATCH 08/14] fix vertex ai test --- tests/pass_through_tests/test_vertex_ai.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index d2c047b2771..0b980b33424 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -78,7 +78,7 @@ LITE_LLM_ENDPOINT = "http://localhost:4000" @pytest.mark.asyncio() async def test_basic_vertex_ai_pass_through_with_spendlog(): - + load_vertex_ai_credentials() vertexai.init( project="adroit-crow-413218", location="us-central1", From 06857d108df072790e50fbfdb8b70e4e80a8b02a Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 30 Aug 2024 17:02:24 -0700 Subject: [PATCH 09/14] fix /spend logs call --- .circleci/config.yml | 2 +- tests/pass_through_tests/test_vertex_ai.py | 6 ++++-- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/.circleci/config.yml b/.circleci/config.yml index 6a99d2a4493..bddba054afe 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -251,7 +251,7 @@ jobs: command: | pwd ls - python -m pytest -vv tests/ -x --junitxml=test-results/junit.xml --durations=5 --ignore=tests/otel_tests --ignore=tests/pass_through_tests + python -m pytest -s -vv tests/ -x --junitxml=test-results/junit.xml --durations=5 --ignore=tests/otel_tests --ignore=tests/pass_through_tests no_output_timeout: 120m # Store test results diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index 0b980b33424..675557098f6 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -10,6 +10,7 @@ import tempfile import json import os import pytest +import asyncio # Path to your service account JSON file @@ -59,13 +60,13 @@ def load_vertex_ai_credentials(): async def call_spend_logs_endpoint(): """ Call this - curl -X GET "http://0.0.0.0:4000/spend/logs?start_date={}" -H "Authorization: Bearer sk-1234" + curl -X GET "http://0.0.0.0:4000/spend/logs" -H "Authorization: Bearer sk-1234" """ import datetime import requests todays_date = datetime.datetime.now().strftime("%Y-%m-%d") - url = f"http://0.0.0.0:4000/spend/logs?start_date={todays_date}" + url = f"http://0.0.0.0:4000/global/spend/logs" headers = {"Authorization": f"Bearer sk-1234"} response = requests.get(url, headers=headers) print("response from call_spend_logs_endpoint", response) @@ -91,6 +92,7 @@ async def test_basic_vertex_ai_pass_through_with_spendlog(): print("response", response) + await asyncio.sleep(3) _spend_logs_response = await call_spend_logs_endpoint() print("spend logs response", _spend_logs_response) From b35bfb0302c9b85127362469c93822b12062a11f Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 31 Aug 2024 08:22:27 -0700 Subject: [PATCH 10/14] fix cost tracking for vertex ai native --- litellm/proxy/proxy_config.yaml | 2 +- litellm/proxy/proxy_server.py | 6 ++-- tests/pass_through_tests/test_vertex_ai.py | 32 ++++++++++++++++++---- 3 files changed, 30 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 708f0e27ca3..f9252a4d540 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -21,4 +21,4 @@ router_settings: general_settings: master_key: sk-1234 - custom_auth: example_config_yaml.custom_auth.user_api_key_auth \ No newline at end of file + custom_auth: example_config_yaml.custom_auth_basic.user_api_key_auth \ No newline at end of file diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e0dbcebfc91..80e015e2a9f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -693,10 +693,10 @@ def load_from_azure_key_vault(use_azure_key_vault: bool = False): def cost_tracking(): global prisma_client, custom_db_client if prisma_client is not None or custom_db_client is not None: - if isinstance(litellm.success_callback, list): + if isinstance(litellm._async_success_callback, list): verbose_proxy_logger.debug("setting litellm success callback to track cost") - if (_PROXY_track_cost_callback) not in litellm.success_callback: # type: ignore - litellm.success_callback.append(_PROXY_track_cost_callback) # type: ignore + if (_PROXY_track_cost_callback) not in litellm._async_success_callback: # type: ignore + litellm._async_success_callback.append(_PROXY_track_cost_callback) # type: ignore async def _PROXY_failure_handler( diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index 675557098f6..542105eb453 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -66,12 +66,24 @@ async def call_spend_logs_endpoint(): import requests todays_date = datetime.datetime.now().strftime("%Y-%m-%d") - url = f"http://0.0.0.0:4000/global/spend/logs" + url = f"http://0.0.0.0:4000/global/spend/logs?api_key=best-api-key-ever" headers = {"Authorization": f"Bearer sk-1234"} response = requests.get(url, headers=headers) print("response from call_spend_logs_endpoint", response) - return response + json_response = response.json() + + # get spend for today + """ + json response looks like this + + [{'date': '2024-08-30', 'spend': 0.00016600000000000002, 'api_key': 'best-api-key-ever'}] + """ + + todays_date = datetime.datetime.now().strftime("%Y-%m-%d") + for spend_log in json_response: + if spend_log["date"] == todays_date: + return spend_log["spend"] LITE_LLM_ENDPOINT = "http://localhost:4000" @@ -79,7 +91,10 @@ LITE_LLM_ENDPOINT = "http://localhost:4000" @pytest.mark.asyncio() async def test_basic_vertex_ai_pass_through_with_spendlog(): - load_vertex_ai_credentials() + + spend_before = await call_spend_logs_endpoint() or 0.0 + # load_vertex_ai_credentials() + vertexai.init( project="adroit-crow-413218", location="us-central1", @@ -92,8 +107,13 @@ async def test_basic_vertex_ai_pass_through_with_spendlog(): print("response", response) - await asyncio.sleep(3) - _spend_logs_response = await call_spend_logs_endpoint() - print("spend logs response", _spend_logs_response) + await asyncio.sleep(20) + spend_after = await call_spend_logs_endpoint() + print("spend_after", spend_after) + assert ( + spend_after > spend_before + ), "Spend should be greater than before. spend_before: {}, spend_after: {}".format( + spend_before, spend_after + ) pass From 9e557ed07236ecfb6ab81ca019ef427e02309639 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 31 Aug 2024 08:39:52 -0700 Subject: [PATCH 11/14] fix test --- tests/pass_through_tests/test_vertex_ai.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index 542105eb453..25fd60c72e3 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -93,7 +93,7 @@ LITE_LLM_ENDPOINT = "http://localhost:4000" async def test_basic_vertex_ai_pass_through_with_spendlog(): spend_before = await call_spend_logs_endpoint() or 0.0 - # load_vertex_ai_credentials() + load_vertex_ai_credentials() vertexai.init( project="adroit-crow-413218", From b8bc44847995dad9907f565e6948e95a9ccb943d Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 31 Aug 2024 09:42:58 -0700 Subject: [PATCH 12/14] ci/cd run again --- litellm/tests/test_completion.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/tests/test_completion.py b/litellm/tests/test_completion.py index baf528d8bc5..8fdf722f06b 100644 --- a/litellm/tests/test_completion.py +++ b/litellm/tests/test_completion.py @@ -23,7 +23,7 @@ from litellm import RateLimitError, Timeout, completion, completion_cost, embedd from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.prompt_templates.factory import anthropic_messages_pt -# litellm.num_retries=3 +# litellm.num_retries = 3 litellm.cache = None litellm.success_callback = [] user_message = "Write a short poem about the sky" From 2e0ee8c72fe22760d54359a82b915acd7edda02c Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 31 Aug 2024 14:48:52 -0700 Subject: [PATCH 13/14] skip end of life model in test --- litellm/tests/test_streaming.py | 1 + 1 file changed, 1 insertion(+) diff --git a/litellm/tests/test_streaming.py b/litellm/tests/test_streaming.py index d2ef8aafc7b..1b8b4e08527 100644 --- a/litellm/tests/test_streaming.py +++ b/litellm/tests/test_streaming.py @@ -1545,6 +1545,7 @@ def test_completion_bedrock_claude_stream(): # test_completion_bedrock_claude_stream() +@pytest.mark.skip(reason="model end of life") def test_completion_bedrock_ai21_stream(): try: litellm.set_verbose = False From 9a3873b9edcfcb8477a3f05504ed38dfac329035 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 31 Aug 2024 15:02:56 -0700 Subject: [PATCH 14/14] mark flaky test as flaky --- litellm/tests/test_custom_callback_input.py | 1 + 1 file changed, 1 insertion(+) diff --git a/litellm/tests/test_custom_callback_input.py b/litellm/tests/test_custom_callback_input.py index 739365a702e..1d53682923e 100644 --- a/litellm/tests/test_custom_callback_input.py +++ b/litellm/tests/test_custom_callback_input.py @@ -976,6 +976,7 @@ async def test_async_embedding_bedrock(): # CACHING ## Test Azure - completion, embedding @pytest.mark.asyncio +@pytest.mark.flaky(retries=3, delay=1) async def test_async_completion_azure_caching(): litellm.set_verbose = True customHandler_caching = CompletionCustomHandler()