From 29509a48f82dd35dc25ca28cc298b651a47d5b1e Mon Sep 17 00:00:00 2001 From: ishaan-jaff Date: Thu, 5 Oct 2023 11:03:36 -0700 Subject: [PATCH] ollama default api_base to http://localhost:11434 --- litellm/main.py | 15 ++++++------ litellm/tests/test_ollama_local.py | 37 +++++++++++++++++++++--------- 2 files changed, 34 insertions(+), 18 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index 1cb7e300f5b..a79dec1c755 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1041,10 +1041,11 @@ def completion( ## RESPONSE OBJECT response = model_response elif custom_llm_provider == "ollama": - endpoint = ( - litellm.api_base - or api_base - or "http://localhost:11434" + api_base = ( + litellm.api_base or + api_base or + "http://localhost:11434" + ) if model in litellm.custom_prompt_dict: # check if the model has a registered custom prompt @@ -1060,13 +1061,13 @@ def completion( ## LOGGING logging.pre_call( - input=prompt, api_key=None, additional_args={"endpoint": endpoint, "custom_prompt_dict": litellm.custom_prompt_dict} + input=prompt, api_key=None, additional_args={"api_base": api_base, "custom_prompt_dict": litellm.custom_prompt_dict} ) if kwargs.get('acompletion', False) == True: - async_generator = ollama.async_get_ollama_response_stream(endpoint, model, prompt) + async_generator = ollama.async_get_ollama_response_stream(api_base, model, prompt) return async_generator - generator = ollama.get_ollama_response_stream(endpoint, model, prompt) + generator = ollama.get_ollama_response_stream(api_base, model, prompt) if optional_params.get("stream", False) == True: # assume all ollama responses are streamed return generator diff --git a/litellm/tests/test_ollama_local.py b/litellm/tests/test_ollama_local.py index a61b318306d..9692ab844d9 100644 --- a/litellm/tests/test_ollama_local.py +++ b/litellm/tests/test_ollama_local.py @@ -16,18 +16,33 @@ # user_message = "respond in 20 words. who are you?" # messages = [{ "content": user_message,"role": "user"}] -# # def test_completion_ollama(): -# # try: -# # response = completion( -# # model="ollama/llama2", -# # messages=messages, -# # api_base="http://localhost:11434" -# # ) -# # print(response) -# # except Exception as e: -# # pytest.fail(f"Error occurred: {e}") +# def test_completion_ollama(): +# try: +# response = completion( +# model="ollama/llama2", +# messages=messages, +# max_tokens=200, +# request_timeout = 10, -# # test_completion_ollama() +# ) +# print(response) +# except Exception as e: +# pytest.fail(f"Error occurred: {e}") + +# test_completion_ollama() + +# def test_completion_ollama_with_api_base(): +# try: +# response = completion( +# model="ollama/llama2", +# messages=messages, +# api_base="http://localhost:11434" +# ) +# print(response) +# except Exception as e: +# pytest.fail(f"Error occurred: {e}") + +# test_completion_ollama_with_api_base() # # def test_completion_ollama_stream(): # # user_message = "what is litellm?"