From 9a9f8826935d5156c4e4a69b0ee72b5b1caa011e Mon Sep 17 00:00:00 2001 From: Timothy Lowrimore Date: Wed, 23 Jul 2025 17:06:11 -0600 Subject: [PATCH] adds test for chat tranformation --- litellm/llms/heroku/chat/transformation.py | 93 ++++++++++++++----- .../heroku/test_heroku_chat_transformation.py | 33 +++++++ 2 files changed, 101 insertions(+), 25 deletions(-) create mode 100644 tests/test_litellm/llms/heroku/test_heroku_chat_transformation.py diff --git a/litellm/llms/heroku/chat/transformation.py b/litellm/llms/heroku/chat/transformation.py index a6bceb525a..6626de6900 100644 --- a/litellm/llms/heroku/chat/transformation.py +++ b/litellm/llms/heroku/chat/transformation.py @@ -1,28 +1,71 @@ -from typing import Optional, List, Union -from litellm.llms.base_llm.chat.transformation import BaseConfig -from litellm.types.llms.openai import AllMessageValues +""" +Heroku Chat Completions API -class HerokuChatConfig(BaseConfig): - def validate_environment( +this is OpenAI compatible - no translation needed / occurs +""" +import os + +from typing import Optional, List, Union +from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig + +# Base error class for Heroku +class HerokuError(Exception): + pass + +class HerokuChatConfig(OpenAIGPTConfig): + max_tokens: Optional[int] = None + stop: Optional[List[str]] = None + stream: Optional[bool] = None + temperature: Optional[float] = None + tool_choice: Optional[str] = None + tools: Optional[list] = None + top_p: Optional[int] = None + + def __init__( self, - headers: dict, - model: str, - messages: List[AllMessageValues], - optional_params: dict, - litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> dict: - headers.update({"Authorization": f"Bearer {api_key}"}) - return headers - - def get_complete_url( - self, - api_base: Optional[str], - api_key: Optional[str], - model: str, - optional_params: dict, - litellm_params: dict, + max_tokens: Optional[int] = None, + stop: Optional[List[str]] = None, stream: Optional[bool] = None, - ) -> str: - return f"{api_base}/v1/chat/completions" \ No newline at end of file + temperature: Optional[float] = None, + tool_choice: Optional[str] = None, + tools: Optional[list] = None, + top_p: Optional[int] = None, + ) -> None: + locals_ = locals().copy() + for key, value in locals_.items(): + if key != "self" and value is not None: + setattr(self.__class__, key, value) + + def get_supported_openai_params(self, model: str) -> list: + return [ + "max_tokens", + "stop", + "stream", + "temperature", + "tool_choice", + "top_p", + "tools", + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + supported_openai_params = self.get_supported_openai_params(model=model) + for param, value in non_default_params.items(): + if param in supported_openai_params: + optional_params[param] = value + return optional_params + + def get_complete_url(self, api_base: Optional[str], api_key: Optional[str], model: str, optional_params: dict, litellm_params: dict, stream: Optional[bool] = None) -> str: + api_base = api_base or os.getenv("HEROKU_API_BASE") + if not api_base: + raise HerokuError("No api base was set. Please provide an api_base, or set the HEROKU_API_BASE environment variable.") + + if not api_base.endswith("/v1/chat/completions"): + api_base = f"{api_base}/v1/chat/completions" + + return api_base \ No newline at end of file diff --git a/tests/test_litellm/llms/heroku/test_heroku_chat_transformation.py b/tests/test_litellm/llms/heroku/test_heroku_chat_transformation.py new file mode 100644 index 0000000000..140105c2a2 --- /dev/null +++ b/tests/test_litellm/llms/heroku/test_heroku_chat_transformation.py @@ -0,0 +1,33 @@ +import os +import pytest +from litellm.llms.custom_httpx.http_handler import HTTPHandler +from unittest.mock import patch +from litellm.llms.heroku.chat.transformation import HerokuChatConfig + +class TestHerokuChatConfig: + def test_default_api_base(self): + """Test that default API base is used when none is provided""" + config = HerokuChatConfig() + headers = {} + api_key = "fake-heroku-key" + + # Call validate_environment without specifying api_base + result = config.validate_environment( + headers=headers, + model="claude-3-5-haiku", + messages=[{"role": "user", "content": "Hey"}], + optional_params={}, + litellm_params={}, + api_key=api_key, + api_base=None, # Not providing api_base + ) + + # set env var for api_base + os.environ["HEROKU_API_BASE"] = "https://mia.heroku.com" + + print('****************************************') + print(config.get_complete_url(api_base=None, api_key=api_key, model="claude-3-5-haiku", optional_params={}, litellm_params={}, stream=False)) + print('****************************************') + # Verify headers are still set correctly + assert result["Authorization"] == f"Bearer {api_key}" + assert result["Content-Type"] == "application/json" \ No newline at end of file