Fixed json parsing acros differnet LLM providers

This commit is contained in:
parshvadaftari
2025-09-16 04:14:38 +05:30
parent a93a7ea6cf
commit c1ee71ad3f
3 changed files with 60 additions and 20 deletions
+18 -12
View File
@@ -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."""
+21 -6
View File
@@ -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."
+21 -2
View File
@@ -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 = []