From c1ee71ad3f5872befeb0898f8b5de33af77b231b Mon Sep 17 00:00:00 2001 From: parshvadaftari Date: Tue, 16 Sep 2025 04:14:38 +0530 Subject: [PATCH] Fixed json parsing acros differnet LLM providers --- mem0/llms/aws_bedrock.py | 30 ++++++++++++++++++------------ mem0/llms/ollama.py | 27 +++++++++++++++++++++------ mem0/memory/main.py | 23 +++++++++++++++++++++-- 3 files changed, 60 insertions(+), 20 deletions(-) diff --git a/mem0/llms/aws_bedrock.py b/mem0/llms/aws_bedrock.py index b29cf59b8..ce10fb9c8 100644 --- a/mem0/llms/aws_bedrock.py +++ b/mem0/llms/aws_bedrock.py @@ -12,6 +12,7 @@ except ImportError: from mem0.configs.llms.base import BaseLlmConfig from mem0.configs.llms.aws_bedrock import AWSBedrockConfig from mem0.llms.base import LLMBase +from mem0.memory.utils import extract_json logger = logging.getLogger(__name__) @@ -371,7 +372,7 @@ class AWSBedrockLLM(LLMBase): processed_response["tool_calls"].append( { "name": item["toolUse"]["name"], - "arguments": item["toolUse"]["input"], + "arguments": json.loads(extract_json(json.dumps(item["toolUse"]["input"]))), } ) @@ -575,21 +576,26 @@ class AWSBedrockLLM(LLMBase): return self._parse_response(response) else: - prompt = self._format_messages(messages) + # For other providers and legacy Amazon models (like Titan) + if self.provider == "amazon": + # Legacy Amazon models need string formatting, not array formatting + prompt = self._format_messages_generic(messages) + else: + prompt = self._format_messages(messages) input_body = self._prepare_input(prompt) - # Convert to JSON - body = json.dumps(input_body) + # Convert to JSON + body = json.dumps(input_body) - # Make API call - response = self.client.invoke_model( - body=body, - modelId=self.config.model, - accept="application/json", - contentType="application/json", - ) + # Make API call + response = self.client.invoke_model( + body=body, + modelId=self.config.model, + accept="application/json", + contentType="application/json", + ) - return self._parse_response(response) + return self._parse_response(response) def list_available_models(self) -> List[Dict[str, Any]]: """List all available models in the current region.""" diff --git a/mem0/llms/ollama.py b/mem0/llms/ollama.py index 9c5b0f36f..369f6916c 100644 --- a/mem0/llms/ollama.py +++ b/mem0/llms/ollama.py @@ -8,6 +8,7 @@ except ImportError: from mem0.configs.llms.base import BaseLlmConfig from mem0.configs.llms.ollama import OllamaConfig from mem0.llms.base import LLMBase +from mem0.memory.utils import extract_json class OllamaLLM(LLMBase): @@ -49,20 +50,32 @@ class OllamaLLM(LLMBase): Returns: str or dict: The processed response. """ + # Get the content from response + if isinstance(response, dict): + content = response["message"]["content"] + else: + content = response.message.content + if tools: processed_response = { - "content": response["message"]["content"] if isinstance(response, dict) else response.message.content, + "content": content, "tool_calls": [], } # Ollama doesn't support tool calls in the same way, so we return the content return processed_response else: - # Handle both dict and object responses - if isinstance(response, dict): - return response["message"]["content"] - else: - return response.message.content + # For JSON responses, try to clean up the content using extract_json + if hasattr(self, '_expecting_json') and self._expecting_json: + try: + # Try to extract clean JSON from the response + cleaned_content = extract_json(content) + return cleaned_content + except: + # If extraction fails, return original content + pass + + return content def generate_response( self, @@ -92,7 +105,9 @@ class OllamaLLM(LLMBase): } # Handle JSON response format by modifying the system prompt + self._expecting_json = False if response_format and response_format.get("type") == "json_object": + self._expecting_json = True # Add JSON format instruction to the last message or create a system message if messages and messages[-1]["role"] == "user": messages[-1]["content"] += "\n\nPlease respond with valid JSON only." diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 8a352d53f..a4c3d9dbd 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 capture_event from mem0.memory.utils import ( + extract_json, get_fact_retrieval_messages, parse_messages, parse_vision_messages, @@ -371,7 +372,16 @@ class Memory(MemoryBase): try: response = remove_code_blocks(response) - new_retrieved_facts = json.loads(response)["facts"] + if not response.strip(): + new_retrieved_facts = [] + else: + try: + # First try direct JSON parsing + new_retrieved_facts = json.loads(response)["facts"] + except json.JSONDecodeError: + # Try extracting JSON from response using built-in function + extracted_json = extract_json(response) + new_retrieved_facts = json.loads(extracted_json)["facts"] except Exception as e: logger.error(f"Error in new_retrieved_facts: {e}") new_retrieved_facts = [] @@ -1368,7 +1378,16 @@ class AsyncMemory(MemoryBase): ) try: response = remove_code_blocks(response) - new_retrieved_facts = json.loads(response)["facts"] + if not response.strip(): + new_retrieved_facts = [] + else: + try: + # First try direct JSON parsing + new_retrieved_facts = json.loads(response)["facts"] + except json.JSONDecodeError: + # Try extracting JSON from response using built-in function + extracted_json = extract_json(response) + new_retrieved_facts = json.loads(extracted_json)["facts"] except Exception as e: logger.error(f"Error in new_retrieved_facts: {e}") new_retrieved_facts = []