fix(llms): send max_completion_tokens for the GPT-5 family across providers (#5547)
This commit is contained in:
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user