From 35fe30aabd817fe9956686d2e75a34a068c1693b Mon Sep 17 00:00:00 2001 From: dhilip_binny <33201880+DhilipBinny@users.noreply.github.com> Date: Tue, 17 Mar 2026 00:27:57 +0800 Subject: [PATCH] fix: ensure JSON instruction in prompts for json_object response format (#3559) (#4271) --- mem0/memory/main.py | 7 + mem0/memory/utils.py | 25 ++++ tests/memory/test_json_prompt_fix.py | 214 +++++++++++++++++++++++++++ tests/test_main.py | 2 +- 4 files changed, 247 insertions(+), 1 deletion(-) create mode 100644 tests/memory/test_json_prompt_fix.py diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 62438502a..40fd3f2df 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -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}], diff --git a/mem0/memory/utils.py b/mem0/memory/utils.py index 8c11705c8..b451c0f6e 100644 --- a/mem0/memory/utils.py +++ b/mem0/memory/utils.py @@ -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: diff --git a/tests/memory/test_json_prompt_fix.py b/tests/memory/test_json_prompt_fix.py new file mode 100644 index 000000000..25c3c9a25 --- /dev/null +++ b/tests/memory/test_json_prompt_fix.py @@ -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 diff --git a/tests/test_main.py b/tests/test_main.py index 55ad661ff..e179f3715 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -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"}