mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(fireworks_ai/): support passing response_format + tool call in same message
Addresses https://github.com/BerriAI/litellm/issues/7135
This commit is contained in:
parent
feefd18614
commit
6a30dc6929
9 changed files with 182 additions and 45 deletions
|
|
@ -631,6 +631,7 @@ _openai_like_providers: List = [
|
|||
"predibase",
|
||||
"databricks",
|
||||
"watsonx",
|
||||
"fireworks_ai",
|
||||
] # private helper. similar to openai but require some custom auth / endpoint handling, so can't use the openai sdk
|
||||
# well supported replicate llms
|
||||
replicate_models: List = [
|
||||
|
|
|
|||
|
|
@ -489,7 +489,6 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
|
|||
elif custom_llm_provider == "fireworks_ai":
|
||||
# fireworks is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.fireworks.ai/inference/v1
|
||||
(
|
||||
model,
|
||||
api_base,
|
||||
dynamic_api_key,
|
||||
) = litellm.FireworksAIConfig()._get_openai_compatible_provider_info(
|
||||
|
|
|
|||
|
|
@ -1,13 +1,20 @@
|
|||
import json
|
||||
import types
|
||||
from typing import Literal, Optional, Tuple, Union
|
||||
from typing import List, Literal, Optional, Tuple, Union
|
||||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionToolParam,
|
||||
ChatCompletionToolParamFunctionChunk,
|
||||
)
|
||||
|
||||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from ...openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
from ..embed.fireworks_ai_transformation import FireworksAIEmbeddingConfig
|
||||
|
||||
|
||||
class FireworksAIConfig(OpenAIGPTConfig):
|
||||
class FireworksAIConfig(OpenAILikeChatConfig):
|
||||
"""
|
||||
Reference: https://docs.fireworks.ai/api-reference/post-chatcompletions
|
||||
|
||||
|
|
@ -86,6 +93,7 @@ class FireworksAIConfig(OpenAIGPTConfig):
|
|||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
replace_max_completion_tokens_with_max_tokens: bool = True,
|
||||
) -> dict:
|
||||
|
||||
supported_openai_params = self.get_supported_openai_params(model=model)
|
||||
|
|
@ -97,28 +105,41 @@ class FireworksAIConfig(OpenAIGPTConfig):
|
|||
else:
|
||||
# pass through the value of tool choice
|
||||
optional_params["tool_choice"] = value
|
||||
elif (
|
||||
param == "response_format" and value.get("type", None) == "json_schema"
|
||||
):
|
||||
optional_params["response_format"] = {
|
||||
"type": "json_object",
|
||||
"schema": value["json_schema"]["schema"],
|
||||
}
|
||||
elif param == "response_format":
|
||||
# fireworks ai does not allow response_format + tools in same request
|
||||
## 1. if tools are provided, add response_format as a tool (Similar to anthropic/bedrock)
|
||||
## 2. if tools are not provided, pass through the response_format
|
||||
if value.get("type", None) == "json_schema":
|
||||
if non_default_params.get("tools", None) is None:
|
||||
optional_params["response_format"] = {
|
||||
"type": "json_object",
|
||||
"schema": value["json_schema"]["schema"],
|
||||
}
|
||||
elif non_default_params.get("tools", None) is not None:
|
||||
tool = self._create_json_tool_call_for_response_format(
|
||||
json_schema=value["json_schema"]["schema"]
|
||||
)
|
||||
optional_params = self._add_tools_to_optional_params(
|
||||
optional_params, [tool]
|
||||
)
|
||||
optional_params["json_mode"] = True
|
||||
elif param == "max_completion_tokens":
|
||||
optional_params["max_tokens"] = value
|
||||
elif param == "tools":
|
||||
optional_params = self._add_tools_to_optional_params(
|
||||
optional_params, value
|
||||
)
|
||||
elif param in supported_openai_params:
|
||||
if value is not None:
|
||||
optional_params[param] = value
|
||||
return optional_params
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, model: str, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> Tuple[str, Optional[str], Optional[str]]:
|
||||
if FireworksAIEmbeddingConfig().is_fireworks_embedding_model(model=model):
|
||||
# fireworks embeddings models do not require accounts/fireworks prefix https://docs.fireworks.ai/api-reference/creates-an-embedding-vector-representing-the-input-text
|
||||
pass
|
||||
elif not model.startswith("accounts/"):
|
||||
model = f"accounts/fireworks/models/{model}"
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: Optional[str] = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
api_base = (
|
||||
api_base
|
||||
or get_secret_str("FIREWORKS_API_BASE")
|
||||
|
|
@ -130,4 +151,22 @@ class FireworksAIConfig(OpenAIGPTConfig):
|
|||
or get_secret_str("FIREWORKSAI_API_KEY")
|
||||
or get_secret_str("FIREWORKS_AI_TOKEN")
|
||||
)
|
||||
return model, api_base, dynamic_api_key
|
||||
return api_base, dynamic_api_key
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
if not model.startswith("accounts/"):
|
||||
model = f"accounts/fireworks/models/{model}"
|
||||
return super().transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -40,7 +40,7 @@ def get_base_model_for_pricing(model_name: str) -> str:
|
|||
if params_billion <= 16.0:
|
||||
return "fireworks-ai-up-to-16b"
|
||||
elif params_billion <= 80.0:
|
||||
return "fireworks-ai-16b-80b"
|
||||
return "fireworks-ai-16.1b-to-80b"
|
||||
|
||||
# If no matches, return the original model_name
|
||||
return "fireworks-ai-default"
|
||||
|
|
|
|||
|
|
@ -207,7 +207,6 @@ class OpenAILikeChatHandler(OpenAILikeBase):
|
|||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
print(f"e.response.text: {e.response.text}")
|
||||
raise OpenAILikeError(
|
||||
status_code=e.response.status_code,
|
||||
message=e.response.text,
|
||||
|
|
@ -215,7 +214,6 @@ class OpenAILikeChatHandler(OpenAILikeBase):
|
|||
except httpx.TimeoutException:
|
||||
raise OpenAILikeError(status_code=408, message="Timeout error occurred.")
|
||||
except Exception as e:
|
||||
print(f"e: {e}")
|
||||
raise OpenAILikeError(status_code=500, message=str(e))
|
||||
|
||||
return OpenAILikeChatConfig._transform_response(
|
||||
|
|
@ -249,8 +247,8 @@ class OpenAILikeChatHandler(OpenAILikeBase):
|
|||
api_key: Optional[str],
|
||||
logging_obj,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
acompletion=None,
|
||||
litellm_params=None,
|
||||
logger_fn=None,
|
||||
headers: Optional[dict] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
|
|
@ -274,27 +272,38 @@ class OpenAILikeChatHandler(OpenAILikeBase):
|
|||
)
|
||||
|
||||
stream: bool = optional_params.pop("stream", None) or False
|
||||
extra_body = optional_params.pop("extra_body", {})
|
||||
json_mode = optional_params.pop("json_mode", None)
|
||||
optional_params.pop("max_retries", None)
|
||||
if not fake_stream:
|
||||
optional_params["stream"] = stream
|
||||
|
||||
if messages is not None and custom_llm_provider is not None:
|
||||
provider_config = ProviderConfigManager.get_provider_chat_config(
|
||||
model=model, provider=LlmProviders(custom_llm_provider)
|
||||
)
|
||||
provider_config = ProviderConfigManager.get_provider_chat_config(
|
||||
model=model, provider=LlmProviders(custom_llm_provider)
|
||||
)
|
||||
|
||||
if messages is not None:
|
||||
|
||||
if isinstance(provider_config, OpenAIGPTConfig) or isinstance(
|
||||
provider_config, OpenAIConfig
|
||||
):
|
||||
messages = provider_config._transform_messages(messages)
|
||||
|
||||
data = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
**optional_params,
|
||||
**extra_body,
|
||||
}
|
||||
if isinstance(provider_config, OpenAILikeChatConfig):
|
||||
data = provider_config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
else: # ensures 'extra_body' is correctly handled
|
||||
data = OpenAILikeChatConfig().transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
OpenAI-like chat completion transformation
|
||||
"""
|
||||
|
||||
import json
|
||||
import types
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
|
|
@ -9,8 +10,14 @@ import httpx
|
|||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionAssistantMessage
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionAssistantMessage,
|
||||
ChatCompletionToolParam,
|
||||
ChatCompletionToolParamFunctionChunk,
|
||||
)
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
from ....utils import _remove_additional_properties, _remove_strict_from_schema
|
||||
|
|
@ -30,6 +37,50 @@ class OpenAILikeChatConfig(OpenAIGPTConfig):
|
|||
) # vllm does not require an api key
|
||||
return api_base, dynamic_api_key
|
||||
|
||||
def _create_json_tool_call_for_response_format(
|
||||
self,
|
||||
json_schema: Optional[dict] = None,
|
||||
) -> ChatCompletionToolParam:
|
||||
"""
|
||||
Handles creating a tool call for getting responses in JSON format.
|
||||
|
||||
Args:
|
||||
json_schema (Optional[dict]): The JSON schema the response should be in
|
||||
|
||||
Returns:
|
||||
ChatCompletionToolParam: The tool call to send to an OpenAI-like provider to get responses in JSON format
|
||||
"""
|
||||
_input_schema: dict = {}
|
||||
|
||||
if json_schema is None:
|
||||
# Anthropic raises a 400 BadRequest error if properties is passed as None
|
||||
# see usage with additionalProperties (Example 5) https://github.com/anthropics/anthropic-cookbook/blob/main/tool_use/extracting_structured_json.ipynb
|
||||
_input_schema["additionalProperties"] = True
|
||||
_input_schema["properties"] = {}
|
||||
else:
|
||||
_input_schema["properties"] = {"values": json_schema}
|
||||
|
||||
_tool = ChatCompletionToolParam(
|
||||
type="function",
|
||||
function=ChatCompletionToolParamFunctionChunk(
|
||||
name=RESPONSE_FORMAT_TOOL_NAME,
|
||||
parameters=_input_schema,
|
||||
),
|
||||
)
|
||||
return _tool
|
||||
|
||||
def _add_tools_to_optional_params(
|
||||
self, optional_params: dict, tools: List[ChatCompletionToolParam]
|
||||
) -> dict:
|
||||
if "tools" not in optional_params:
|
||||
optional_params["tools"] = tools
|
||||
else:
|
||||
optional_params["tools"] = [
|
||||
*optional_params["tools"],
|
||||
*tools,
|
||||
]
|
||||
return optional_params
|
||||
|
||||
@staticmethod
|
||||
def _convert_tool_response_to_message(
|
||||
message: ChatCompletionAssistantMessage, json_mode: bool
|
||||
|
|
@ -53,11 +104,31 @@ class OpenAILikeChatConfig(OpenAIGPTConfig):
|
|||
if _tool_calls is None or len(_tool_calls) != 1:
|
||||
return message
|
||||
|
||||
message["content"] = _tool_calls[0]["function"].get("arguments") or ""
|
||||
message["tool_calls"] = None
|
||||
if (
|
||||
"name" in _tool_calls[0]["function"]
|
||||
and _tool_calls[0]["function"]["name"] == RESPONSE_FORMAT_TOOL_NAME
|
||||
):
|
||||
message["content"] = _tool_calls[0]["function"].get("arguments") or ""
|
||||
message["tool_calls"] = None
|
||||
|
||||
return message
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
extra_body = optional_params.pop("extra_body", {})
|
||||
return {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
**optional_params,
|
||||
**extra_body,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _transform_response(
|
||||
model: str,
|
||||
|
|
@ -75,7 +146,6 @@ class OpenAILikeChatConfig(OpenAIGPTConfig):
|
|||
custom_llm_provider: str,
|
||||
base_model: Optional[str],
|
||||
) -> ModelResponse:
|
||||
print(f"response: {response}")
|
||||
response_json = response.json()
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
|
|
|
|||
|
|
@ -1524,7 +1524,10 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
or custom_llm_provider == "mistral"
|
||||
or custom_llm_provider == "openai"
|
||||
or custom_llm_provider == "together_ai"
|
||||
or custom_llm_provider in litellm.openai_compatible_providers
|
||||
or (
|
||||
custom_llm_provider in litellm.openai_compatible_providers
|
||||
and custom_llm_provider not in litellm._openai_like_providers
|
||||
)
|
||||
or "ft:gpt-3.5-turbo" in model # finetune gpt-3.5-turbo
|
||||
): # allow user to make an openai call with a custom base
|
||||
# note: if a user sets a custom base - we should ensure this works
|
||||
|
|
@ -2041,6 +2044,24 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
)
|
||||
return response
|
||||
response = model_response
|
||||
elif custom_llm_provider == "fireworks_ai":
|
||||
model_response = openai_like_chat_completion.completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=api_base,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
encoding=encoding,
|
||||
api_key=api_key,
|
||||
logging_obj=logging,
|
||||
custom_llm_provider="fireworks_ai",
|
||||
custom_prompt_dict=custom_prompt_dict,
|
||||
)
|
||||
|
||||
response = model_response
|
||||
elif custom_llm_provider == "databricks":
|
||||
api_base = (
|
||||
api_base # for databricks we check in get_llm_provider and pass in the api base from there
|
||||
|
|
|
|||
|
|
@ -1830,17 +1830,14 @@ def supports_function_calling(
|
|||
|
||||
def _supports_factory(model: str, custom_llm_provider: Optional[str], key: str) -> bool:
|
||||
"""
|
||||
Check if the given model supports function calling and return a boolean value.
|
||||
Check if the given model supports 'key' and return a boolean value.
|
||||
|
||||
Parameters:
|
||||
model (str): The model name to be checked.
|
||||
custom_llm_provider (Optional[str]): The provider to be checked.
|
||||
|
||||
Returns:
|
||||
bool: True if the model supports function calling, False otherwise.
|
||||
|
||||
Raises:
|
||||
Exception: If the given model is not found or there's an error in retrieval.
|
||||
bool: True if the model supports 'key', False otherwise.
|
||||
"""
|
||||
try:
|
||||
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
|
|
@ -1855,9 +1852,10 @@ def _supports_factory(model: str, custom_llm_provider: Optional[str], key: str)
|
|||
return True
|
||||
return False
|
||||
except Exception as e:
|
||||
raise Exception(
|
||||
verbose_logger.error(
|
||||
f"Model not found or error in checking {key} support. You passed model={model}, custom_llm_provider={custom_llm_provider}. Error: {str(e)}"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def supports_audio_input(model: str, custom_llm_provider: Optional[str] = None) -> bool:
|
||||
|
|
|
|||
|
|
@ -78,7 +78,7 @@ def test_map_response_format():
|
|||
class TestFireworksAIChatCompletion(BaseLLMChatTest):
|
||||
def get_base_completion_call_args(self) -> dict:
|
||||
return {
|
||||
"model": "fireworks_ai/accounts/fireworks/models/llama-v3p2-11b-vision-instruct"
|
||||
"model": "fireworks_ai/accounts/fireworks/models/llama-v3p3-70b-instruct"
|
||||
}
|
||||
|
||||
def test_tool_call_no_arguments(self, tool_call_no_arguments):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue