fix(llms): accept and forward **kwargs in Together/LangChain/Sarvam providers (#5556)

This commit is contained in:
Yash Singh
2026-06-15 12:23:04 +05:30
committed by GitHub
parent a1eefc31bc
commit 66901d7393
6 changed files with 92 additions and 2 deletions
+14
View File
@@ -115,6 +115,20 @@ def test_generate_response_with_tools(mock_langchain_model):
assert response["tool_calls"][0]["arguments"] == {"data": "Today is a sunny day."}
def test_generate_response_forwards_extra_kwargs(mock_langchain_model):
"""Per the LLMBase contract, extra model kwargs must be accepted and forwarded to
the underlying LangChain model's ``invoke``."""
config = BaseLlmConfig(model=mock_langchain_model, temperature=0.7, max_tokens=100, api_key="test-api-key")
llm = LangchainLLM(config)
messages = [{"role": "user", "content": "Hello"}]
response = llm.generate_response(messages, frequency_penalty=0.5)
mock_langchain_model.invoke.assert_called_once()
assert mock_langchain_model.invoke.call_args.kwargs["frequency_penalty"] == 0.5
assert response == "This is a test response"
def test_invalid_model():
"""Test that LangchainLLM raises an error with an invalid model."""
config = BaseLlmConfig(model="not-a-valid-model-instance", temperature=0.7, max_tokens=100, api_key="test-api-key")
+46
View File
@@ -0,0 +1,46 @@
from unittest.mock import Mock, patch
import pytest
from mem0.configs.llms.base import BaseLlmConfig
from mem0.llms.sarvam import SarvamLLM
@pytest.fixture
def sarvam_llm():
config = BaseLlmConfig(model="sarvam-m", temperature=0.7, max_tokens=100, top_p=1.0, api_key="test-api-key")
return SarvamLLM(config)
def _mock_post(content="Hello there!"):
mock_response = Mock()
mock_response.raise_for_status.return_value = None
mock_response.json.return_value = {"choices": [{"message": {"content": content}}]}
return mock_response
def test_generate_response_returns_content(sarvam_llm):
with patch("mem0.llms.sarvam.requests.post", return_value=_mock_post("Hi!")) as mock_post:
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello, how are you?"},
]
response = sarvam_llm.generate_response(messages)
assert response == "Hi!"
sent_payload = mock_post.call_args.kwargs["json"]
assert sent_payload["model"] == "sarvam-m"
assert sent_payload["messages"] == messages
assert sent_payload["temperature"] == 0.7
def test_generate_response_forwards_extra_kwargs(sarvam_llm):
"""Per the LLMBase contract, extra provider-specific kwargs must be accepted and
forwarded into the Sarvam request payload."""
with patch("mem0.llms.sarvam.requests.post", return_value=_mock_post("Hi!")) as mock_post:
messages = [{"role": "user", "content": "Hello"}]
response = sarvam_llm.generate_response(messages, frequency_penalty=0.5)
assert response == "Hi!"
sent_payload = mock_post.call_args.kwargs["json"]
assert sent_payload["frequency_penalty"] == 0.5
+18
View File
@@ -84,3 +84,21 @@ def test_generate_response_with_tools(mock_together_client):
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_generate_response_forwards_extra_kwargs(mock_together_client):
"""Per the LLMBase contract, extra provider-specific kwargs must be accepted and
forwarded to the Together client (matching openai/deepseek/vllm/xai behavior)."""
config = BaseLlmConfig(model="mistralai/Mixtral-8x7B-Instruct-v0.1", temperature=0.7, max_tokens=100, top_p=1.0)
llm = TogetherLLM(config)
messages = [{"role": "user", "content": "Hello"}]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="Hi"))]
mock_together_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages, frequency_penalty=0.5)
assert response == "Hi"
call_kwargs = mock_together_client.chat.completions.create.call_args.kwargs
assert call_kwargs["frequency_penalty"] == 0.5