fix(llms): send max_completion_tokens for the GPT-5 family across providers (#5547)

This commit is contained in:
Davide Leopardi
2026-06-15 08:34:19 +02:00
committed by GitHub
parent 3951ad4705
commit de471799d1
6 changed files with 144 additions and 3 deletions
+4 -1
View File
@@ -76,9 +76,12 @@ class AzureOpenAIStructuredLLM(LLMBase):
"model": self.config.model,
"messages": messages,
"temperature": self.config.temperature,
"max_tokens": self.config.max_tokens,
"top_p": self.config.top_p,
}
if self._uses_max_completion_tokens(self.config.model):
params["max_completion_tokens"] = self.config.max_tokens
else:
params["max_tokens"] = self.config.max_tokens
if response_format:
params["response_format"] = response_format
if tools:
+25 -1
View File
@@ -79,6 +79,25 @@ class LLMBase(ABC):
return False
def _uses_max_completion_tokens(self, model: str) -> bool:
"""
Check if the model expects ``max_completion_tokens`` instead of ``max_tokens``.
The whole GPT-5 family (gpt-5.4-mini, gpt-5.4-nano, gpt-5.5, ...) rejects the
legacy ``max_tokens`` parameter on the Chat Completions API and requires
``max_completion_tokens``. Older models (gpt-4.x, gpt-3.5, etc.) still accept
``max_tokens``.
Args:
model: The model name to check
Returns:
bool: True if the model requires ``max_completion_tokens``
"""
# Strip provider prefixes (e.g. "openai/gpt-5.4-mini" -> "gpt-5.4-mini")
base_model = (model or "").lower().rsplit("/", 1)[-1]
return base_model.startswith("gpt-5")
def _get_supported_params(self, **kwargs) -> Dict:
"""
Get parameters that are supported by the current model.
@@ -141,10 +160,15 @@ class LLMBase(ABC):
"""
params = {
"temperature": self.config.temperature,
"max_tokens": self.config.max_tokens,
"top_p": self.config.top_p,
}
model = getattr(self.config, "model", "")
if self._uses_max_completion_tokens(model):
params["max_completion_tokens"] = self.config.max_tokens
else:
params["max_tokens"] = self.config.max_tokens
# Add provider-specific parameters from kwargs
params.update(kwargs)
+4 -1
View File
@@ -74,9 +74,12 @@ class LiteLLM(LLMBase):
"model": self.config.model,
"messages": messages,
"temperature": self.config.temperature,
"max_tokens": self.config.max_tokens,
"top_p": self.config.top_p,
}
if self._uses_max_completion_tokens(self.config.model):
params["max_completion_tokens"] = self.config.max_tokens
else:
params["max_tokens"] = self.config.max_tokens
if response_format:
params["response_format"] = response_format
if tools: # TODO: Remove tools if no issues found with new memory addition logic
@@ -119,6 +119,44 @@ def test_generate_response_without_tools(mock_azure_openai):
assert response == "Hello there!"
@mock.patch("mem0.llms.azure_openai_structured.AzureOpenAI")
def test_generate_response_gpt5_uses_max_completion_tokens(mock_azure_openai):
mock_client = Mock()
mock_azure_openai.return_value = mock_client
config = DummyConfig(model="gpt-5.4-mini", azure_kwargs=DummyAzureKwargs(api_key="real-key"))
llm = AzureOpenAIStructuredLLM(config)
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="Hi"))]
mock_client.chat.completions.create.return_value = mock_response
llm.generate_response([{"role": "user", "content": "Hi"}])
_, kwargs = mock_client.chat.completions.create.call_args
assert kwargs["max_completion_tokens"] == 256
assert "max_tokens" not in kwargs
@mock.patch("mem0.llms.azure_openai_structured.AzureOpenAI")
def test_generate_response_legacy_model_uses_max_tokens(mock_azure_openai):
mock_client = Mock()
mock_azure_openai.return_value = mock_client
config = DummyConfig(model="gpt-4.1", azure_kwargs=DummyAzureKwargs(api_key="real-key"))
llm = AzureOpenAIStructuredLLM(config)
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="Hi"))]
mock_client.chat.completions.create.return_value = mock_response
llm.generate_response([{"role": "user", "content": "Hi"}])
_, kwargs = mock_client.chat.completions.create.call_args
assert kwargs["max_tokens"] == 256
assert "max_completion_tokens" not in kwargs
@mock.patch("mem0.llms.azure_openai_structured.AzureOpenAI")
def test_generate_response_with_tools(mock_azure_openai):
mock_client = Mock()
+34
View File
@@ -44,6 +44,40 @@ def test_generate_response_without_tools(mock_litellm):
assert response == "I'm doing well, thank you for asking!"
def test_generate_response_gpt5_uses_max_completion_tokens(mock_litellm):
config = BaseLlmConfig(model="gpt-5.4-mini", temperature=0.7, max_tokens=100, top_p=1)
llm = litellm.LiteLLM(config)
messages = [{"role": "user", "content": "Hello"}]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="Hi"))]
mock_litellm.completion.return_value = mock_response
mock_litellm.supports_function_calling.return_value = True
llm.generate_response(messages)
_, kwargs = mock_litellm.completion.call_args
assert kwargs["max_completion_tokens"] == 100
assert "max_tokens" not in kwargs
def test_generate_response_legacy_model_uses_max_tokens(mock_litellm):
config = BaseLlmConfig(model="gpt-4.1", temperature=0.7, max_tokens=100, top_p=1)
llm = litellm.LiteLLM(config)
messages = [{"role": "user", "content": "Hello"}]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="Hi"))]
mock_litellm.completion.return_value = mock_response
mock_litellm.supports_function_calling.return_value = True
llm.generate_response(messages)
_, kwargs = mock_litellm.completion.call_args
assert kwargs["max_tokens"] == 100
assert "max_completion_tokens" not in kwargs
def test_generate_response_with_tools(mock_litellm):
config = BaseLlmConfig(model="gpt-4.1-nano-2025-04-14", temperature=0.7, max_tokens=100, top_p=1)
llm = litellm.LiteLLM(config)
+39
View File
@@ -375,6 +375,45 @@ def test_is_reasoning_model_override_generates_correct_params(mock_openai_client
assert "temperature" not in call_kwargs
def test_gpt5_uses_max_completion_tokens(mock_openai_client):
"""gpt-5.x (non-reasoning) must send max_completion_tokens, not max_tokens.
The GPT-5 family rejects the legacy max_tokens param on Chat Completions and
requires max_completion_tokens. Regression test for
https://github.com/mem0ai/mem0/issues/5054
"""
config = OpenAIConfig(model="gpt-5.4-mini", temperature=0.7, max_tokens=100, top_p=1.0)
llm = OpenAILLM(config)
messages = [{"role": "user", "content": "Hello"}]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="ok"))]
mock_openai_client.chat.completions.create.return_value = mock_response
llm.generate_response(messages)
call_kwargs = mock_openai_client.chat.completions.create.call_args[1]
assert call_kwargs.get("max_completion_tokens") == 100
assert "max_tokens" not in call_kwargs
def test_gpt4_uses_max_tokens(mock_openai_client):
"""Older models (gpt-4.x) keep using max_tokens — guards against regressions."""
config = OpenAIConfig(model="gpt-4.1-nano-2025-04-14", temperature=0.7, max_tokens=100, top_p=1.0)
llm = OpenAILLM(config)
messages = [{"role": "user", "content": "Hello"}]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="ok"))]
mock_openai_client.chat.completions.create.return_value = mock_response
llm.generate_response(messages)
call_kwargs = mock_openai_client.chat.completions.create.call_args[1]
assert call_kwargs.get("max_tokens") == 100
assert "max_completion_tokens" not in call_kwargs
def test_callback_with_tools(mock_openai_client):
mock_callback = Mock()
config = OpenAIConfig(model="gpt-4.1-nano-2025-04-14", response_callback=mock_callback)