This commit is contained in:
@@ -26,6 +26,7 @@ from mem0.memory.setup import mem0_dir, setup_config
|
||||
from mem0.memory.storage import SQLiteManager
|
||||
from mem0.memory.telemetry import MEM0_TELEMETRY, capture_event
|
||||
from mem0.memory.utils import (
|
||||
ensure_json_instruction,
|
||||
extract_json,
|
||||
get_fact_retrieval_messages,
|
||||
parse_messages,
|
||||
@@ -432,6 +433,9 @@ class Memory(MemoryBase):
|
||||
is_agent_memory = self._should_use_agent_memory_extraction(messages, metadata)
|
||||
system_prompt, user_prompt = get_fact_retrieval_messages(parsed_messages, is_agent_memory)
|
||||
|
||||
# Ensure 'json' appears in prompts for json_object response format compatibility
|
||||
system_prompt, user_prompt = ensure_json_instruction(system_prompt, user_prompt)
|
||||
|
||||
response = self.llm.generate_response(
|
||||
messages=[
|
||||
{"role": "system", "content": system_prompt},
|
||||
@@ -1460,6 +1464,9 @@ class AsyncMemory(MemoryBase):
|
||||
is_agent_memory = self._should_use_agent_memory_extraction(messages, metadata)
|
||||
system_prompt, user_prompt = get_fact_retrieval_messages(parsed_messages, is_agent_memory)
|
||||
|
||||
# Ensure 'json' appears in prompts for json_object response format compatibility
|
||||
system_prompt, user_prompt = ensure_json_instruction(system_prompt, user_prompt)
|
||||
|
||||
response = await asyncio.to_thread(
|
||||
self.llm.generate_response,
|
||||
messages=[{"role": "system", "content": system_prompt}, {"role": "user", "content": user_prompt}],
|
||||
|
||||
@@ -29,6 +29,31 @@ def get_fact_retrieval_messages_legacy(message):
|
||||
return FACT_RETRIEVAL_PROMPT, f"Input:\n{message}"
|
||||
|
||||
|
||||
def ensure_json_instruction(system_prompt, user_prompt):
|
||||
"""Ensure the word 'json' appears in the prompts when using json_object response format.
|
||||
|
||||
OpenAI's API requires the word 'json' to appear in the messages when
|
||||
response_format is set to {"type": "json_object"}. When users provide a
|
||||
custom_fact_extraction_prompt that doesn't include 'json', this causes a
|
||||
400 error. This function appends a JSON format instruction to the system
|
||||
prompt if 'json' is not already present in either prompt.
|
||||
|
||||
Args:
|
||||
system_prompt: The system prompt string
|
||||
user_prompt: The user prompt string
|
||||
|
||||
Returns:
|
||||
tuple: (system_prompt, user_prompt) with JSON instruction added if needed
|
||||
"""
|
||||
combined = (system_prompt + user_prompt).lower()
|
||||
if "json" not in combined:
|
||||
system_prompt += (
|
||||
"\n\nYou must return your response in valid JSON format "
|
||||
"with a 'facts' key containing an array of strings."
|
||||
)
|
||||
return system_prompt, user_prompt
|
||||
|
||||
|
||||
def parse_messages(messages):
|
||||
response = ""
|
||||
for msg in messages:
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
"""
|
||||
Tests for issue #3559: Custom prompts crash with response_format json_object
|
||||
when the word 'json' is not present in the prompt.
|
||||
|
||||
OpenAI API requires the word 'json' to appear in messages when using
|
||||
response_format: {"type": "json_object"}. Custom fact extraction prompts
|
||||
may not include this word, causing BadRequestError.
|
||||
|
||||
This tests the ensure_json_instruction utility function and verifies
|
||||
the fix is applied in both sync and async code paths.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from mem0.memory.utils import ensure_json_instruction
|
||||
|
||||
|
||||
class TestEnsureJsonInstruction:
|
||||
"""Tests for the ensure_json_instruction utility function."""
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Core behavior: append when missing, skip when present
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_appends_when_json_missing_from_both_prompts(self):
|
||||
"""When neither prompt contains 'json', instruction is appended to system prompt."""
|
||||
system, user = ensure_json_instruction(
|
||||
"Extract facts from the conversation and return them as a list.",
|
||||
"Input:\nuser: Hi my name is John",
|
||||
)
|
||||
assert "json" in system.lower()
|
||||
assert "facts" in system.lower()
|
||||
|
||||
def test_no_change_when_json_in_system_prompt(self):
|
||||
"""When system prompt already contains 'json', no modification."""
|
||||
original = "Extract facts and return in json format."
|
||||
system, user = ensure_json_instruction(original, "Input:\nuser: Hi")
|
||||
assert system == original
|
||||
|
||||
def test_no_change_when_json_in_user_prompt(self):
|
||||
"""When user prompt contains 'json', no modification to system prompt."""
|
||||
original_system = "Extract facts from the conversation."
|
||||
original_user = "Input (respond in json):\nuser: Hi"
|
||||
system, user = ensure_json_instruction(original_system, original_user)
|
||||
assert system == original_system
|
||||
|
||||
def test_user_prompt_never_modified(self):
|
||||
"""The user prompt should never be modified regardless of content."""
|
||||
original_user = "Input:\nuser: I like pizza"
|
||||
_, user = ensure_json_instruction("Extract facts.", original_user)
|
||||
assert user == original_user
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Case insensitivity
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_case_insensitive_lowercase(self):
|
||||
original = "Return results in json format."
|
||||
system, _ = ensure_json_instruction(original, "Input:\nuser: Hi")
|
||||
assert system == original
|
||||
|
||||
def test_case_insensitive_uppercase(self):
|
||||
original = "Return results in JSON format."
|
||||
system, _ = ensure_json_instruction(original, "Input:\nuser: Hi")
|
||||
assert system == original
|
||||
|
||||
def test_case_insensitive_mixed(self):
|
||||
original = "Return results in Json format."
|
||||
system, _ = ensure_json_instruction(original, "Input:\nuser: Hi")
|
||||
assert system == original
|
||||
|
||||
def test_case_insensitive_in_user_prompt(self):
|
||||
original_system = "Extract facts."
|
||||
system, _ = ensure_json_instruction(original_system, "Return JSON.\nuser: Hi")
|
||||
assert system == original_system
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Parametrized: various custom prompts
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"prompt,should_append",
|
||||
[
|
||||
# Prompts WITHOUT json — should append
|
||||
("Extract all facts from the conversation.", True),
|
||||
("You are a memory extractor. Return facts as a list.", True),
|
||||
("Analyze the input and find key information.", True),
|
||||
("Return data in structured format.", True),
|
||||
("List the user preferences.", True),
|
||||
# Prompts WITH json — should NOT append
|
||||
("Extract facts and return in json format.", False),
|
||||
("Return a json object with facts.", False),
|
||||
("Output must be valid JSON.", False),
|
||||
("Respond with a JSON array of facts.", False),
|
||||
("Format: json output expected.", False),
|
||||
],
|
||||
)
|
||||
def test_various_custom_prompts(self, prompt, should_append):
|
||||
user_prompt = "Input:\nuser: Hi my name is John"
|
||||
system, _ = ensure_json_instruction(prompt, user_prompt)
|
||||
|
||||
if should_append:
|
||||
assert system != prompt, f"Expected JSON instruction to be appended for: {prompt}"
|
||||
assert "json" in system.lower()
|
||||
else:
|
||||
assert system == prompt, f"Did not expect modification for: {prompt}"
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Edge cases
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_empty_system_prompt(self):
|
||||
"""Empty system prompt should get JSON instruction."""
|
||||
system, _ = ensure_json_instruction("", "Input:\nuser: test")
|
||||
assert "json" in system.lower()
|
||||
|
||||
def test_whitespace_only_system_prompt(self):
|
||||
"""Whitespace-only prompt should get JSON instruction."""
|
||||
system, _ = ensure_json_instruction(" \n ", "Input:\nuser: test")
|
||||
assert "json" in system.lower()
|
||||
|
||||
def test_preserves_original_prompt_content(self):
|
||||
"""The fix should only append, never modify the original prompt content."""
|
||||
original = "Extract all user preferences and habits from the conversation."
|
||||
system, _ = ensure_json_instruction(original, "Input:\nuser: I like pizza")
|
||||
assert system.startswith(original)
|
||||
assert len(system) > len(original)
|
||||
|
||||
def test_appended_instruction_mentions_facts_key(self):
|
||||
"""The appended instruction should guide the model to use the 'facts' key."""
|
||||
system, _ = ensure_json_instruction(
|
||||
"Extract information.", "Input:\nuser: test"
|
||||
)
|
||||
assert "facts" in system.lower()
|
||||
|
||||
def test_idempotent_when_already_has_json(self):
|
||||
"""Calling ensure_json_instruction twice doesn't double-append."""
|
||||
system1, user1 = ensure_json_instruction(
|
||||
"Extract facts.", "Input:\nuser: test"
|
||||
)
|
||||
system2, user2 = ensure_json_instruction(system1, user1)
|
||||
assert system1 == system2
|
||||
assert user1 == user2
|
||||
|
||||
def test_json_in_curly_braces_not_detected(self):
|
||||
"""A prompt with JSON-like structure but no 'json' word should get instruction.
|
||||
e.g. '{"facts": [...]}' contains the characters j,s,o,n but not the word 'json'."""
|
||||
prompt = 'Return format: {"facts": [...]}'
|
||||
# This contains the substring "json" inside the key name — let's check
|
||||
if "json" in prompt.lower():
|
||||
# If it does contain json, it won't be modified
|
||||
system, _ = ensure_json_instruction(prompt, "Input:\nuser: test")
|
||||
assert system == prompt
|
||||
else:
|
||||
system, _ = ensure_json_instruction(prompt, "Input:\nuser: test")
|
||||
assert system != prompt
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Default prompts verification
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_default_prompts_already_contain_json(self):
|
||||
"""Built-in prompts already contain 'json', so ensure_json_instruction is a no-op."""
|
||||
from mem0.configs.prompts import (
|
||||
FACT_RETRIEVAL_PROMPT,
|
||||
USER_MEMORY_EXTRACTION_PROMPT,
|
||||
AGENT_MEMORY_EXTRACTION_PROMPT,
|
||||
)
|
||||
|
||||
for name, prompt in [
|
||||
("FACT_RETRIEVAL_PROMPT", FACT_RETRIEVAL_PROMPT),
|
||||
("USER_MEMORY_EXTRACTION_PROMPT", USER_MEMORY_EXTRACTION_PROMPT),
|
||||
("AGENT_MEMORY_EXTRACTION_PROMPT", AGENT_MEMORY_EXTRACTION_PROMPT),
|
||||
]:
|
||||
assert "json" in prompt.lower(), (
|
||||
f"{name} should contain 'json' — "
|
||||
"if this fails, the default prompts have changed"
|
||||
)
|
||||
# ensure_json_instruction should be a no-op for defaults
|
||||
system, _ = ensure_json_instruction(prompt, "Input:\nuser: test")
|
||||
assert system == prompt, f"ensure_json_instruction modified {name} unexpectedly"
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Integration: verify fix is wired into both sync and async paths
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_fix_applied_in_sync_memory_class(self):
|
||||
"""Verify the ensure_json_instruction call exists in Memory._add_to_vector_store."""
|
||||
import inspect
|
||||
from mem0.memory.main import Memory
|
||||
|
||||
source = inspect.getsource(Memory._add_to_vector_store)
|
||||
assert "ensure_json_instruction" in source, (
|
||||
"ensure_json_instruction not found in Memory._add_to_vector_store (sync)"
|
||||
)
|
||||
|
||||
def test_fix_applied_in_async_memory_class(self):
|
||||
"""Verify the ensure_json_instruction call exists in AsyncMemory._add_to_vector_store."""
|
||||
import inspect
|
||||
from mem0.memory.main import AsyncMemory
|
||||
|
||||
source = inspect.getsource(AsyncMemory._add_to_vector_store)
|
||||
assert "ensure_json_instruction" in source, (
|
||||
"ensure_json_instruction not found in AsyncMemory._add_to_vector_store (async)"
|
||||
)
|
||||
|
||||
def test_import_exists_in_main(self):
|
||||
"""Verify ensure_json_instruction is imported in main.py."""
|
||||
import inspect
|
||||
import mem0.memory.main as main_module
|
||||
|
||||
source = inspect.getsource(main_module)
|
||||
assert "from mem0.memory.utils import" in source
|
||||
assert "ensure_json_instruction" in source
|
||||
+1
-1
@@ -60,7 +60,7 @@ def memory_custom_instance():
|
||||
|
||||
config = MemoryConfig(
|
||||
version="v1.1",
|
||||
custom_fact_extraction_prompt="custom prompt extracting memory",
|
||||
custom_fact_extraction_prompt="custom prompt extracting memory in json format",
|
||||
custom_update_memory_prompt="custom prompt determining memory update",
|
||||
)
|
||||
config.graph_store.config = {"some_config": "value"}
|
||||
|
||||
Reference in New Issue
Block a user