fix(llms/xai): forward tools, add XAIConfig, parse tool_calls (#5190)

This commit is contained in:
Aarkin Karnik
2026-06-10 12:46:27 +05:30
committed by GitHub
parent a86c87236d
commit 652193d599
4 changed files with 344 additions and 19 deletions
+55
View File
@@ -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
+71 -13
View File
@@ -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)
+9 -6
View File
@@ -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),
+209
View File
@@ -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)