fix(azure): stop mutating and corrupting caller messages in content rewrite (#5731)

This commit is contained in:
Davide Leopardi
2026-06-22 08:15:46 +02:00
committed by GitHub
parent a48f34cf77
commit 8a786bf72d
4 changed files with 157 additions and 11 deletions
+27 -6
View File
@@ -1,3 +1,4 @@
import copy
import json
import os
from typing import Dict, List, Optional, Union
@@ -69,6 +70,26 @@ class AzureOpenAILLM(LLMBase):
default_headers=default_headers,
)
@staticmethod
def _rewrite_assistant_keyword(messages):
"""
Return a copy of ``messages`` with the word "assistant" replaced by "ai"
in the last message's textual content.
Azure's content management policy can flag the literal word "assistant",
which makes ``add`` fail (see issue #2636). The rewrite targets that
trigger without mutating the caller's messages and without assuming the
content is a string, so multimodal (list) content passes through untouched.
"""
if not messages:
return messages
messages = copy.deepcopy(messages)
last_content = messages[-1].get("content")
if isinstance(last_content, str):
messages[-1]["content"] = last_content.replace("assistant", "ai")
return messages
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
@@ -121,14 +142,14 @@ class AzureOpenAILLM(LLMBase):
str: The generated response.
"""
user_prompt = messages[-1]["content"]
user_prompt = user_prompt.replace("assistant", "ai")
messages[-1]["content"] = user_prompt
# Azure's "Indirect Attacks" content filter can flag the literal word
# "assistant" in the prompt, so it is rewritten to "ai" before the request.
# Work on a copy so the caller's messages are left untouched and string-only
# content is handled without breaking multimodal (list) content.
messages = self._rewrite_assistant_keyword(messages)
params = self._get_supported_params(messages=messages, **kwargs)
# Add model and messages
params.update({
"model": self.config.model,
+26 -5
View File
@@ -1,3 +1,4 @@
import copy
import json
import os
from typing import Dict, List, Optional
@@ -66,11 +67,11 @@ class AzureOpenAIStructuredLLM(LLMBase):
str: The generated response.
"""
user_prompt = messages[-1]["content"]
user_prompt = user_prompt.replace("assistant", "ai")
messages[-1]["content"] = user_prompt
# Azure's "Indirect Attacks" content filter can flag the literal word
# "assistant" in the prompt, so it is rewritten to "ai" before the request.
# Work on a copy so the caller's messages are left untouched and string-only
# content is handled without breaking multimodal (list) content.
messages = self._rewrite_assistant_keyword(messages)
is_reasoning = self._is_reasoning_model(self.config.model)
params = {
@@ -101,6 +102,26 @@ class AzureOpenAIStructuredLLM(LLMBase):
response = self.client.chat.completions.create(**params)
return self._parse_response(response, tools)
@staticmethod
def _rewrite_assistant_keyword(messages):
"""
Return a copy of ``messages`` with the word "assistant" replaced by "ai"
in the last message's textual content.
Azure's content management policy can flag the literal word "assistant",
which makes ``add`` fail (see issue #2636). The rewrite targets that
trigger without mutating the caller's messages and without assuming the
content is a string, so multimodal (list) content passes through untouched.
"""
if not messages:
return messages
messages = copy.deepcopy(messages)
last_content = messages[-1].get("content")
if isinstance(last_content, str):
messages[-1]["content"] = last_content.replace("assistant", "ai")
return messages
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
+46
View File
@@ -135,6 +135,52 @@ def test_generate_response_without_response_format(mock_openai_client):
assert response == "Why did the chicken cross the road?"
def test_generate_response_does_not_mutate_caller_messages(mock_openai_client):
config = AzureOpenAIConfig(model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P)
llm = AzureOpenAILLM(config)
messages = [{"role": "user", "content": "my assistant helps me schedule meetings"}]
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)
assert messages[-1]["content"] == "my assistant helps me schedule meetings"
def test_generate_response_rewrites_assistant_keyword_for_model_only(mock_openai_client):
config = AzureOpenAIConfig(model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P)
llm = AzureOpenAILLM(config)
messages = [{"role": "user", "content": "my assistant helps me"}]
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)
sent_messages = mock_openai_client.chat.completions.create.call_args[1]["messages"]
assert sent_messages[-1]["content"] == "my ai helps me"
assert messages[-1]["content"] == "my assistant helps me"
def test_generate_response_handles_multimodal_content(mock_openai_client):
config = AzureOpenAIConfig(model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P)
llm = AzureOpenAILLM(config)
messages = [{"role": "user", "content": [{"type": "text", "text": "describe my assistant"}]}]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="ok"))]
mock_openai_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages)
assert response == "ok"
sent_messages = mock_openai_client.chat.completions.create.call_args[1]["messages"]
assert sent_messages[-1]["content"] == [{"type": "text", "text": "describe my assistant"}]
def test_reasoning_model_with_reasoning_effort(mock_openai_client):
"""Test that reasoning_effort is passed to the API for Azure reasoning models."""
config = AzureOpenAIConfig(model="o3-mini", reasoning_effort="low")
@@ -264,3 +264,61 @@ def test_regular_model_sends_sampling_params(mock_azure_openai):
assert "max_tokens" in call_kwargs # standard sampling params still forwarded
assert "top_p" in call_kwargs
assert call_kwargs["model"] == "gpt-4o"
@mock.patch("mem0.llms.azure_openai_structured.AzureOpenAI")
def test_generate_response_does_not_mutate_caller_messages(mock_azure_openai):
mock_client = Mock()
mock_azure_openai.return_value = mock_client
config = DummyConfig(model="test-model", azure_kwargs=DummyAzureKwargs(api_key="real-key"))
llm = AzureOpenAIStructuredLLM(config)
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="ok"))]
mock_client.chat.completions.create.return_value = mock_response
messages = [{"role": "user", "content": "my assistant helps me schedule meetings"}]
llm.generate_response(messages)
assert messages[-1]["content"] == "my assistant helps me schedule meetings"
@mock.patch("mem0.llms.azure_openai_structured.AzureOpenAI")
def test_generate_response_rewrites_assistant_keyword_for_model_only(mock_azure_openai):
mock_client = Mock()
mock_azure_openai.return_value = mock_client
config = DummyConfig(model="test-model", azure_kwargs=DummyAzureKwargs(api_key="real-key"))
llm = AzureOpenAIStructuredLLM(config)
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="ok"))]
mock_client.chat.completions.create.return_value = mock_response
messages = [{"role": "user", "content": "my assistant helps me"}]
llm.generate_response(messages)
sent_messages = mock_client.chat.completions.create.call_args[1]["messages"]
assert sent_messages[-1]["content"] == "my ai helps me"
assert messages[-1]["content"] == "my assistant helps me"
@mock.patch("mem0.llms.azure_openai_structured.AzureOpenAI")
def test_generate_response_handles_multimodal_content(mock_azure_openai):
mock_client = Mock()
mock_azure_openai.return_value = mock_client
config = DummyConfig(model="test-model", azure_kwargs=DummyAzureKwargs(api_key="real-key"))
llm = AzureOpenAIStructuredLLM(config)
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="ok"))]
mock_client.chat.completions.create.return_value = mock_response
messages = [{"role": "user", "content": [{"type": "text", "text": "describe my assistant"}]}]
response = llm.generate_response(messages)
assert response == "ok"
sent_messages = mock_client.chat.completions.create.call_args[1]["messages"]
assert sent_messages[-1]["content"] == [{"type": "text", "text": "describe my assistant"}]