diff --git a/litellm/__init__.py b/litellm/__init__.py index 155bca0479c..dcce13c97e1 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -48,6 +48,8 @@ cache: Optional[Cache] = None # cache object <- use this - https://docs.litellm. model_alias_map: Dict[str, str] = {} model_group_alias_map: Dict[str, str] = {} max_budget: float = 0.0 # set the max budget across all providers +_openai_completion_params = ["functions", "function_call", "temperature", "temperature", "top_p", "n", "stream", "stop", "max_tokens", "presence_penalty", "frequency_penalty", "logit_bias", "user", "request_timeout", "api_base", "api_version", "api_key", "deployment_id", "organization", "base_url", "default_headers", "timeout", "response_format", "seed", "tools", "tool_choice", "max_retries"] +_litellm_completion_params = ["metadata", "acompletion", "caching", "mock_response", "api_key", "api_version", "api_base", "force_timeout", "logger_fn", "verbose", "custom_llm_provider", "litellm_logging_obj", "litellm_call_id", "use_client", "id", "fallbacks", "azure", "headers", "model_list", "num_retries", "context_window_fallback_dict", "roles", "final_prompt_value", "bos_token", "eos_token", "request_timeout", "complete_response", "self", "client", "rpm", "tpm", "input_cost_per_token", "output_cost_per_token", "hf_model_name", "model_info", "proxy_server_request", "preset_cache_key"] _current_cost = 0 # private variable, used if max budget is set error_logs: Dict = {} add_function_to_prompt: bool = False # if function calling not supported by api, append function call details to system prompt diff --git a/litellm/caching.py b/litellm/caching.py index 556ab4fb5dc..d3e18270b76 100644 --- a/litellm/caching.py +++ b/litellm/caching.py @@ -232,7 +232,9 @@ class Cache: # sort kwargs by keys, since model: [gpt-4, temperature: 0.2, max_tokens: 200] == [temperature: 0.2, max_tokens: 200, model: gpt-4] completion_kwargs = ["model", "messages", "temperature", "top_p", "n", "stop", "max_tokens", "presence_penalty", "frequency_penalty", "logit_bias", "user", "response_format", "seed", "tools", "tool_choice"] - for param in completion_kwargs: + embedding_kwargs = ["model", "input", "user", "encoding_format"] + combined_kwargs = list(set(completion_kwargs + embedding_kwargs)) + for param in combined_kwargs: # ignore litellm params here if param in kwargs: # check if param == model and model_group is passed in, then override model with model_group diff --git a/litellm/tests/test_custom_callback_input.py b/litellm/tests/test_custom_callback_input.py index 66033308015..a2258a1704b 100644 --- a/litellm/tests/test_custom_callback_input.py +++ b/litellm/tests/test_custom_callback_input.py @@ -5,7 +5,7 @@ from datetime import datetime import pytest sys.path.insert(0, os.path.abspath('../..')) from typing import Optional, Literal, List, Union -from litellm import completion, embedding +from litellm import completion, embedding, Cache import litellm from litellm.integrations.custom_logger import CustomLogger @@ -14,6 +14,7 @@ from litellm.integrations.custom_logger import CustomLogger ## 2: Post-API-Call ## 3: On LiteLLM Call success ## 4: On LiteLLM Call failure +## 5. Caching # Test models ## 1. OpenAI @@ -32,7 +33,7 @@ class CompletionCustomHandler(CustomLogger): # https://docs.litellm.ai/docs/obse def __init__(self): self.errors = [] self.states: Optional[List[Literal["sync_pre_api_call", "async_pre_api_call", "post_api_call", "sync_stream", "async_stream", "sync_success", "async_success", "sync_failure", "async_failure"]]] = [] - + def log_pre_api_call(self, model, messages, kwargs): try: self.states.append("sync_pre_api_call") @@ -126,6 +127,7 @@ class CompletionCustomHandler(CustomLogger): # https://docs.litellm.ai/docs/obse assert isinstance(kwargs['original_response'], (str, litellm.CustomStreamWrapper)) assert isinstance(kwargs['additional_args'], (dict, type(None))) assert isinstance(kwargs['log_event_type'], str) + assert isinstance(kwargs["cache_hit"], Optional[bool]) except: print(f"Assertion Error: {traceback.format_exc()}") self.errors.append(traceback.format_exc()) @@ -197,7 +199,7 @@ class CompletionCustomHandler(CustomLogger): # https://docs.litellm.ai/docs/obse assert isinstance(kwargs['original_response'], (str, litellm.CustomStreamWrapper)) or inspect.isasyncgen(kwargs['original_response']) or inspect.iscoroutine(kwargs['original_response']) assert isinstance(kwargs['additional_args'], (dict, type(None))) assert isinstance(kwargs['log_event_type'], str) - + assert isinstance(kwargs["cache_hit"], Optional[bool]) except: print(f"Assertion Error: {traceback.format_exc()}") self.errors.append(traceback.format_exc()) @@ -577,4 +579,47 @@ async def test_async_embedding_bedrock(): except Exception as e: pytest.fail(f"An exception occurred: {str(e)}") -# asyncio.run(test_async_embedding_bedrock()) \ No newline at end of file +# asyncio.run(test_async_embedding_bedrock()) + +# CACHING +## Test Azure - completion, embedding +@pytest.mark.asyncio +async def test_async_completion_azure_caching(): + customHandler_caching = CompletionCustomHandler() + litellm.cache = Cache(type="redis", host=os.environ['REDIS_HOST'], port=os.environ['REDIS_PORT'], password=os.environ['REDIS_PASSWORD']) + litellm.callbacks = [customHandler_caching] + unique_time = time.time() + response1 = await litellm.acompletion(model="azure/chatgpt-v-2", + messages=[{ + "role": "user", + "content": f"Hi 👋 - i'm async azure {unique_time}" + }], + caching=True) + await asyncio.sleep(1) + print(f"customHandler_caching.states pre-cache hit: {customHandler_caching.states}") + response2 = await litellm.acompletion(model="azure/chatgpt-v-2", + messages=[{ + "role": "user", + "content": f"Hi 👋 - i'm async azure {unique_time}" + }], + caching=True) + await asyncio.sleep(1) # success callbacks are done in parallel + print(f"customHandler_caching.states post-cache hit: {customHandler_caching.states}") + assert len(customHandler_caching.errors) == 0 + assert len(customHandler_caching.states) == 4 # pre, post, success, success + +@pytest.mark.asyncio +async def test_async_embedding_azure_caching(): + customHandler_caching = CompletionCustomHandler() + litellm.cache = Cache(type="redis", host=os.environ['REDIS_HOST'], port=os.environ['REDIS_PORT'], password=os.environ['REDIS_PASSWORD']) + litellm.callbacks = [customHandler_caching] + unique_time = time.time() + response1 = await litellm.aembedding(model="azure/azure-embedding-model", + input=[f"good morning from litellm1 {unique_time}"], + caching=True) + response2 = await litellm.aembedding(model="azure/azure-embedding-model", + input=[f"good morning from litellm1 {unique_time}"], + caching=True) + await asyncio.sleep(1) # success callbacks are done in parallel + assert len(customHandler_caching.errors) == 0 + assert len(customHandler_caching.states) == 4 # pre, post, success, success diff --git a/litellm/tests/test_custom_callback_router.py b/litellm/tests/test_custom_callback_router.py index d9f67d6e3a4..8c31300a146 100644 --- a/litellm/tests/test_custom_callback_router.py +++ b/litellm/tests/test_custom_callback_router.py @@ -150,6 +150,7 @@ class CompletionCustomHandler(CustomLogger): # https://docs.litellm.ai/docs/obse assert isinstance(kwargs['original_response'], (str, litellm.CustomStreamWrapper)) assert isinstance(kwargs['additional_args'], (dict, type(None))) assert isinstance(kwargs['log_event_type'], str) + assert isinstance(kwargs["cache_hit"], Optional[bool]) except: print(f"Assertion Error: {traceback.format_exc()}") self.errors.append(traceback.format_exc()) @@ -213,6 +214,7 @@ class CompletionCustomHandler(CustomLogger): # https://docs.litellm.ai/docs/obse assert isinstance(kwargs['original_response'], (str, litellm.CustomStreamWrapper)) or inspect.isasyncgen(kwargs['original_response']) or inspect.iscoroutine(kwargs['original_response']) assert isinstance(kwargs['additional_args'], (dict, type(None))) assert isinstance(kwargs['log_event_type'], str) + assert isinstance(kwargs["cache_hit"], Optional[bool]) ### ROUTER-SPECIFIC KWARGS assert isinstance(kwargs["litellm_params"]["metadata"], dict) assert isinstance(kwargs["litellm_params"]["metadata"]["model_group"], str) diff --git a/litellm/utils.py b/litellm/utils.py index 80064bb6fdc..a686828a1c0 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -574,8 +574,9 @@ class Logging: self.litellm_call_id = litellm_call_id self.function_id = function_id self.streaming_chunks = [] # for generating complete stream response + self.model_call_details = {} - def update_environment_variables(self, model, user, optional_params, litellm_params): + def update_environment_variables(self, model, user, optional_params, litellm_params, **additional_params): self.optional_params = optional_params self.model = model self.user = user @@ -590,7 +591,8 @@ class Logging: "start_time": self.start_time, "stream": self.stream, "user": user, - **self.optional_params + **self.optional_params, + **additional_params } def _pre_call(self, input, api_key, model=None, additional_args={}): @@ -821,7 +823,7 @@ class Logging: ) pass - def _success_handler_helper_fn(self, result=None, start_time=None, end_time=None): + def _success_handler_helper_fn(self, result=None, start_time=None, end_time=None, cache_hit=None): try: if start_time is None: start_time = self.start_time @@ -829,6 +831,7 @@ class Logging: end_time = datetime.datetime.now() self.model_call_details["log_event_type"] = "successful_api_call" self.model_call_details["end_time"] = end_time + self.model_call_details["cache_hit"] = cache_hit if litellm.max_budget and self.stream: time_diff = (end_time - start_time).total_seconds() @@ -836,10 +839,10 @@ class Logging: litellm._current_cost += litellm.completion_cost(model=self.model, prompt="", completion=result["content"], total_time=float_diff) return start_time, end_time, result - except: - pass + except Exception as e: + print_verbose(f"[Non-Blocking] LiteLLM.Success_Call Error: {str(e)}") - def success_handler(self, result=None, start_time=None, end_time=None, **kwargs): + def success_handler(self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs): print_verbose( f"Logging Details LiteLLM-Success Call" ) @@ -867,7 +870,7 @@ class Logging: if complete_streaming_response: self.model_call_details["complete_streaming_response"] = complete_streaming_response - start_time, end_time, result = self._success_handler_helper_fn(start_time=start_time, end_time=end_time, result=result) + start_time, end_time, result = self._success_handler_helper_fn(start_time=start_time, end_time=end_time, result=result, cache_hit=cache_hit) for callback in litellm.success_callback: try: if callback == "lite_debugger": @@ -1063,7 +1066,7 @@ class Logging: ) pass - async def async_success_handler(self, result=None, start_time=None, end_time=None, **kwargs): + async def async_success_handler(self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs): """ Implementing async callbacks, to handle asyncio event loop issues when custom integrations need to use async functions. """ @@ -1082,7 +1085,7 @@ class Logging: self.streaming_chunks.append(result) if complete_streaming_response: self.model_call_details["complete_streaming_response"] = complete_streaming_response - start_time, end_time, result = self._success_handler_helper_fn(start_time=start_time, end_time=end_time, result=result) + start_time, end_time, result = self._success_handler_helper_fn(start_time=start_time, end_time=end_time, result=result, cache_hit=cache_hit) for callback in litellm._async_success_callback: try: if callback == "cache" and litellm.cache is not None: @@ -1440,6 +1443,7 @@ def client(original_function): model = args[0] if len(args) > 0 else kwargs["model"] call_type = original_function.__name__ if call_type == CallTypes.completion.value or call_type == CallTypes.acompletion.value: + messages = None if len(args) > 1: messages = args[1] elif kwargs.get("messages", None): @@ -1509,11 +1513,12 @@ def client(original_function): if litellm._current_cost > litellm.max_budget: raise BudgetExceededError(current_cost=litellm._current_cost, max_budget=litellm.max_budget) - # [OPTIONAL] CHECK CACHE # remove this after deprecating litellm.caching if (litellm.caching or litellm.caching_with_models) and litellm.cache is None: litellm.cache = Cache() + + # [OPTIONAL] CHECK CACHE print_verbose(f"kwargs[caching]: {kwargs.get('caching', False)}; litellm.cache: {litellm.cache}") # if caching is false, don't run this if (kwargs.get("caching", None) is None and litellm.cache is not None) or kwargs.get("caching", False) == True: # allow users to control returning cached responses from the completion function @@ -1563,11 +1568,6 @@ def client(original_function): # LOG SUCCESS - handle streaming success logging in the _next_ object, remove `handle_success` once it's deprecated print_verbose(f"Wrapper: Completed Call, calling success_handler") threading.Thread(target=logging_obj.success_handler, args=(result, start_time, end_time)).start() - # threading.Thread(target=logging_obj.success_handler, args=(result, start_time, end_time)).start() - my_thread = threading.Thread( - target=handle_success, args=(args, kwargs, result, start_time, end_time) - ) # don't interrupt execution of main thread - my_thread.start() # RETURN RESULT result._response_ms = (end_time - start_time).total_seconds() * 1000 # return response latency in ms like openai return result @@ -1648,13 +1648,22 @@ def client(original_function): call_type = original_function.__name__ if call_type == CallTypes.acompletion.value and isinstance(cached_result, dict): if kwargs.get("stream", False) == True: - return convert_to_streaming_response_async( + cached_result = convert_to_streaming_response_async( response_object=cached_result, ) else: - return convert_to_model_response_object(response_object=cached_result, model_response_object=ModelResponse()) - else: - return cached_result + cached_result = convert_to_model_response_object(response_object=cached_result, model_response_object=ModelResponse()) + elif call_type == CallTypes.aembedding.value and isinstance(cached_result, dict): + cached_result = convert_to_model_response_object(response_object=cached_result, model_response_object=EmbeddingResponse(), response_type="embedding") + # LOG SUCCESS + cache_hit = True + end_time = datetime.datetime.now() + model, custom_llm_provider, dynamic_api_key, api_base = litellm.get_llm_provider(model=model, custom_llm_provider=kwargs.get('custom_llm_provider', None), api_base=kwargs.get('api_base', None), api_key=kwargs.get('api_key', None)) + print_verbose(f"Async Wrapper: Completed Call, calling async_success_handler: {logging_obj.async_success_handler}") + logging_obj.update_environment_variables(model=model, user=kwargs.get('user', None), optional_params={}, litellm_params={"logger_fn": kwargs.get('logger_fn', None), "acompletion": True}, input=kwargs.get('messages', ""), api_key=kwargs.get('api_key', None), original_response=str(cached_result), additional_args=None, stream=kwargs.get('stream', False)) + asyncio.create_task(logging_obj.async_success_handler(cached_result, start_time, end_time, cache_hit)) + threading.Thread(target=logging_obj.success_handler, args=(cached_result, start_time, end_time, cache_hit)).start() + return cached_result # MODEL CALL result = await original_function(*args, **kwargs) end_time = datetime.datetime.now() @@ -1672,7 +1681,10 @@ def client(original_function): # [OPTIONAL] ADD TO CACHE if litellm.caching or litellm.caching_with_models or litellm.cache != None: # user init a cache object - litellm.cache.add_cache(result, *args, **kwargs) + if isinstance(result, litellm.ModelResponse) or isinstance(result, litellm.EmbeddingResponse): + litellm.cache.add_cache(result.json(), *args, **kwargs) + else: + litellm.cache.add_cache(result, *args, **kwargs) # LOG SUCCESS - handle streaming success logging in the _next_ object print_verbose(f"Async Wrapper: Completed Call, calling async_success_handler: {logging_obj.async_success_handler}") asyncio.create_task(logging_obj.async_success_handler(result, start_time, end_time))