From eb9e4f178775db30bc840c5f9885e144e28e3002 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 2 Sep 2024 17:05:53 -0700 Subject: [PATCH] track /embedding in spendLogs --- .../vertex_embeddings/embedding_handler.py | 30 +++++++++++++++++++ .../pass_through_endpoints/success_handler.py | 24 +++++++++++++++ 2 files changed, 54 insertions(+) diff --git a/litellm/llms/vertex_ai_and_google_ai_studio/vertex_embeddings/embedding_handler.py b/litellm/llms/vertex_ai_and_google_ai_studio/vertex_embeddings/embedding_handler.py index 4cd5513c4fd..5638c58cdd6 100644 --- a/litellm/llms/vertex_ai_and_google_ai_studio/vertex_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai_and_google_ai_studio/vertex_embeddings/embedding_handler.py @@ -281,3 +281,33 @@ async def async_embedding( ) setattr(model_response, "usage", usage) return model_response + + +async def transform_vertex_response_to_openai( + response: dict, model: str, model_response: litellm.EmbeddingResponse +) -> litellm.EmbeddingResponse: + + _predictions = response["predictions"] + + embedding_response = [] + input_tokens: int = 0 + for idx, element in enumerate(_predictions): + + embedding = element["embeddings"] + embedding_response.append( + { + "object": "embedding", + "index": idx, + "embedding": embedding["values"], + } + ) + input_tokens += embedding["statistics"]["token_count"] + + model_response.object = "list" + model_response.data = embedding_response + model_response.model = model + usage = Usage( + prompt_tokens=input_tokens, completion_tokens=0, total_tokens=input_tokens + ) + setattr(model_response, "usage", usage) + return model_response diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 618f68659ea..5ed6a194844 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -97,6 +97,30 @@ class PassThroughEndpointLogging: 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, + ) + elif "predict" in url_route: + from litellm.llms.vertex_ai_and_google_ai_studio.vertex_embeddings.embedding_handler import ( + transform_vertex_response_to_openai, + ) + + model = self.extract_model_from_url(url_route) + _json_response = httpx_response.json() + + litellm_model_response = await transform_vertex_response_to_openai( + response=_json_response, + model=model, + model_response=litellm.EmbeddingResponse(), + ) + + litellm_model_response.model = model + 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,