mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #12992 from tlowrimore-heroku/heroku-llms
Heroku llms
This commit is contained in:
commit
2269ea7f31
13 changed files with 427 additions and 1 deletions
|
|
@ -344,6 +344,7 @@ curl 'http://0.0.0.0:4000/key/generate' \
|
|||
| [Novita AI](https://novita.ai/models/llm?utm_source=github_litellm&utm_medium=github_readme&utm_campaign=github_link) | ✅ | ✅ | ✅ | ✅ | | |
|
||||
| [Featherless AI](https://docs.litellm.ai/docs/providers/featherless_ai) | ✅ | ✅ | ✅ | ✅ | | |
|
||||
| [Nebius AI Studio](https://docs.litellm.ai/docs/providers/nebius) | ✅ | ✅ | ✅ | ✅ | ✅ | |
|
||||
| [Heroku](https://docs.litellm.ai/docs/providers/heroku) | ✅ | ✅ | | | | |
|
||||
|
||||
[**Read the Docs**](https://docs.litellm.ai/docs/)
|
||||
|
||||
|
|
|
|||
76
docs/my-website/docs/providers/heroku.md
Normal file
76
docs/my-website/docs/providers/heroku.md
Normal file
|
|
@ -0,0 +1,76 @@
|
|||
# Heroku
|
||||
|
||||
## Provision a Model
|
||||
|
||||
To use Heroku with LiteLLM, [configure a Heroku app and attach a supported model](https://devcenter.heroku.com/articles/heroku-inference#provision-access-to-an-ai-model-resource).
|
||||
|
||||
|
||||
## Supported Models
|
||||
|
||||
Heroku for LiteLLM supports various [chat](https://devcenter.heroku.com/articles/heroku-inference-api-v1-chat-completions) models:
|
||||
|
||||
| Model | Region |
|
||||
|-----------------------------------|---------|
|
||||
| [`heroku/claude-sonnet-4`](https://devcenter.heroku.com/articles/heroku-inference-api-model-claude-4-sonnet) | US, EU |
|
||||
| [`heroku/claude-3-7-sonnet`](https://devcenter.heroku.com/articles/heroku-inference-api-model-claude-3-7-sonnet) | US, EU |
|
||||
| [`heroku/claude-3-5-sonnet-latest`](https://devcenter.heroku.com/articles/heroku-inference-api-model-claude-3-5-sonnet-latest) | US |
|
||||
| [`heroku/claude-3-5-haiku`](https://devcenter.heroku.com/articles/heroku-inference-api-model-claude-3-5-haiku) | US |
|
||||
| [`heroku/claude-3`](https://devcenter.heroku.com/articles/heroku-inference-api-model-claude-3-haiku) | EU |
|
||||
|
||||
## Environment Variables
|
||||
|
||||
When you attach a model to a Heroku app, three config variables are set:
|
||||
|
||||
- `INFERENCE_KEY`: The API key used for authenticating requests to the model.
|
||||
- `INFERENCE_MODEL_ID`: The name of the model, for example`claude-3-5-haiku`.
|
||||
- `INFERENCE_URL`: The base URL for calling the model.
|
||||
|
||||
Both `INFERENCE_KEY` and `INFERENCE_URL` are required to make calls to your model.
|
||||
|
||||
For more information on these variables, see the [Heroku documentation](https://devcenter.heroku.com/articles/heroku-inference#model-resource-config-vars).
|
||||
|
||||
## Usage Examples
|
||||
### Using Config Variables
|
||||
|
||||
Heroku uses the following LiteLLM API config variables:
|
||||
|
||||
- `HEROKU_API_KEY`: This value corresponds to [LiteLLM's `api_key` param](https://docs.litellm.ai/docs/set_keys#litellmapi_key). Set this variable to the value of Heroku's `INFERENCE_KEY` config variable.
|
||||
- `HEROKU_API_BASE`: This value corresponds to [LiteLLM's `api_base` param](https://docs.litellm.ai/docs/set_keys#litellmapi_base). Set this variable to the value of Heroku's `INFERENCE_URL` config variable.
|
||||
|
||||
In this example, we don't explicitly pass the `api_key` and `api_base` variables. Instead, we set the config variables which Heroku will use:
|
||||
|
||||
```python
|
||||
import os
|
||||
from litellm import completion
|
||||
|
||||
os.environ["HEROKU_API_BASE"] = "https://us.inference.heroku.com"
|
||||
os.environ["HEROKU_API_KEY"] = "fake-heroku-key"
|
||||
|
||||
response = completion(
|
||||
model="heroku/claude-3-5-haiku",
|
||||
messages=[
|
||||
{"role": "user", "content": "write code for saying hey from LiteLLM"}
|
||||
]
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
> Include the `heroku/` prefix in the model name so LiteLLM knows the model provider to use.
|
||||
|
||||
### Explicitly Setting `api_key` and `api_base`
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
response = completion(
|
||||
model="heroku/claude-sonnet-4",
|
||||
api_key="fake-heroku-key",
|
||||
api_base="https://us.inference.heroku.com",
|
||||
messages=[
|
||||
{"role": "user", "content": "write code for saying hey from LiteLLM"}
|
||||
],
|
||||
)
|
||||
```
|
||||
|
||||
> Include the `heroku/` prefix in the model name so LiteLLM knows the model provider to use.
|
||||
|
|
@ -482,6 +482,7 @@ const sidebars = {
|
|||
"providers/nebius",
|
||||
"providers/dashscope",
|
||||
"providers/bytez",
|
||||
"providers/heroku",
|
||||
"providers/oci",
|
||||
"providers/datarobot",
|
||||
],
|
||||
|
|
|
|||
|
|
@ -239,6 +239,7 @@ novita_api_key: Optional[str] = None
|
|||
snowflake_key: Optional[str] = None
|
||||
gradient_ai_api_key: Optional[str] = None
|
||||
nebius_key: Optional[str] = None
|
||||
heroku_key: Optional[str] = None
|
||||
cometapi_key: Optional[str] = None
|
||||
common_cloud_provider_auth_params: dict = {
|
||||
"params": ["project", "region_name", "token"],
|
||||
|
|
@ -482,6 +483,7 @@ azure_ai_models: Set = set()
|
|||
jina_ai_models: Set = set()
|
||||
voyage_models: Set = set()
|
||||
infinity_models: Set = set()
|
||||
heroku_models: Set = set()
|
||||
databricks_models: Set = set()
|
||||
cloudflare_models: Set = set()
|
||||
codestral_models: Set = set()
|
||||
|
|
@ -710,6 +712,8 @@ def add_known_models():
|
|||
deepgram_models.add(key)
|
||||
elif value.get("litellm_provider") == "elevenlabs":
|
||||
elevenlabs_models.add(key)
|
||||
elif value.get("litellm_provider") == "heroku":
|
||||
heroku_models.add(key)
|
||||
elif value.get("litellm_provider") == "dashscope":
|
||||
dashscope_models.add(key)
|
||||
elif value.get("litellm_provider") == "moonshot":
|
||||
|
|
@ -821,6 +825,7 @@ model_list = list(
|
|||
| recraft_models
|
||||
| cometapi_models
|
||||
| oci_models
|
||||
| heroku_models
|
||||
| vercel_ai_gateway_models
|
||||
| volcengine_models
|
||||
)
|
||||
|
|
@ -893,6 +898,7 @@ models_by_provider: dict = {
|
|||
"featherless_ai": featherless_ai_models,
|
||||
"deepgram": deepgram_models,
|
||||
"elevenlabs": elevenlabs_models,
|
||||
"heroku": heroku_models,
|
||||
"dashscope": dashscope_models,
|
||||
"moonshot": moonshot_models,
|
||||
"v0": v0_models,
|
||||
|
|
@ -1220,6 +1226,7 @@ from .llms.azure.azure import (
|
|||
AzureOpenAIError,
|
||||
AzureOpenAIAssistantsAPIConfig,
|
||||
)
|
||||
from .llms.heroku.chat.transformation import HerokuChatConfig
|
||||
from .llms.cometapi.chat.transformation import CometAPIConfig
|
||||
from .llms.azure.chat.gpt_transformation import AzureOpenAIConfig
|
||||
from .llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config
|
||||
|
|
|
|||
|
|
@ -306,6 +306,7 @@ LITELLM_CHAT_PROVIDERS = [
|
|||
"dashscope",
|
||||
"moonshot",
|
||||
"v0",
|
||||
"heroku",
|
||||
"oci",
|
||||
"morph",
|
||||
"lambda_ai",
|
||||
|
|
|
|||
|
|
@ -365,6 +365,8 @@ def get_llm_provider( # noqa: PLR0915
|
|||
# bytez models
|
||||
elif model.startswith("bytez/"):
|
||||
custom_llm_provider = "bytez"
|
||||
elif model.startswith("heroku/"):
|
||||
custom_llm_provider = "heroku"
|
||||
# cometapi models
|
||||
elif model.startswith("cometapi/"):
|
||||
custom_llm_provider = "cometapi"
|
||||
|
|
@ -704,6 +706,13 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
|
|||
) = litellm.NscaleConfig()._get_openai_compatible_provider_info(
|
||||
api_base=api_base, api_key=api_key
|
||||
)
|
||||
elif custom_llm_provider == "heroku":
|
||||
(
|
||||
api_base,
|
||||
dynamic_api_key,
|
||||
) = litellm.HerokuChatConfig()._get_openai_compatible_provider_info(
|
||||
api_base, api_key
|
||||
)
|
||||
elif custom_llm_provider == "dashscope":
|
||||
(
|
||||
api_base,
|
||||
|
|
|
|||
67
litellm/llms/heroku/chat/transformation.py
Normal file
67
litellm/llms/heroku/chat/transformation.py
Normal file
|
|
@ -0,0 +1,67 @@
|
|||
"""
|
||||
Heroku Chat Completions API
|
||||
|
||||
this is OpenAI compatible - no translation needed / occurs
|
||||
"""
|
||||
import os
|
||||
|
||||
from typing import Optional, List, Tuple, Union, Coroutine, Any, Literal, overload
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
handle_messages_with_content_list_to_str_conversion,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
# Base error class for Heroku
|
||||
class HerokuError(Exception):
|
||||
pass
|
||||
|
||||
class HerokuChatConfig(OpenAIGPTConfig):
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: Literal[True]
|
||||
) -> Coroutine[Any, Any, List[AllMessageValues]]:
|
||||
...
|
||||
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self,
|
||||
messages: List[AllMessageValues],
|
||||
model: str,
|
||||
is_async: Literal[False] = False,
|
||||
) -> List[AllMessageValues]:
|
||||
...
|
||||
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: bool = False
|
||||
) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]:
|
||||
"""
|
||||
Heroku does not support content in list format.
|
||||
See: https://devcenter.heroku.com/articles/heroku-inference-api-v1-chat-completions#content-object
|
||||
"""
|
||||
messages = handle_messages_with_content_list_to_str_conversion(messages)
|
||||
if is_async:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=True
|
||||
)
|
||||
else:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=False
|
||||
)
|
||||
|
||||
def _get_openai_compatible_provider_info(self, api_base: Optional[str], api_key: Optional[str]) -> Tuple[Optional[str], Optional[str]]:
|
||||
api_base = api_base or os.getenv("HEROKU_API_BASE")
|
||||
api_key = api_key or os.getenv("HEROKU_API_KEY")
|
||||
|
||||
return api_base, api_key
|
||||
|
||||
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, _ = self._get_openai_compatible_provider_info(api_base, api_key)
|
||||
|
||||
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
|
||||
|
|
@ -150,8 +150,9 @@ from .llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
|||
from .llms.custom_llm import CustomLLM, custom_chat_llm_router
|
||||
from .llms.databricks.embed.handler import DatabricksEmbeddingHandler
|
||||
from .llms.deprecated_providers import aleph_alpha, palm
|
||||
from .llms.gemini.common_utils import get_api_key_from_env
|
||||
from .llms.groq.chat.handler import GroqChatCompletion
|
||||
from .llms.heroku.chat.transformation import HerokuChatConfig
|
||||
from .llms.gemini.common_utils import get_api_key_from_env
|
||||
from .llms.huggingface.embedding.handler import HuggingFaceEmbedding
|
||||
from .llms.nlp_cloud.chat.handler import completion as nlp_cloud_chat_completion
|
||||
from .llms.oci.chat.transformation import OCIChatConfig
|
||||
|
|
@ -256,6 +257,7 @@ base_llm_http_handler = BaseLLMHTTPHandler()
|
|||
base_llm_aiohttp_handler = BaseLLMAIOHTTPHandler()
|
||||
sagemaker_chat_completion = SagemakerChatHandler()
|
||||
bytez_transformation = BytezChatConfig()
|
||||
heroku_transformation = HerokuChatConfig()
|
||||
oci_transformation = OCIChatConfig()
|
||||
####### COMPLETION ENDPOINTS ################
|
||||
|
||||
|
|
@ -1773,6 +1775,35 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
additional_args={"headers": headers},
|
||||
)
|
||||
raise e
|
||||
elif custom_llm_provider == "heroku":
|
||||
try:
|
||||
response = base_llm_http_handler.completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
headers=headers,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
acompletion=acompletion,
|
||||
logging_obj=logging,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
encoding=encoding,
|
||||
stream=stream,
|
||||
provider_config=provider_config,
|
||||
)
|
||||
except Exception as e:
|
||||
logging.post_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
original_response=str(e),
|
||||
additional_args={"headers": headers},
|
||||
)
|
||||
raise e
|
||||
|
||||
elif custom_llm_provider == "xai":
|
||||
## COMPLETION CALL
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -19958,6 +19958,38 @@
|
|||
"supports_system_messages": true,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
"heroku/claude-4-sonnet": {
|
||||
"max_tokens": 8192,
|
||||
"litellm_provider": "heroku",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"heroku/claude-3-7-sonnet": {
|
||||
"max_tokens": 8192,
|
||||
"litellm_provider": "heroku",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"heroku/claude-3-5-sonnet-latest": {
|
||||
"max_tokens": 8192,
|
||||
"litellm_provider": "heroku",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"heroku/claude-3-5-haiku": {
|
||||
"max_tokens": 4096,
|
||||
"litellm_provider": "heroku",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"vercel_ai_gateway/alibaba/qwen3-coder": {
|
||||
"max_tokens": 262144,
|
||||
"input_cost_per_token": 4e-07,
|
||||
|
|
|
|||
|
|
@ -2349,6 +2349,7 @@ class LlmProviders(str, Enum):
|
|||
PG_VECTOR = "pg_vector"
|
||||
HYPERBOLIC = "hyperbolic"
|
||||
RECRAFT = "recraft"
|
||||
HEROKU = "heroku"
|
||||
AIML = "aiml"
|
||||
COMETAPI = "cometapi"
|
||||
OCI = "oci"
|
||||
|
|
|
|||
|
|
@ -7076,6 +7076,8 @@ class ProviderConfigManager:
|
|||
return litellm.GradientAIConfig()
|
||||
elif litellm.LlmProviders.NSCALE == provider:
|
||||
return litellm.NscaleConfig()
|
||||
elif litellm.LlmProviders.HEROKU == provider:
|
||||
return litellm.HerokuChatConfig()
|
||||
elif litellm.LlmProviders.OCI == provider:
|
||||
return litellm.OCIChatConfig()
|
||||
elif litellm.LlmProviders.HYPERBOLIC == provider:
|
||||
|
|
|
|||
|
|
@ -19958,6 +19958,38 @@
|
|||
"supports_system_messages": true,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
"heroku/claude-4-sonnet": {
|
||||
"max_tokens": 8192,
|
||||
"litellm_provider": "heroku",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"heroku/claude-3-7-sonnet": {
|
||||
"max_tokens": 8192,
|
||||
"litellm_provider": "heroku",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"heroku/claude-3-5-sonnet-latest": {
|
||||
"max_tokens": 8192,
|
||||
"litellm_provider": "heroku",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"heroku/claude-3-5-haiku": {
|
||||
"max_tokens": 4096,
|
||||
"litellm_provider": "heroku",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"vercel_ai_gateway/alibaba/qwen3-coder": {
|
||||
"max_tokens": 262144,
|
||||
"input_cost_per_token": 4e-07,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,166 @@
|
|||
import os
|
||||
import pytest
|
||||
import litellm
|
||||
from litellm import completion
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from unittest.mock import patch
|
||||
from litellm.llms.heroku.chat.transformation import HerokuChatConfig
|
||||
|
||||
os.environ["HEROKU_API_BASE"] = "https://us.inference.heroku.com"
|
||||
os.environ["HEROKU_API_KEY"] = "fake-heroku-key"
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
# Verify headers are still set correctly
|
||||
assert result["Authorization"] == f"Bearer {api_key}"
|
||||
assert result["Content-Type"] == "application/json"
|
||||
|
||||
@pytest.mark.respx()
|
||||
def test_heroku_chat_mock(self, respx_mock):
|
||||
"""Test that the Heroku chat API is called correctly"""
|
||||
|
||||
litellm.disable_aiohttp_transport = True
|
||||
|
||||
model = "heroku/claude-3-5-haiku"
|
||||
model_name = "claude-3-5-haiku"
|
||||
|
||||
respx_mock.post("https://us.inference.heroku.com/v1/chat/completions").respond(
|
||||
json={
|
||||
"id": "chatcmpl-123",
|
||||
"object": "chat.completion",
|
||||
"created": 1677652288,
|
||||
"model": model_name,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "It's me, Mia! How are you?",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 9,
|
||||
"completion_tokens": 12,
|
||||
"total_tokens": 21,
|
||||
},
|
||||
},
|
||||
status_code=200,
|
||||
)
|
||||
|
||||
response = completion(
|
||||
model=model,
|
||||
messages=[
|
||||
{"role": "user", "content": "write code for saying hey from LiteLLM"}
|
||||
],
|
||||
extended_thinking={ "enabled": True, "include_reasoning":True }
|
||||
)
|
||||
|
||||
# Verify the request was made with correct headers
|
||||
assert len(respx_mock.calls) == 1
|
||||
request = respx_mock.calls[0].request
|
||||
|
||||
assert request.headers["Authorization"] == f"Bearer {os.environ['HEROKU_API_KEY']}"
|
||||
assert request.headers["Content-Type"] == "application/json"
|
||||
|
||||
assert response.choices[0].message.content == "It's me, Mia! How are you?"
|
||||
|
||||
@pytest.mark.respx()
|
||||
def test_heroku_tool_calling(self, respx_mock):
|
||||
"""Test that the Heroku tool calling API is called correctly"""
|
||||
config = HerokuChatConfig()
|
||||
headers = {}
|
||||
api_key = "fake-heroku-key"
|
||||
|
||||
litellm.disable_aiohttp_transport = True
|
||||
|
||||
model = "heroku/claude-4-sonnet"
|
||||
|
||||
respx_mock.post("https://us.inference.heroku.com/v1/chat/completions").respond(
|
||||
json={
|
||||
"id": "chatcmpl-1859428879fc791b17d73",
|
||||
"object": "chat.completion",
|
||||
"created": 1754506683,
|
||||
"model": "claude-4-sonnet",
|
||||
"system_fingerprint": "heroku-inf-cp42st",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"refusal": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "tooluse_dV3Vtnb-S9-Z_YFicSv2Gw",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_current_weather",
|
||||
"arguments": "{\"location\":\"Portland, OR\"}"
|
||||
}
|
||||
}
|
||||
],
|
||||
"content": "Let me check the current weather in Portland for you."
|
||||
},
|
||||
"finish_reason": "tool_calls"
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 354,
|
||||
"completion_tokens": 69,
|
||||
"total_tokens": 423
|
||||
}
|
||||
},
|
||||
status_code=200,
|
||||
)
|
||||
|
||||
response = completion(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "What's the weather in Portland?"}],
|
||||
tools=[{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_current_weather",
|
||||
"description": "Get the current weather in a given location",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {
|
||||
"type": "string",
|
||||
"description": "The city and state, e.g. Portland, OR"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"location"
|
||||
]
|
||||
}
|
||||
}
|
||||
}],
|
||||
tool_choice="auto",
|
||||
)
|
||||
print(response)
|
||||
assert response.choices[0].message.content == "Let me check the current weather in Portland for you."
|
||||
assert response.choices[0].message.tool_calls[0].id == "tooluse_dV3Vtnb-S9-Z_YFicSv2Gw"
|
||||
assert response.choices[0].message.tool_calls[0].type == "function"
|
||||
assert response.choices[0].message.tool_calls[0].function.name == "get_current_weather"
|
||||
assert response.choices[0].message.tool_calls[0].function.arguments == "{\"location\":\"Portland, OR\"}"
|
||||
|
||||
assert response.usage.prompt_tokens == 354
|
||||
assert response.usage.completion_tokens == 69
|
||||
assert response.usage.total_tokens == 423
|
||||
Loading…
Add table
Reference in a new issue