From 44b1b219115f9d4804ca9dec2b82441059a3c128 Mon Sep 17 00:00:00 2001 From: David Manouchehri Date: Tue, 7 May 2024 21:20:15 +0000 Subject: [PATCH] feat(utils.py) - Add OIDC caching for Google Cloud Run and GitHub Actions. --- litellm/utils.py | 23 +++++++++++++++++++---- 1 file changed, 19 insertions(+), 4 deletions(-) diff --git a/litellm/utils.py b/litellm/utils.py index 652406f64e8..f5d3b974b6d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -34,6 +34,8 @@ from dataclasses import ( import litellm._service_logger # for storing API inputs, outputs, and metadata from litellm.llms.custom_httpx.http_handler import HTTPHandler +from litellm.caching import DualCache +oidc_cache = DualCache() try: # this works in python 3.8 @@ -9291,11 +9293,15 @@ def get_secret( # Example: oidc/google/https://bedrock-runtime.us-east-1.amazonaws.com/model/stability.stable-diffusion-xl-v1/invoke if secret_name.startswith("oidc/"): - secret_name = secret_name.replace("oidc/", "") - oidc_provider, oidc_aud = secret_name.split("/", 1) + secret_name_split = secret_name.replace("oidc/", "") + oidc_provider, oidc_aud = secret_name_split.split("/", 1) # TODO: Add caching for HTTP requests match oidc_provider: case "google": + oidc_token = oidc_cache.get_cache(key=secret_name) + if oidc_token is not None: + return oidc_token + client = HTTPHandler(timeout=httpx.Timeout(timeout=600.0, connect=5.0)) # https://cloud.google.com/compute/docs/instances/verifying-instance-identity#request_signature response = client.get( @@ -9304,7 +9310,9 @@ def get_secret( headers={"Metadata-Flavor": "Google"}, ) if response.status_code == 200: - return response.text + oidc_token = response.text + oidc_cache.set_cache(key=secret_name, value=oidc_token, ttl=3600 - 60) + return oidc_token else: raise ValueError("Google OIDC provider failed") case "circleci": @@ -9325,6 +9333,11 @@ def get_secret( actions_id_token_request_token = os.getenv("ACTIONS_ID_TOKEN_REQUEST_TOKEN") if actions_id_token_request_url is None or actions_id_token_request_token is None: raise ValueError("ACTIONS_ID_TOKEN_REQUEST_URL or ACTIONS_ID_TOKEN_REQUEST_TOKEN not found in environment") + + oidc_token = oidc_cache.get_cache(key=secret_name) + if oidc_token is not None: + return oidc_token + client = HTTPHandler(timeout=httpx.Timeout(timeout=600.0, connect=5.0)) response = client.get( actions_id_token_request_url, @@ -9335,7 +9348,9 @@ def get_secret( }, ) if response.status_code == 200: - return response.text['value'] + oidc_token = response.text['value'] + oidc_cache.set_cache(key=secret_name, value=oidc_token, ttl=300 - 5) + return oidc_token else: raise ValueError("Github OIDC provider failed") case _: