mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(fireworks_ai): translate NIM/vLLM extras on the text completion path
Mirror the chat extras translation for /v1/completions, adapted to the typed OpenAI SDK: anything completions.create() rejects (reasoning_effort, response_format, fireworks-native extras) rides inside extra_body, which the SDK merges server-side. Top-level reasoning_effort and response_format are moved into extra_body (they raised TypeError before), truncate aliases, chat_template_kwargs effort keys, and guided_* resolve into extra_body fields, and the strip set removes the rest. Verified live: /v1/completions rejects prompt_truncate_len, so both truncate names are stripped on this path rather than renamed.
This commit is contained in:
parent
0f15b471c4
commit
2cf5b04ace
2 changed files with 323 additions and 1 deletions
|
|
@ -1,11 +1,24 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.llms.openai import AllMessageValues, OpenAITextCompletionUserMessage
|
||||
from litellm.utils import supports_reasoning
|
||||
|
||||
from ...base_llm.completion.transformation import BaseTextCompletionConfig
|
||||
from ...openai.completion.utils import _transform_prompt
|
||||
from ..chat.transformation import (
|
||||
_EFFORT_KWARG_KEYS,
|
||||
_NIM_VLLM_STRIP_PARAMS,
|
||||
FireworksAIConfig,
|
||||
_effort_from_chat_template_kwargs,
|
||||
)
|
||||
from ..common_utils import FireworksAIMixin
|
||||
|
||||
_TEXT_COMPLETION_STRIP_PARAMS: Final = (
|
||||
frozenset({"truncate_prompt_tokens", "prompt_truncate_len"}) | _NIM_VLLM_STRIP_PARAMS
|
||||
)
|
||||
|
||||
|
||||
class FireworksAITextCompletionConfig(FireworksAIMixin, BaseTextCompletionConfig):
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
|
|
@ -41,6 +54,107 @@ class FireworksAITextCompletionConfig(FireworksAIMixin, BaseTextCompletionConfig
|
|||
optional_params[k] = v
|
||||
return optional_params
|
||||
|
||||
def map_extra_body_params(
|
||||
self, optional_params: Mapping[str, object], model: str
|
||||
) -> dict: # mutable-ok: returned dict is spread into the OpenAI SDK call as kwargs
|
||||
raw_extra_body: Final = optional_params.get("extra_body")
|
||||
initial_body: Final = (
|
||||
dict(raw_extra_body) if isinstance(raw_extra_body, dict) else {} # mutable-ok: JSON request body
|
||||
)
|
||||
stripped_body: Final = self._strip_unsupported_params(initial_body, model)
|
||||
moved_body: Final = self._move_native_params_into_extra_body(stripped_body, optional_params)
|
||||
effort_body: Final = self._translate_chat_template_kwargs(moved_body, optional_params, model)
|
||||
final_body: Final = self._translate_guided_into_extra_body(effort_body, optional_params)
|
||||
base: Final = { # mutable-ok: JSON request body
|
||||
k: v for k, v in optional_params.items() if k not in ("extra_body", "response_format", "reasoning_effort")
|
||||
}
|
||||
if final_body:
|
||||
base["extra_body"] = final_body
|
||||
return base
|
||||
|
||||
@staticmethod
|
||||
def _strip_unsupported_params(
|
||||
extra_body: Mapping[str, object], model: str
|
||||
) -> dict: # mutable-ok: JSON request body
|
||||
stripped: Final = tuple(sorted(k for k in extra_body if k in _TEXT_COMPLETION_STRIP_PARAMS))
|
||||
if stripped:
|
||||
verbose_logger.debug(
|
||||
"fireworks_ai does not support NIM/vLLM params %s for model=%s; dropping them from the request.",
|
||||
stripped,
|
||||
model,
|
||||
)
|
||||
return { # mutable-ok: JSON request body
|
||||
k: v for k, v in extra_body.items() if k not in _TEXT_COMPLETION_STRIP_PARAMS
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _move_native_params_into_extra_body(
|
||||
extra_body: Mapping[str, object], optional_params: Mapping[str, object]
|
||||
) -> dict: # mutable-ok: JSON request body
|
||||
moved: Final = dict(extra_body) # mutable-ok: JSON request body
|
||||
for key in ("response_format", "reasoning_effort"):
|
||||
value = optional_params.get(key)
|
||||
if value is None:
|
||||
continue
|
||||
if key in moved:
|
||||
verbose_logger.debug("fireworks_ai overriding extra_body.%s with the top-level %s.", key, key)
|
||||
moved[key] = value
|
||||
return moved
|
||||
|
||||
def _translate_chat_template_kwargs(
|
||||
self, extra_body: Mapping[str, object], optional_params: Mapping[str, object], model: str
|
||||
) -> dict: # mutable-ok: JSON request body
|
||||
chat_template_kwargs: Final = extra_body.get("chat_template_kwargs")
|
||||
if chat_template_kwargs is None:
|
||||
return dict(extra_body) # mutable-ok: JSON request body
|
||||
result: Final = { # mutable-ok: JSON request body
|
||||
k: v for k, v in extra_body.items() if k != "chat_template_kwargs"
|
||||
}
|
||||
if not isinstance(chat_template_kwargs, dict):
|
||||
verbose_logger.debug(
|
||||
"fireworks_ai dropping chat_template_kwargs for model=%s; expected an object, got %s.",
|
||||
model,
|
||||
type(chat_template_kwargs).__name__,
|
||||
)
|
||||
return result
|
||||
other_keys: Final = tuple(sorted(k for k in chat_template_kwargs if k not in _EFFORT_KWARG_KEYS))
|
||||
if other_keys:
|
||||
verbose_logger.debug(
|
||||
"fireworks_ai does not support chat_template_kwargs keys %s for model=%s; dropping them.",
|
||||
other_keys,
|
||||
model,
|
||||
)
|
||||
effort: Final = _effort_from_chat_template_kwargs(chat_template_kwargs)
|
||||
if effort is None:
|
||||
return result
|
||||
if "reasoning_effort" in result or "thinking" in optional_params:
|
||||
verbose_logger.debug(
|
||||
"fireworks_ai ignoring chat_template_kwargs; explicit reasoning_effort/thinking takes precedence."
|
||||
)
|
||||
return result
|
||||
if not supports_reasoning(model=model, custom_llm_provider="fireworks_ai"):
|
||||
verbose_logger.debug(
|
||||
"fireworks_ai model %r does not support reasoning; dropping chat_template_kwargs effort keys.",
|
||||
model,
|
||||
)
|
||||
return result
|
||||
return {**result, "reasoning_effort": effort} # mutable-ok: JSON request body
|
||||
|
||||
@staticmethod
|
||||
def _translate_guided_into_extra_body(
|
||||
extra_body: Mapping[str, object], optional_params: Mapping[str, object]
|
||||
) -> dict: # mutable-ok: JSON request body
|
||||
guided_response_format: Final = FireworksAIConfig._translate_guided_params(extra_body, optional_params)
|
||||
remaining: Final = { # mutable-ok: JSON request body
|
||||
k: v for k, v in extra_body.items() if k not in ("guided_json", "guided_grammar", "guided_choice")
|
||||
}
|
||||
if guided_response_format:
|
||||
return { # mutable-ok: JSON request body
|
||||
**remaining,
|
||||
guided_response_format[0][0]: guided_response_format[0][1],
|
||||
}
|
||||
return remaining
|
||||
|
||||
def transform_text_completion_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -48,6 +162,7 @@ class FireworksAITextCompletionConfig(FireworksAIMixin, BaseTextCompletionConfig
|
|||
optional_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
translated_params: Final = self.map_extra_body_params(optional_params=optional_params, model=model)
|
||||
prompt: Final = _transform_prompt(messages=messages)
|
||||
|
||||
if not model.startswith("accounts/") and "#" not in model:
|
||||
|
|
@ -56,6 +171,6 @@ class FireworksAITextCompletionConfig(FireworksAIMixin, BaseTextCompletionConfig
|
|||
data: Final = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
**optional_params,
|
||||
**translated_params,
|
||||
}
|
||||
return data
|
||||
|
|
|
|||
|
|
@ -0,0 +1,207 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
from litellm.llms.fireworks_ai.completion.transformation import (
|
||||
FireworksAITextCompletionConfig,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def force_local_model_cost(monkeypatch):
|
||||
"""Force local model cost map usage for all tests in this file."""
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
|
||||
|
||||
litellm.model_cost = get_model_cost_map(url=litellm.model_cost_map_url)
|
||||
|
||||
|
||||
_REASONING_MODEL = "fireworks_ai/accounts/fireworks/models/glm-5p1"
|
||||
_NON_REASONING_MODEL = "fireworks_ai/accounts/fireworks/models/llama-v3-70b-instruct"
|
||||
|
||||
|
||||
def test_map_extra_body_params_strips_truncate_params():
|
||||
"""
|
||||
prompt_truncate_len is accepted on chat completions but rejected by
|
||||
/v1/completions ("Extra inputs are not permitted"), so both the NIM/vLLM
|
||||
name and the Fireworks name must be stripped on the text completion path.
|
||||
"""
|
||||
config = FireworksAITextCompletionConfig()
|
||||
result = config.map_extra_body_params(
|
||||
{"extra_body": {"truncate_prompt_tokens": 4096, "prompt_truncate_len": 2048}},
|
||||
_REASONING_MODEL,
|
||||
)
|
||||
assert result == {}
|
||||
|
||||
|
||||
def test_map_extra_body_params_chat_template_kwargs_effort():
|
||||
config = FireworksAITextCompletionConfig()
|
||||
disabled = config.map_extra_body_params(
|
||||
{"extra_body": {"chat_template_kwargs": {"enable_thinking": False}}},
|
||||
_REASONING_MODEL,
|
||||
)
|
||||
assert disabled == {"extra_body": {"reasoning_effort": "none"}}
|
||||
|
||||
enabled = config.map_extra_body_params(
|
||||
{"extra_body": {"chat_template_kwargs": {"enable_thinking": True}}},
|
||||
_REASONING_MODEL,
|
||||
)
|
||||
assert enabled == {}
|
||||
|
||||
budget = config.map_extra_body_params(
|
||||
{"extra_body": {"chat_template_kwargs": {"reasoning_budget": 512}}},
|
||||
_REASONING_MODEL,
|
||||
)
|
||||
assert budget == {"extra_body": {"reasoning_effort": 512}}
|
||||
|
||||
low = config.map_extra_body_params(
|
||||
{"extra_body": {"chat_template_kwargs": {"low_effort": True}}},
|
||||
_REASONING_MODEL,
|
||||
)
|
||||
assert low == {"extra_body": {"reasoning_effort": "low"}}
|
||||
|
||||
|
||||
def test_map_extra_body_params_chat_template_kwargs_dropped_for_non_reasoning_model():
|
||||
config = FireworksAITextCompletionConfig()
|
||||
result = config.map_extra_body_params(
|
||||
{"extra_body": {"chat_template_kwargs": {"reasoning_budget": 512}}},
|
||||
_NON_REASONING_MODEL,
|
||||
)
|
||||
assert result == {}
|
||||
|
||||
|
||||
def test_map_extra_body_params_top_level_reasoning_effort_moves_into_extra_body():
|
||||
"""
|
||||
The OpenAI SDK completions.create() rejects a top-level reasoning_effort
|
||||
kwarg, so it must ride inside extra_body (and win over kwargs-derived effort).
|
||||
"""
|
||||
config = FireworksAITextCompletionConfig()
|
||||
result = config.map_extra_body_params(
|
||||
{
|
||||
"reasoning_effort": "high",
|
||||
"extra_body": {"chat_template_kwargs": {"enable_thinking": False}},
|
||||
},
|
||||
_REASONING_MODEL,
|
||||
)
|
||||
assert result == {"extra_body": {"reasoning_effort": "high"}}
|
||||
assert "reasoning_effort" not in {
|
||||
k for k in result if k != "extra_body"
|
||||
}
|
||||
|
||||
|
||||
def test_map_extra_body_params_top_level_response_format_moves_into_extra_body():
|
||||
config = FireworksAITextCompletionConfig()
|
||||
native = {"type": "json_object"}
|
||||
result = config.map_extra_body_params(
|
||||
{
|
||||
"response_format": native,
|
||||
"extra_body": {"response_format": {"type": "json_schema"}},
|
||||
},
|
||||
_REASONING_MODEL,
|
||||
)
|
||||
assert result == {"extra_body": {"response_format": native}}
|
||||
|
||||
|
||||
def test_map_extra_body_params_guided_params():
|
||||
config = FireworksAITextCompletionConfig()
|
||||
schema = {"type": "object", "properties": {"x": {"type": "string"}}}
|
||||
guided_json = config.map_extra_body_params(
|
||||
{"extra_body": {"guided_json": schema}}, _REASONING_MODEL
|
||||
)
|
||||
assert guided_json == {
|
||||
"extra_body": {
|
||||
"response_format": {
|
||||
"type": "json_schema",
|
||||
"json_schema": {"name": "response", "schema": schema},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
guided_choice = config.map_extra_body_params(
|
||||
{"extra_body": {"guided_choice": ["yes", "no"]}}, _REASONING_MODEL
|
||||
)
|
||||
assert guided_choice == {
|
||||
"extra_body": {
|
||||
"response_format": {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "choice",
|
||||
"schema": {"type": "string", "enum": ["yes", "no"]},
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def test_map_extra_body_params_guided_native_response_format_wins():
|
||||
config = FireworksAITextCompletionConfig()
|
||||
native = {"type": "json_object"}
|
||||
result = config.map_extra_body_params(
|
||||
{
|
||||
"response_format": native,
|
||||
"extra_body": {"guided_json": {"type": "object"}},
|
||||
},
|
||||
_REASONING_MODEL,
|
||||
)
|
||||
assert result == {"extra_body": {"response_format": native}}
|
||||
|
||||
|
||||
def test_map_extra_body_params_strips_unsupported_and_preserves_passthrough():
|
||||
config = FireworksAITextCompletionConfig()
|
||||
result = config.map_extra_body_params(
|
||||
{
|
||||
"extra_body": {
|
||||
"min_tokens": 10,
|
||||
"top_k": 40,
|
||||
"best_of": 2,
|
||||
"include_reasoning": True,
|
||||
"nvext": {"verbosity": 1},
|
||||
}
|
||||
},
|
||||
_REASONING_MODEL,
|
||||
)
|
||||
assert result == {"extra_body": {"min_tokens": 10, "top_k": 40}}
|
||||
|
||||
|
||||
def test_transform_text_completion_request_keeps_sdk_rejected_keys_in_extra_body():
|
||||
"""
|
||||
The request data is spread into the typed OpenAI SDK completions.create(),
|
||||
so anything the SDK does not accept (reasoning_effort, response_format,
|
||||
prompt_truncate_len, fireworks-native extras) must live inside extra_body
|
||||
or the call raises TypeError before it reaches Fireworks.
|
||||
"""
|
||||
config = FireworksAITextCompletionConfig()
|
||||
data = config.transform_text_completion_request(
|
||||
model="glm-5p1",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={
|
||||
"max_tokens": 10,
|
||||
"reasoning_effort": "low",
|
||||
"extra_body": {
|
||||
"truncate_prompt_tokens": 4096,
|
||||
"chat_template_kwargs": {"low_effort": True},
|
||||
"best_of": 2,
|
||||
"top_k": 40,
|
||||
},
|
||||
},
|
||||
headers={},
|
||||
)
|
||||
assert data["model"] == "accounts/fireworks/models/glm-5p1"
|
||||
assert data["prompt"] == "hi"
|
||||
assert data["max_tokens"] == 10
|
||||
assert "reasoning_effort" not in data
|
||||
assert data["extra_body"]["reasoning_effort"] == "low"
|
||||
assert data["extra_body"]["top_k"] == 40
|
||||
assert "truncate_prompt_tokens" not in data["extra_body"]
|
||||
assert "prompt_truncate_len" not in data["extra_body"]
|
||||
assert "chat_template_kwargs" not in data["extra_body"]
|
||||
assert "best_of" not in data["extra_body"]
|
||||
assert "response_format" not in data
|
||||
Loading…
Add table
Reference in a new issue