From 900cf2c2eea90510443982575ddc600741024de8 Mon Sep 17 00:00:00 2001 From: Toni Engelhardt Date: Thu, 14 Sep 2023 16:13:31 +0100 Subject: [PATCH] simplify mock logic Adds shortcut for the mock_completion method. --- litellm/main.py | 56 +++++++++++++++++++++++++++---------------------- 1 file changed, 31 insertions(+), 25 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index 575a49f0dc6..f80d33e3219 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -71,6 +71,32 @@ async def acompletion(*args, **kwargs): else: return response +## Use this in your testing pipeline, if you need to mock an LLM response +def mock_completion(model: str, messages: List, stream: bool = False, mock_response: str = "This is a mock request", **kwargs): + try: + model_response = ModelResponse() + if stream: # return a generator object, iterate through the text in chunks of 3 char / chunk + for i in range(0, len(mock_response), 3): + completion_obj = {"role": "assistant", "content": mock_response[i: i+3]} + yield { + "choices": + [ + { + "delta": completion_obj, + "finish_reason": None + }, + ] + } + else: + ## RESPONSE OBJECT + completion_response = "This is a mock request" + model_response["choices"][0]["message"]["content"] = completion_response + model_response["created"] = time.time() + model_response["model"] = "MockResponse" + return model_response + except: + raise Exception("Mock completion response failed") + @client @timeout( # type: ignore 600 @@ -95,6 +121,7 @@ def completion( # Optional liteLLM function params *, return_async=False, + mock_response: Optional[str] = None, api_key: Optional[str] = None, api_version: Optional[str] = None, api_base: Optional[str] = None, @@ -116,6 +143,10 @@ def completion( caching = False, cache_params = {}, # optional to specify metadata for caching ) -> ModelResponse: + # If `mock_response` is set, execute the `mock_completion` method instead. + if mock_response: + return mock_completion(model, messages, stream=stream, mock_response=mock_response) + args = locals() try: logging = litellm_logging_obj @@ -951,31 +982,6 @@ def batch_completion( results = [future.result() for future in completions] return results -## Use this in your testing pipeline, if you need to mock an LLM response -def mock_completion(model: str, messages: List, stream: bool = False, mock_response: str = "This is a mock request", **kwargs): - try: - model_response = ModelResponse() - if stream: # return a generator object, iterate through the text in chunks of 3 char / chunk - for i in range(0, len(mock_response), 3): - completion_obj = {"role": "assistant", "content": mock_response[i: i+3]} - yield { - "choices": - [ - { - "delta": completion_obj, - "finish_reason": None - }, - ] - } - else: - ## RESPONSE OBJECT - completion_response = "This is a mock request" - model_response["choices"][0]["message"]["content"] = completion_response - model_response["created"] = time.time() - model_response["model"] = "MockResponse" - return model_response - except: - raise Exception("Mock completion response failed") ### EMBEDDING ENDPOINTS #################### @client @timeout( # type: ignore