diff --git a/mem0/llms/ollama.py b/mem0/llms/ollama.py index 3a5fabb2e..74d59c0d1 100644 --- a/mem0/llms/ollama.py +++ b/mem0/llms/ollama.py @@ -1,3 +1,4 @@ +import json from typing import Dict, List, Optional, Union try: @@ -8,6 +9,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): @@ -61,7 +63,28 @@ class OllamaLLM(LLMBase): "tool_calls": [], } - # Ollama doesn't support tool calls in the same way, so we return the content + if isinstance(response, dict): + raw_calls = response.get("message", {}).get("tool_calls") or [] + else: + raw_calls = getattr(response.message, "tool_calls", None) or [] + + for tool_call in raw_calls: + if isinstance(tool_call, dict): + fn = tool_call.get("function", {}) + name = fn.get("name", "") + arguments = fn.get("arguments", {}) + else: + fn = getattr(tool_call, "function", None) + name = getattr(fn, "name", "") if fn else "" + arguments = getattr(fn, "arguments", {}) if fn else {} + + if isinstance(arguments, str): + arguments = json.loads(extract_json(arguments)) + + processed_response["tool_calls"].append( + {"name": name, "arguments": arguments} + ) + return processed_response else: return content @@ -113,5 +136,8 @@ class OllamaLLM(LLMBase): # Remove OpenAI-specific parameters that Ollama doesn't support params.pop("max_tokens", None) # Ollama uses different parameter names + if tools: + params["tools"] = tools + response = self.client.chat(**params) return self._parse_response(response, tools) diff --git a/tests/llms/test_ollama.py b/tests/llms/test_ollama.py index 0f1e6ac3e..cfbb8bfdf 100644 --- a/tests/llms/test_ollama.py +++ b/tests/llms/test_ollama.py @@ -32,3 +32,111 @@ def test_generate_response_without_tools(mock_ollama_client): model="llama3.1:70b", messages=messages, options={"temperature": 0.7, "num_predict": 100, "top_p": 1.0} ) assert response == "I'm doing well, thank you for asking!" + + +def test_generate_response_with_tools_passes_tools_to_client(mock_ollama_client): + """Tools should be forwarded to ollama client.chat().""" + config = OllamaConfig(model="llama3.1:70b", temperature=0.1, max_tokens=100, top_p=1.0) + llm = OllamaLLM(config) + messages = [{"role": "user", "content": "Extract entities from: Alice works at UCSD"}] + tools = [ + { + "type": "function", + "function": { + "name": "extract_entities", + "description": "Extract entities", + "parameters": {"type": "object", "properties": {"entities": {"type": "array"}}}, + }, + } + ] + + mock_response = { + "message": { + "content": "", + "tool_calls": [ + { + "function": { + "name": "extract_entities", + "arguments": {"entities": [{"name": "Alice"}, {"name": "UCSD"}]}, + } + } + ], + } + } + mock_ollama_client.chat.return_value = mock_response + + response = llm.generate_response(messages, tools=tools) + + # Verify tools were passed to client.chat + call_kwargs = mock_ollama_client.chat.call_args + assert "tools" in call_kwargs.kwargs or (len(call_kwargs.args) > 0 and "tools" in call_kwargs[1]) + assert call_kwargs[1]["tools"] == tools + + # Verify tool_calls were parsed correctly + assert response["tool_calls"] == [ + {"name": "extract_entities", "arguments": {"entities": [{"name": "Alice"}, {"name": "UCSD"}]}} + ] + + +def test_generate_response_with_tools_no_tool_calls_in_response(mock_ollama_client): + """When model returns content without tool_calls, tool_calls should be empty list.""" + config = OllamaConfig(model="llama3.1:70b", temperature=0.1, max_tokens=100, top_p=1.0) + llm = OllamaLLM(config) + messages = [{"role": "user", "content": "Hello"}] + tools = [{"type": "function", "function": {"name": "noop", "parameters": {}}}] + + mock_response = {"message": {"content": "I cannot use tools for this.", "tool_calls": []}} + mock_ollama_client.chat.return_value = mock_response + + response = llm.generate_response(messages, tools=tools) + + assert response["content"] == "I cannot use tools for this." + assert response["tool_calls"] == [] + + +def test_generate_response_with_tools_string_arguments(mock_ollama_client): + """When tool_call arguments come as JSON string, they should be parsed.""" + config = OllamaConfig(model="llama3.1:70b", temperature=0.1, max_tokens=100, top_p=1.0) + llm = OllamaLLM(config) + messages = [{"role": "user", "content": "test"}] + tools = [{"type": "function", "function": {"name": "test_fn", "parameters": {}}}] + + mock_response = { + "message": { + "content": "", + "tool_calls": [ + {"function": {"name": "test_fn", "arguments": '{"key": "value"}'}} + ], + } + } + mock_ollama_client.chat.return_value = mock_response + + response = llm.generate_response(messages, tools=tools) + + assert response["tool_calls"] == [{"name": "test_fn", "arguments": {"key": "value"}}] + + +def test_parse_response_with_tools_object_style(mock_ollama_client): + """Test _parse_response with object-style response (non-dict).""" + config = OllamaConfig(model="llama3.1:70b") + llm = OllamaLLM(config) + + # Simulate object-style response + mock_fn = Mock() + mock_fn.name = "extract" + mock_fn.arguments = {"entities": ["Alice"]} + + mock_tool_call = Mock() + mock_tool_call.function = mock_fn + + mock_message = Mock() + mock_message.content = "" + mock_message.tool_calls = [mock_tool_call] + + mock_response = Mock() + mock_response.message = mock_message + + tools = [{"type": "function", "function": {"name": "extract"}}] + result = llm._parse_response(mock_response, tools) + + assert result["tool_calls"] == [{"name": "extract", "arguments": {"entities": ["Alice"]}}]