mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
fix: use base model params for JSON providers
This commit is contained in:
parent
ee7c7e14f3
commit
b953695723
3 changed files with 123 additions and 37 deletions
|
|
@ -3,9 +3,8 @@ Dynamic configuration class generator for JSON-based providers.
|
|||
"""
|
||||
|
||||
from collections.abc import Coroutine
|
||||
from typing import Any, Final, Literal, overload
|
||||
from typing import Any, Final, Literal, Protocol, overload, runtime_checkable
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
handle_messages_with_content_list_to_str_conversion,
|
||||
)
|
||||
|
|
@ -17,6 +16,20 @@ from litellm.types.llms.openai import AllMessageValues
|
|||
from .json_loader import SimpleProviderConfig
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class BaseModelAwareConfig(Protocol):
|
||||
supports_base_model_hint: bool
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
base_model: str | None = None,
|
||||
) -> dict: ...
|
||||
|
||||
|
||||
def create_config_class(provider: SimpleProviderConfig):
|
||||
"""Generate config class dynamically from JSON configuration"""
|
||||
|
||||
|
|
@ -89,37 +102,31 @@ def create_config_class(provider: SimpleProviderConfig):
|
|||
|
||||
return api_base
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
"""Get supported OpenAI params, excluding tool-related params for models
|
||||
that don't support function calling."""
|
||||
def _get_supported_openai_params_for_model(self, model: str) -> list:
|
||||
from litellm.utils import supports_function_calling, supports_reasoning
|
||||
|
||||
supported_params: Final = super().get_supported_openai_params(model=model)
|
||||
tool_params: Final = ("tools", "tool_choice", "function_call", "functions", "parallel_tool_calls")
|
||||
params_without_tools: Final = tuple(
|
||||
param for param in super().get_supported_openai_params(model=model) if param not in tool_params
|
||||
)
|
||||
params_with_tools: Final = tuple(dict.fromkeys((*params_without_tools, *tool_params)))
|
||||
supported_params: Final = (
|
||||
params_with_tools
|
||||
if supports_function_calling(model=model, custom_llm_provider=provider.slug)
|
||||
else params_without_tools
|
||||
)
|
||||
if (
|
||||
supports_reasoning(model=model, custom_llm_provider=provider.slug)
|
||||
and "reasoning_effort" not in supported_params
|
||||
):
|
||||
return [*supported_params, "reasoning_effort"]
|
||||
return list(supported_params)
|
||||
|
||||
_supports_fc: Final = supports_function_calling(model=model, custom_llm_provider=provider.slug)
|
||||
|
||||
if not _supports_fc:
|
||||
tool_params: Final = [
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"function_call",
|
||||
"functions",
|
||||
"parallel_tool_calls",
|
||||
]
|
||||
for param in tool_params:
|
||||
if param in supported_params:
|
||||
supported_params.remove(param)
|
||||
verbose_logger.debug(
|
||||
"Model %s on provider %s does not support function calling — removed tool-related params from supported params.",
|
||||
model,
|
||||
provider.slug,
|
||||
)
|
||||
|
||||
_supports_reasoning: Final = supports_reasoning(model=model, custom_llm_provider=provider.slug)
|
||||
if _supports_reasoning and "reasoning_effort" not in supported_params:
|
||||
supported_params.append("reasoning_effort")
|
||||
|
||||
return supported_params
|
||||
def get_supported_openai_params(self, model: str, base_model: str | None = None) -> list:
|
||||
supported_params: Final = self._get_supported_openai_params_for_model(model)
|
||||
if not base_model or base_model == model:
|
||||
return supported_params
|
||||
return list(dict.fromkeys([*supported_params, *self._get_supported_openai_params_for_model(base_model)]))
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
|
|
@ -127,10 +134,11 @@ def create_config_class(provider: SimpleProviderConfig):
|
|||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
base_model: str | None = None,
|
||||
) -> dict:
|
||||
"""Apply parameter mappings and constraints"""
|
||||
|
||||
supported_params: Final = self.get_supported_openai_params(model)
|
||||
supported_params: Final = self.get_supported_openai_params(model, base_model=base_model)
|
||||
|
||||
# Apply supported params
|
||||
for param, value in non_default_params.items():
|
||||
|
|
@ -167,6 +175,7 @@ def create_config_class(provider: SimpleProviderConfig):
|
|||
def custom_llm_provider(self) -> str | None:
|
||||
return provider.slug
|
||||
|
||||
JSONProviderConfig.supports_base_model_hint = True
|
||||
return JSONProviderConfig
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -4714,12 +4714,24 @@ def get_optional_params(
|
|||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif provider_config is not None:
|
||||
optional_params = provider_config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
drop_params_value: Final = bool(drop_params)
|
||||
from litellm.llms.openai_like.dynamic_config import BaseModelAwareConfig
|
||||
|
||||
if isinstance(provider_config, BaseModelAwareConfig):
|
||||
optional_params = provider_config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=drop_params_value,
|
||||
base_model=base_model,
|
||||
)
|
||||
else:
|
||||
optional_params = provider_config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=drop_params_value,
|
||||
)
|
||||
else: # assume passing in params for openai-like api
|
||||
optional_params = litellm.OpenAILikeChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
|
|
|
|||
|
|
@ -19,6 +19,71 @@ def _isolate_generated_class_cache():
|
|||
dynamic_config._responses_config_cache.clear()
|
||||
|
||||
|
||||
class TestBaseModelParamSupport:
|
||||
TOOLS = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "lookup",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"query": {"type": "string"}},
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
def test_generated_chat_config_uses_base_model_for_supported_params(self, local_model_cost_map):
|
||||
config = dynamic_config.create_config_class(_provider("publicai", base_class="openai_gpt"))()
|
||||
|
||||
endpoint_params = config.get_supported_openai_params(model="ep-publicai")
|
||||
assert "tools" not in endpoint_params
|
||||
assert "reasoning_effort" not in endpoint_params
|
||||
|
||||
instruct_params = config.get_supported_openai_params(
|
||||
model="ep-publicai",
|
||||
base_model="publicai/allenai/Olmo-3-7B-Instruct",
|
||||
)
|
||||
assert "tools" in instruct_params
|
||||
assert "reasoning_effort" not in instruct_params
|
||||
|
||||
thinking_params = config.get_supported_openai_params(
|
||||
model="ep-publicai",
|
||||
base_model="publicai/allenai/Olmo-3-7B-Think",
|
||||
)
|
||||
assert "tools" in thinking_params
|
||||
assert "reasoning_effort" in thinking_params
|
||||
|
||||
def test_get_optional_params_passes_base_model_to_json_provider(self, local_model_cost_map):
|
||||
from litellm.utils import get_optional_params
|
||||
|
||||
optional_params = get_optional_params(
|
||||
model="ep-publicai",
|
||||
custom_llm_provider="publicai",
|
||||
tools=self.TOOLS,
|
||||
reasoning_effort="high",
|
||||
base_model="publicai/allenai/Olmo-3-7B-Think",
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
assert optional_params["tools"] == self.TOOLS
|
||||
assert optional_params["reasoning_effort"] == "high"
|
||||
assert "base_model" not in optional_params
|
||||
|
||||
def test_non_json_provider_does_not_receive_base_model_kwarg(self):
|
||||
from litellm.utils import get_optional_params
|
||||
|
||||
optional_params = get_optional_params(
|
||||
model="a2a/test-agent",
|
||||
custom_llm_provider="a2a",
|
||||
tools=self.TOOLS,
|
||||
base_model="publicai/allenai/Olmo-3-7B-Think",
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
assert "base_model" not in optional_params
|
||||
|
||||
|
||||
class TestClassCaching:
|
||||
def test_same_slug_returns_the_identical_class_object(self):
|
||||
provider = _provider("cache_same_slug")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue