adds test for chat tranformation

This commit is contained in:
Timothy Lowrimore 2025-07-23 17:06:11 -06:00
parent 4295f3972a
commit 9a9f882693
2 changed files with 101 additions and 25 deletions

View File

@ -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"
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

View File

@ -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"