fix(azure): stop mutating and corrupting caller messages in content rewrite (#5731)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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"}]
|
||||
|
||||
Reference in New Issue
Block a user