From 652193d59908da52bef412ec3ac3cb0242f42f79 Mon Sep 17 00:00:00 2001 From: Aarkin Karnik <56022539+Aarkin7@users.noreply.github.com> Date: Wed, 10 Jun 2026 12:46:27 +0530 Subject: [PATCH] fix(llms/xai): forward tools, add XAIConfig, parse tool_calls (#5190) --- mem0/configs/llms/xai.py | 55 +++++++++++ mem0/llms/xai.py | 84 +++++++++++++--- mem0/utils/factory.py | 15 +-- tests/llms/test_xai.py | 209 +++++++++++++++++++++++++++++++++++++++ 4 files changed, 344 insertions(+), 19 deletions(-) create mode 100644 mem0/configs/llms/xai.py create mode 100644 tests/llms/test_xai.py diff --git a/mem0/configs/llms/xai.py b/mem0/configs/llms/xai.py new file mode 100644 index 000000000..7a262df7c --- /dev/null +++ b/mem0/configs/llms/xai.py @@ -0,0 +1,55 @@ +from typing import Optional + +from mem0.configs.llms.base import BaseLlmConfig + + +class XAIConfig(BaseLlmConfig): + """ + Configuration class for X.AI (Grok) provider parameters. + Inherits from BaseLlmConfig and adds X.AI-specific settings. + """ + + def __init__( + self, + # Base parameters + model: Optional[str] = None, + temperature: float = 0.1, + api_key: Optional[str] = None, + max_tokens: int = 2000, + top_p: float = 0.1, + top_k: int = 1, + enable_vision: bool = False, + vision_details: Optional[str] = "auto", + http_client_proxies: Optional[dict] = None, + # X.AI-specific parameters + xai_base_url: Optional[str] = None, + ): + """ + Initialize X.AI configuration. + + Args: + model: X.AI / Grok model to use, defaults to None + temperature: Controls randomness, defaults to 0.1 + api_key: X.AI API key, defaults to None + max_tokens: Maximum tokens to generate, defaults to 2000 + top_p: Nucleus sampling parameter, defaults to 0.1 + top_k: Top-k sampling parameter, defaults to 1 + enable_vision: Enable vision capabilities, defaults to False + vision_details: Vision detail level, defaults to "auto" + http_client_proxies: HTTP client proxy settings, defaults to None + xai_base_url: X.AI API base URL, defaults to None + """ + super().__init__( + model=model, + temperature=temperature, + api_key=api_key, + max_tokens=max_tokens, + top_p=top_p, + top_k=top_k, + enable_vision=enable_vision, + vision_details=vision_details, + http_client_proxies=http_client_proxies, + ) + + # X.AI-specific parameters + self.xai_base_url = xai_base_url diff --git a/mem0/llms/xai.py b/mem0/llms/xai.py index a918ac41c..8c38e1d07 100644 --- a/mem0/llms/xai.py +++ b/mem0/llms/xai.py @@ -1,14 +1,36 @@ +import json import os -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Union from openai import OpenAI from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.xai import XAIConfig from mem0.llms.base import LLMBase +from mem0.memory.utils import extract_json class XAILLM(LLMBase): - def __init__(self, config: Optional[BaseLlmConfig] = None): + def __init__(self, config: Optional[Union[BaseLlmConfig, XAIConfig, Dict]] = None): + # Convert to XAIConfig if needed + if config is None: + config = XAIConfig() + elif isinstance(config, dict): + config = XAIConfig(**config) + elif isinstance(config, BaseLlmConfig) and not isinstance(config, XAIConfig): + # Convert BaseLlmConfig to XAIConfig so xai_base_url is available + config = XAIConfig( + model=config.model, + temperature=config.temperature, + api_key=config.api_key, + max_tokens=config.max_tokens, + top_p=config.top_p, + top_k=config.top_k, + enable_vision=config.enable_vision, + vision_details=config.vision_details, + http_client_proxies=config.http_client, + ) + super().__init__(config) if not self.config.model: @@ -18,35 +40,71 @@ class XAILLM(LLMBase): base_url = self.config.xai_base_url or os.getenv("XAI_API_BASE") or "https://api.x.ai/v1" self.client = OpenAI(api_key=api_key, base_url=base_url) + def _parse_response(self, response, tools): + """ + Process the response based on whether tools are used or not. + + Args: + response: The raw response from API. + tools: The list of tools provided in the request. + + Returns: + str or dict: The processed response. + """ + if tools: + processed_response = { + "content": response.choices[0].message.content, + "tool_calls": [], + } + + if response.choices[0].message.tool_calls: + for tool_call in response.choices[0].message.tool_calls: + processed_response["tool_calls"].append( + { + "name": tool_call.function.name, + "arguments": json.loads(extract_json(tool_call.function.arguments)), + } + ) + + return processed_response + else: + return response.choices[0].message.content + def generate_response( self, messages: List[Dict[str, str]], response_format=None, tools: Optional[List[Dict]] = None, tool_choice: str = "auto", + **kwargs, ): """ - Generate a response based on the given messages using XAI. + Generate a response based on the given messages using X.AI (Grok). Args: messages (list): List of message dicts containing 'role' and 'content'. - response_format (str or object, optional): Format of the response. Defaults to "text". + response_format (str or object, optional): Format of the response. Defaults to None. tools (list, optional): List of tools that the model can call. Defaults to None. tool_choice (str, optional): Tool choice method. Defaults to "auto". + **kwargs: Additional X.AI-specific parameters. Returns: - str: The generated response. + str or dict: The generated response. A string when tools are not requested; + a dict ``{"content": ..., "tool_calls": [...]}`` when tools are requested. """ - params = { - "model": self.config.model, - "messages": messages, - "temperature": self.config.temperature, - "max_tokens": self.config.max_tokens, - "top_p": self.config.top_p, - } + params = self._get_supported_params(messages=messages, **kwargs) + params.update( + { + "model": self.config.model, + "messages": messages, + } + ) if response_format: params["response_format"] = response_format + if tools: + params["tools"] = tools + params["tool_choice"] = tool_choice response = self.client.chat.completions.create(**params) - return response.choices[0].message.content + return self._parse_response(response, tools) diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py index 754ea0988..487c10850 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -7,17 +7,18 @@ from mem0.configs.llms.aws_bedrock import AWSBedrockConfig from mem0.configs.llms.azure import AzureOpenAIConfig from mem0.configs.llms.base import BaseLlmConfig from mem0.configs.llms.deepseek import DeepSeekConfig -from mem0.configs.llms.minimax import MinimaxConfig from mem0.configs.llms.lmstudio import LMStudioConfig +from mem0.configs.llms.minimax import MinimaxConfig from mem0.configs.llms.ollama import OllamaConfig from mem0.configs.llms.openai import OpenAIConfig from mem0.configs.llms.vllm import VllmConfig +from mem0.configs.llms.xai import XAIConfig from mem0.configs.rerankers.base import BaseRerankerConfig from mem0.configs.rerankers.cohere import CohereRerankerConfig +from mem0.configs.rerankers.huggingface import HuggingFaceRerankerConfig +from mem0.configs.rerankers.llm import LLMRerankerConfig from mem0.configs.rerankers.sentence_transformer import SentenceTransformerRerankerConfig from mem0.configs.rerankers.zero_entropy import ZeroEntropyRerankerConfig -from mem0.configs.rerankers.llm import LLMRerankerConfig -from mem0.configs.rerankers.huggingface import HuggingFaceRerankerConfig from mem0.embeddings.mock import MockEmbeddings @@ -48,7 +49,7 @@ class LlmFactory: "gemini": ("mem0.llms.gemini.GeminiLLM", BaseLlmConfig), "deepseek": ("mem0.llms.deepseek.DeepSeekLLM", DeepSeekConfig), "minimax": ("mem0.llms.minimax.MiniMaxLLM", MinimaxConfig), - "xai": ("mem0.llms.xai.XAILLM", BaseLlmConfig), + "xai": ("mem0.llms.xai.XAILLM", XAIConfig), "sarvam": ("mem0.llms.sarvam.SarvamLLM", BaseLlmConfig), "lmstudio": ("mem0.llms.lmstudio.LMStudioLLM", LMStudioConfig), "vllm": ("mem0.llms.vllm.VllmLLM", VllmConfig), @@ -209,7 +210,6 @@ class VectorStoreFactory: return instance - class RerankerFactory: """ Factory for creating reranker instances with appropriate configurations. @@ -219,7 +219,10 @@ class RerankerFactory: # Provider mappings with their config classes provider_to_class = { "cohere": ("mem0.reranker.cohere_reranker.CohereReranker", CohereRerankerConfig), - "sentence_transformer": ("mem0.reranker.sentence_transformer_reranker.SentenceTransformerReranker", SentenceTransformerRerankerConfig), + "sentence_transformer": ( + "mem0.reranker.sentence_transformer_reranker.SentenceTransformerReranker", + SentenceTransformerRerankerConfig, + ), "zero_entropy": ("mem0.reranker.zero_entropy_reranker.ZeroEntropyReranker", ZeroEntropyRerankerConfig), "llm_reranker": ("mem0.reranker.llm_reranker.LLMReranker", LLMRerankerConfig), "huggingface": ("mem0.reranker.huggingface_reranker.HuggingFaceReranker", HuggingFaceRerankerConfig), diff --git a/tests/llms/test_xai.py b/tests/llms/test_xai.py new file mode 100644 index 000000000..5143064ec --- /dev/null +++ b/tests/llms/test_xai.py @@ -0,0 +1,209 @@ +import os +from unittest.mock import Mock, patch + +import pytest + +from mem0.configs.llms.base import BaseLlmConfig +from mem0.configs.llms.xai import XAIConfig +from mem0.llms.xai import XAILLM +from mem0.utils.factory import LlmFactory + + +@pytest.fixture +def mock_xai_client(): + with patch("mem0.llms.xai.OpenAI") as mock_openai: + mock_client = Mock() + mock_openai.return_value = mock_client + yield mock_client + + +def test_xai_llm_base_url(): + # case1: default + config = XAIConfig(model="grok-2-latest", api_key="api_key") + llm = XAILLM(config) + assert str(llm.client.base_url) == "https://api.x.ai/v1/" + + # case2: XAI_API_BASE env var + os.environ["XAI_API_BASE"] = "https://api.provider.com/v1" + config = XAIConfig(model="grok-2-latest", api_key="api_key") + llm = XAILLM(config) + assert str(llm.client.base_url) == "https://api.provider.com/v1/" + + # case3: config.xai_base_url wins over env + config = XAIConfig( + model="grok-2-latest", + api_key="api_key", + xai_base_url="https://api.config.com/v1", + ) + llm = XAILLM(config) + assert str(llm.client.base_url) == "https://api.config.com/v1/" + + +def test_xai_accepts_base_llm_config(): + # Used to AttributeError on self.config.xai_base_url because the factory + # wired XAI with plain BaseLlmConfig. + llm = XAILLM(BaseLlmConfig(model="grok-2-latest", api_key="k")) + assert isinstance(llm.config, XAIConfig) + assert llm.config.xai_base_url is None + + +def test_generate_response_without_tools(mock_xai_client): + config = XAIConfig(model="grok-2-latest", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key") + llm = XAILLM(config) + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Hello, how are you?"}, + ] + + mock_response = Mock() + mock_response.choices = [Mock(message=Mock(content="I'm doing well, thank you for asking!"))] + mock_xai_client.chat.completions.create.return_value = mock_response + + response = llm.generate_response(messages) + + mock_xai_client.chat.completions.create.assert_called_once_with( + model="grok-2-latest", + messages=messages, + temperature=0.7, + max_tokens=100, + top_p=1.0, + ) + assert response == "I'm doing well, thank you for asking!" + + +def test_generate_response_with_tools(mock_xai_client): + config = XAIConfig(model="grok-2-latest", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key") + llm = XAILLM(config) + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Add a new memory: Today is a sunny day."}, + ] + tools = [ + { + "type": "function", + "function": { + "name": "add_memory", + "description": "Add a memory", + "parameters": { + "type": "object", + "properties": {"data": {"type": "string", "description": "Data to add to memory"}}, + "required": ["data"], + }, + }, + } + ] + + mock_response = Mock() + mock_message = Mock() + mock_message.content = "I've added the memory for you." + mock_tool_call = Mock() + mock_tool_call.function.name = "add_memory" + mock_tool_call.function.arguments = '{"data": "Today is a sunny day."}' + mock_message.tool_calls = [mock_tool_call] + mock_response.choices = [Mock(message=mock_message)] + mock_xai_client.chat.completions.create.return_value = mock_response + + response = llm.generate_response(messages, tools=tools) + + mock_xai_client.chat.completions.create.assert_called_once_with( + model="grok-2-latest", + messages=messages, + temperature=0.7, + max_tokens=100, + top_p=1.0, + tools=tools, + tool_choice="auto", + ) + + assert response["content"] == "I've added the memory for you." + assert response["tool_calls"] == [{"name": "add_memory", "arguments": {"data": "Today is a sunny day."}}] + + +def test_empty_tools_list_not_forwarded(mock_xai_client): + # tools=[] would otherwise get rejected by some OpenAI-compatible backends + config = XAIConfig(model="grok-2-latest", api_key="api_key") + llm = XAILLM(config) + + mock_response = Mock() + mock_response.choices = [Mock(message=Mock(content="hello"))] + mock_xai_client.chat.completions.create.return_value = mock_response + + response = llm.generate_response([{"role": "user", "content": "hi"}], tools=[]) + + call_kwargs = mock_xai_client.chat.completions.create.call_args.kwargs + assert "tools" not in call_kwargs + assert "tool_choice" not in call_kwargs + assert response == "hello" + + +def test_tools_requested_but_model_returns_no_calls(mock_xai_client): + # Model can decline to call any tool even when offered + config = XAIConfig(model="grok-2-latest", api_key="api_key") + llm = XAILLM(config) + tools = [{"type": "function", "function": {"name": "f", "parameters": {"type": "object", "properties": {}}}}] + + mock_message = Mock(content="just a chat reply") + mock_message.tool_calls = None + mock_xai_client.chat.completions.create.return_value = Mock(choices=[Mock(message=mock_message)]) + + response = llm.generate_response([{"role": "user", "content": "x"}], tools=tools) + assert response == {"content": "just a chat reply", "tool_calls": []} + + +def test_generate_response_with_response_format(mock_xai_client): + config = XAIConfig(model="grok-2-latest", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key") + llm = XAILLM(config) + messages = [ + {"role": "system", "content": "You are a memory extraction assistant."}, + {"role": "user", "content": "I like hiking on weekends."}, + ] + + mock_response = Mock() + mock_response.choices = [Mock(message=Mock(content='{"facts": ["User likes hiking on weekends"]}'))] + mock_xai_client.chat.completions.create.return_value = mock_response + + response = llm.generate_response(messages, response_format={"type": "json_object"}) + + mock_xai_client.chat.completions.create.assert_called_once_with( + model="grok-2-latest", + messages=messages, + temperature=0.7, + max_tokens=100, + top_p=1.0, + response_format={"type": "json_object"}, + ) + assert response == '{"facts": ["User likes hiking on weekends"]}' + + +def test_generate_response_without_response_format(mock_xai_client): + config = XAIConfig(model="grok-2-latest", api_key="api_key") + llm = XAILLM(config) + + mock_response = Mock() + mock_response.choices = [Mock(message=Mock(content="hi"))] + mock_xai_client.chat.completions.create.return_value = mock_response + + llm.generate_response([{"role": "user", "content": "hi"}]) + + call_kwargs = mock_xai_client.chat.completions.create.call_args.kwargs + assert "response_format" not in call_kwargs + + +def test_factory_creates_xai_from_dict(): + with patch("mem0.llms.xai.OpenAI") as mock_openai: + mock_openai.return_value = Mock() + llm = LlmFactory.create( + "xai", + {"model": "grok-2-latest", "api_key": "k", "xai_base_url": "https://example.com/v1"}, + ) + assert isinstance(llm, XAILLM) + assert isinstance(llm.config, XAIConfig) + assert llm.config.xai_base_url == "https://example.com/v1" + + +def test_factory_creates_xai_from_base_config(): + # Legacy callers still hand the factory a plain BaseLlmConfig + with patch("mem0.llms.xai.OpenAI") as mock_openai: + mock_openai.return_value = Mock() + llm = LlmFactory.create("xai", BaseLlmConfig(model="grok-2-latest", api_key="k")) + assert isinstance(llm.config, XAIConfig)