Fixed json parsing acros differnet LLM providers
This commit is contained in:
+18
-12
@@ -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
@@ -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
@@ -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 = []
|
||||
|
||||
Reference in New Issue
Block a user