feat: add MiniMax LLM provider (#4132)
This commit is contained in:
@@ -273,7 +273,7 @@ config = MemoryConfig(
|
||||
|
||||
### Supported Providers
|
||||
|
||||
#### LLM Providers (19 supported)
|
||||
#### LLM Providers (20 supported)
|
||||
- **openai** - OpenAI GPT models (default)
|
||||
- **anthropic** - Claude models
|
||||
- **gemini** - Google Gemini
|
||||
@@ -284,6 +284,7 @@ config = MemoryConfig(
|
||||
- **azure_openai** - Azure OpenAI
|
||||
- **litellm** - LiteLLM proxy
|
||||
- **deepseek** - DeepSeek models
|
||||
- **minimax** - MiniMax models
|
||||
- **xai** - xAI models
|
||||
- **sarvam** - Sarvam AI
|
||||
- **lmstudio** - LM Studio local server
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
from typing import Optional
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
|
||||
|
||||
class MinimaxConfig(BaseLlmConfig):
|
||||
"""
|
||||
Configuration class for MiniMax-specific parameters.
|
||||
Inherits from BaseLlmConfig and adds MiniMax-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,
|
||||
# MiniMax-specific parameters
|
||||
minimax_base_url: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Initialize MiniMax configuration.
|
||||
|
||||
Args:
|
||||
model: MiniMax model to use, defaults to None
|
||||
temperature: Controls randomness, defaults to 0.1
|
||||
api_key: MiniMax 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
|
||||
minimax_base_url: MiniMax API base URL, defaults to None
|
||||
"""
|
||||
# Initialize base parameters
|
||||
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,
|
||||
)
|
||||
|
||||
# MiniMax-specific parameters
|
||||
self.minimax_base_url = minimax_base_url
|
||||
@@ -23,6 +23,7 @@ class LlmConfig(BaseModel):
|
||||
"azure_openai_structured",
|
||||
"gemini",
|
||||
"deepseek",
|
||||
"minimax",
|
||||
"xai",
|
||||
"sarvam",
|
||||
"lmstudio",
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
import json
|
||||
import os
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
from openai import OpenAI
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
from mem0.configs.llms.minimax import MinimaxConfig
|
||||
from mem0.llms.base import LLMBase
|
||||
from mem0.memory.utils import extract_json
|
||||
|
||||
|
||||
class MiniMaxLLM(LLMBase):
|
||||
def __init__(self, config: Optional[Union[BaseLlmConfig, MinimaxConfig, Dict]] = None):
|
||||
# Convert to MinimaxConfig if needed
|
||||
if config is None:
|
||||
config = MinimaxConfig()
|
||||
elif isinstance(config, dict):
|
||||
config = MinimaxConfig(**config)
|
||||
elif isinstance(config, BaseLlmConfig) and not isinstance(config, MinimaxConfig):
|
||||
# Convert BaseLlmConfig to MinimaxConfig
|
||||
config = MinimaxConfig(
|
||||
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:
|
||||
self.config.model = "MiniMax-M2.1"
|
||||
|
||||
api_key = self.config.api_key or os.getenv("MINIMAX_API_KEY")
|
||||
base_url = (
|
||||
self.config.minimax_base_url
|
||||
or os.getenv("MINIMAX_API_BASE")
|
||||
or "https://api.minimaxi.io/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 MiniMax.
|
||||
|
||||
Args:
|
||||
messages (list): List of message dicts containing 'role' and 'content'.
|
||||
response_format (str or object, optional): Format of the response. Defaults to "text".
|
||||
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 MiniMax-specific parameters.
|
||||
|
||||
Returns:
|
||||
str: The generated response.
|
||||
"""
|
||||
params = self._get_supported_params(messages=messages, **kwargs)
|
||||
params.update(
|
||||
{
|
||||
"model": self.config.model,
|
||||
"messages": messages,
|
||||
}
|
||||
)
|
||||
|
||||
if tools:
|
||||
params["tools"] = tools
|
||||
params["tool_choice"] = tool_choice
|
||||
|
||||
response = self.client.chat.completions.create(**params)
|
||||
return self._parse_response(response, tools)
|
||||
@@ -6,6 +6,7 @@ from mem0.configs.llms.anthropic import AnthropicConfig
|
||||
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.ollama import OllamaConfig
|
||||
from mem0.configs.llms.openai import OpenAIConfig
|
||||
@@ -45,6 +46,7 @@ class LlmFactory:
|
||||
"azure_openai_structured": ("mem0.llms.azure_openai_structured.AzureOpenAIStructuredLLM", AzureOpenAIConfig),
|
||||
"gemini": ("mem0.llms.gemini.GeminiLLM", BaseLlmConfig),
|
||||
"deepseek": ("mem0.llms.deepseek.DeepSeekLLM", DeepSeekConfig),
|
||||
"minimax": ("mem0.llms.minimax.MiniMaxLLM", MinimaxConfig),
|
||||
"xai": ("mem0.llms.xai.XAILLM", BaseLlmConfig),
|
||||
"sarvam": ("mem0.llms.sarvam.SarvamLLM", BaseLlmConfig),
|
||||
"lmstudio": ("mem0.llms.lmstudio.LMStudioLLM", LMStudioConfig),
|
||||
|
||||
@@ -0,0 +1,169 @@
|
||||
import os
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
from mem0.configs.llms.minimax import MinimaxConfig
|
||||
from mem0.llms.minimax import MiniMaxLLM
|
||||
from mem0.utils.factory import LlmFactory
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_minimax_client():
|
||||
with patch("mem0.llms.minimax.OpenAI") as mock_openai:
|
||||
mock_client = Mock()
|
||||
mock_openai.return_value = mock_client
|
||||
yield mock_client
|
||||
|
||||
|
||||
def test_minimax_llm_default_base_url():
|
||||
"""Default config uses MiniMax official base URL."""
|
||||
config = BaseLlmConfig(
|
||||
model="MiniMax-M2.1", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key"
|
||||
)
|
||||
llm = MiniMaxLLM(config)
|
||||
# OpenAI client may normalize URL with trailing slash
|
||||
assert str(llm.client.base_url).rstrip("/") == "https://api.minimaxi.io/v1"
|
||||
|
||||
|
||||
def test_minimax_llm_env_base_url():
|
||||
"""Config uses MINIMAX_API_BASE env variable when set."""
|
||||
provider_base_url = "https://api.provider.com/v1/"
|
||||
os.environ["MINIMAX_API_BASE"] = provider_base_url
|
||||
try:
|
||||
config = MinimaxConfig(
|
||||
model="MiniMax-M2.1",
|
||||
temperature=0.7,
|
||||
max_tokens=100,
|
||||
top_p=1.0,
|
||||
api_key="api_key",
|
||||
)
|
||||
llm = MiniMaxLLM(config)
|
||||
assert str(llm.client.base_url).rstrip("/") == provider_base_url.rstrip("/")
|
||||
finally:
|
||||
os.environ.pop("MINIMAX_API_BASE", None)
|
||||
|
||||
|
||||
def test_minimax_llm_config_base_url():
|
||||
"""Config uses minimax_base_url when provided."""
|
||||
config_base_url = "https://api.config.com/v1/"
|
||||
config = MinimaxConfig(
|
||||
model="MiniMax-M2.1",
|
||||
temperature=0.7,
|
||||
max_tokens=100,
|
||||
top_p=1.0,
|
||||
api_key="api_key",
|
||||
minimax_base_url=config_base_url,
|
||||
)
|
||||
llm = MiniMaxLLM(config)
|
||||
assert str(llm.client.base_url).rstrip("/") == config_base_url.rstrip("/")
|
||||
|
||||
|
||||
def test_minimax_llm_default_model(mock_minimax_client):
|
||||
"""Default model is MiniMax-M2.1 when not specified."""
|
||||
config = MinimaxConfig(temperature=0.7, max_tokens=100, api_key="api_key")
|
||||
llm = MiniMaxLLM(config)
|
||||
assert llm.config.model == "MiniMax-M2.1"
|
||||
|
||||
|
||||
def test_minimax_llm_env_api_key():
|
||||
"""Uses MINIMAX_API_KEY env when api_key not in config."""
|
||||
os.environ["MINIMAX_API_KEY"] = "env-api-key"
|
||||
try:
|
||||
with patch("mem0.llms.minimax.OpenAI") as mock_openai:
|
||||
mock_client = Mock()
|
||||
mock_openai.return_value = mock_client
|
||||
config = MinimaxConfig(model="MiniMax-M2.1", api_key=None)
|
||||
MiniMaxLLM(config)
|
||||
mock_openai.assert_called_once_with(
|
||||
api_key="env-api-key",
|
||||
base_url="https://api.minimaxi.io/v1",
|
||||
)
|
||||
finally:
|
||||
os.environ.pop("MINIMAX_API_KEY", None)
|
||||
|
||||
|
||||
def test_generate_response_without_tools(mock_minimax_client):
|
||||
"""generate_response returns text when no tools provided."""
|
||||
config = BaseLlmConfig(
|
||||
model="MiniMax-M2.1", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key"
|
||||
)
|
||||
llm = MiniMaxLLM(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_minimax_client.chat.completions.create.return_value = mock_response
|
||||
|
||||
response = llm.generate_response(messages)
|
||||
|
||||
mock_minimax_client.chat.completions.create.assert_called_once_with(
|
||||
model="MiniMax-M2.1", 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_minimax_client):
|
||||
"""generate_response returns tool_calls when tools provided."""
|
||||
config = BaseLlmConfig(
|
||||
model="MiniMax-M2.1", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key"
|
||||
)
|
||||
llm = MiniMaxLLM(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_minimax_client.chat.completions.create.return_value = mock_response
|
||||
|
||||
response = llm.generate_response(messages, tools=tools)
|
||||
|
||||
mock_minimax_client.chat.completions.create.assert_called_once_with(
|
||||
model="MiniMax-M2.1",
|
||||
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 len(response["tool_calls"]) == 1
|
||||
assert response["tool_calls"][0]["name"] == "add_memory"
|
||||
assert response["tool_calls"][0]["arguments"] == {"data": "Today is a sunny day."}
|
||||
|
||||
|
||||
def test_factory_creates_minimax_llm(mock_minimax_client):
|
||||
"""LlmFactory.create returns MiniMaxLLM for provider 'minimax'."""
|
||||
llm = LlmFactory.create("minimax", {"model": "MiniMax-M2.1", "api_key": "test-key"})
|
||||
assert isinstance(llm, MiniMaxLLM)
|
||||
assert llm.config.model == "MiniMax-M2.1"
|
||||
Reference in New Issue
Block a user