fix(bedrock): omit topP for Anthropic Converse; use AWSBedrockConfig in LlmFactory (#4469)

This commit is contained in:
Himanshu
2026-03-25 11:22:52 +05:30
committed by GitHub
parent d1b4b304c7
commit 2e0f91e70d
4 changed files with 383 additions and 46 deletions
+8 -4
View File
@@ -16,7 +16,7 @@ class AWSBedrockConfig(BaseLlmConfig):
model: Optional[str] = None,
temperature: float = 0.1,
max_tokens: int = 2000,
top_p: float = 0.9,
top_p: Optional[float] = None,
top_k: int = 1,
aws_access_key_id: Optional[str] = None,
aws_secret_access_key: Optional[str] = None,
@@ -33,7 +33,8 @@ class AWSBedrockConfig(BaseLlmConfig):
model: Bedrock model identifier (e.g., "amazon.nova-3-mini-20241119-v1:0")
temperature: Controls randomness (0.0 to 2.0)
max_tokens: Maximum tokens to generate
top_p: Nucleus sampling parameter (0.0 to 1.0)
top_p: Nucleus sampling (0.0–1.0). Default None omits topP on Converse
(required for Anthropic, which rejects temperature and topP together).
top_k: Top-k sampling parameter (1 to 40)
aws_access_key_id: AWS access key (optional, uses env vars if not provided)
aws_secret_access_key: AWS secret key (optional, uses env vars if not provided)
@@ -75,13 +76,16 @@ class AWSBedrockConfig(BaseLlmConfig):
def get_model_config(self) -> Dict[str, Any]:
"""Get model-specific configuration parameters."""
base_config = {
base_config: Dict[str, Any] = {
"temperature": self.temperature,
"max_tokens": self.max_tokens,
"top_p": self.top_p,
"top_k": self.top_k,
}
# Only include top_p when explicitly set by the user.
if self.top_p is not None:
base_config["top_p"] = self.top_p
# Add custom model kwargs
base_config.update(self.model_kwargs)
+43 -41
View File
@@ -228,6 +228,12 @@ class AWSBedrockLLM(LLMBase):
return "\n\nHuman: " + "".join(formatted_messages) + "\n\nAssistant:"
def _merge_optional_top_p(self, target: Dict[str, Any], *, key: str = "top_p") -> None:
"""Add nucleus sampling to ``target`` only when ``model_config`` has ``top_p`` set."""
top_p = self.model_config.get("top_p")
if top_p is not None:
target[key] = top_p
def _prepare_input(self, prompt: str) -> Dict[str, Any]:
"""
Prepare input for the current provider's model.
@@ -268,44 +274,38 @@ class AWSBedrockLLM(LLMBase):
"messages": [{"role": "user", "content": prompt}],
"max_tokens": self.model_config.get("max_tokens", 5000),
"temperature": self.model_config.get("temperature", 0.1),
"top_p": self.model_config.get("top_p", 0.9),
}
self._merge_optional_top_p(input_body)
else:
# Legacy Amazon models
input_body = {
"inputText": prompt,
"textGenerationConfig": {
"maxTokenCount": self.model_config.get("max_tokens", 5000),
"topP": self.model_config.get("top_p", 0.9),
"temperature": self.model_config.get("temperature", 0.1),
},
}
# Remove None values
input_body["textGenerationConfig"] = {
k: v for k, v in input_body["textGenerationConfig"].items() if v is not None
text_gen_config: Dict[str, Any] = {
"maxTokenCount": self.model_config.get("max_tokens", 5000),
"temperature": self.model_config.get("temperature", 0.1),
}
self._merge_optional_top_p(text_gen_config, key="topP")
input_body = {"inputText": prompt, "textGenerationConfig": text_gen_config}
elif self.provider == "anthropic":
input_body = {
"messages": [{"role": "user", "content": [{"type": "text", "text": prompt}]}],
"max_tokens": self.model_config.get("max_tokens", 2000),
"temperature": self.model_config.get("temperature", 0.1),
"top_p": self.model_config.get("top_p", 0.9),
"anthropic_version": "bedrock-2023-05-31",
}
self._merge_optional_top_p(input_body)
elif self.provider == "meta":
input_body = {
"prompt": prompt,
"max_gen_len": self.model_config.get("max_tokens", 5000),
"temperature": self.model_config.get("temperature", 0.1),
"top_p": self.model_config.get("top_p", 0.9),
}
self._merge_optional_top_p(input_body)
elif self.provider == "mistral":
input_body = {
"prompt": prompt,
"max_tokens": self.model_config.get("max_tokens", 5000),
"temperature": self.model_config.get("temperature", 0.1),
"top_p": self.model_config.get("top_p", 0.9),
}
self._merge_optional_top_p(input_body)
else:
# Generic case - add all model config parameters
input_body.update(self.model_config)
@@ -479,6 +479,29 @@ class AWSBedrockLLM(LLMBase):
return converse_tools
def _default_max_tokens_for_converse(self) -> int:
"""Default maxTokens if ``max_tokens`` is missing (Nova: 5000, else 2000)."""
model_id = (self.config.model or "").lower()
if self.provider == "amazon" and "nova" in model_id:
return 5000
return 2000
def _build_inference_config(self) -> Dict[str, Any]:
"""Build Converse ``inferenceConfig``. Anthropic allows only one of temperature or topP; we keep temperature and omit topP."""
inference_config: Dict[str, Any] = {
"maxTokens": self.model_config.get("max_tokens", self._default_max_tokens_for_converse()),
"temperature": self.model_config.get("temperature", 0.1),
}
top_p = self.model_config.get("top_p")
if top_p is not None:
if self.provider == "anthropic":
logger.debug("Omitting topP for Anthropic Converse (using temperature); top_p=%s", top_p)
else:
inference_config["topP"] = top_p
return inference_config
def _generate_with_tools(self, messages: List[Dict[str, str]], tools: List[Dict], stream: bool = False) -> Dict[str, Any]:
"""Generate response with tool calling support using correct message format."""
# Format messages for tool-enabled models
@@ -501,11 +524,7 @@ class AWSBedrockLLM(LLMBase):
converse_params = {
"modelId": self.config.model,
"messages": formatted_messages,
"inferenceConfig": {
"maxTokens": self.model_config.get("max_tokens", 2000),
"temperature": self.model_config.get("temperature", 0.1),
"topP": self.model_config.get("top_p", 0.9),
}
"inferenceConfig": self._build_inference_config(),
}
# Add system message if present (for Anthropic)
@@ -531,11 +550,7 @@ class AWSBedrockLLM(LLMBase):
converse_params = {
"modelId": self.config.model,
"messages": formatted_messages,
"inferenceConfig": {
"maxTokens": self.model_config.get("max_tokens", 2000),
"temperature": self.model_config.get("temperature", 0.1),
"topP": self.model_config.get("top_p", 0.9),
}
"inferenceConfig": self._build_inference_config(),
}
# Add system message if present
@@ -554,26 +569,13 @@ class AWSBedrockLLM(LLMBase):
return str(response)
elif self.provider == "amazon" and "nova" in self.config.model.lower():
# Nova models use converse API even without tools
# Nova models use the Converse API even without tools
formatted_messages = self._format_messages_amazon(messages)
input_body = {
"messages": formatted_messages,
"max_tokens": self.model_config.get("max_tokens", 5000),
"temperature": self.model_config.get("temperature", 0.1),
"top_p": self.model_config.get("top_p", 0.9),
}
# Use converse API for Nova models
response = self.client.converse(
modelId=self.config.model,
messages=input_body["messages"],
inferenceConfig={
"maxTokens": input_body["max_tokens"],
"temperature": input_body["temperature"],
"topP": input_body["top_p"],
}
messages=formatted_messages,
inferenceConfig=self._build_inference_config(),
)
return self._parse_response(response)
else:
# For other providers and legacy Amazon models (like Titan)
+2 -1
View File
@@ -3,6 +3,7 @@ from typing import Dict, Optional, Union
from mem0.configs.embeddings.base import BaseEmbedderConfig
from mem0.configs.llms.anthropic import AnthropicConfig
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
@@ -38,7 +39,7 @@ class LlmFactory:
"openai": ("mem0.llms.openai.OpenAILLM", OpenAIConfig),
"groq": ("mem0.llms.groq.GroqLLM", BaseLlmConfig),
"together": ("mem0.llms.together.TogetherLLM", BaseLlmConfig),
"aws_bedrock": ("mem0.llms.aws_bedrock.AWSBedrockLLM", BaseLlmConfig),
"aws_bedrock": ("mem0.llms.aws_bedrock.AWSBedrockLLM", AWSBedrockConfig),
"litellm": ("mem0.llms.litellm.LiteLLM", BaseLlmConfig),
"azure_openai": ("mem0.llms.azure_openai.AzureOpenAILLM", AzureOpenAIConfig),
"openai_structured": ("mem0.llms.openai_structured.OpenAIStructuredLLM", OpenAIConfig),
+330
View File
@@ -0,0 +1,330 @@
from unittest.mock import MagicMock, patch
import pytest
from mem0.configs.llms.aws_bedrock import AWSBedrockConfig
from mem0.llms.aws_bedrock import AWSBedrockLLM, extract_provider
from mem0.utils.factory import LlmFactory
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@pytest.fixture
def mock_boto3():
"""Patch boto3 so no real AWS calls are made during unit tests."""
with patch("mem0.llms.aws_bedrock.boto3") as mock_b3:
runtime_client = MagicMock()
bedrock_client = MagicMock()
bedrock_client.list_foundation_models.return_value = {"modelSummaries": []}
def _client(service, **kwargs):
if service == "bedrock-runtime":
return runtime_client
return bedrock_client
mock_b3.client.side_effect = _client
yield runtime_client
def _make_llm(model: str, mock_boto3, **kwargs) -> AWSBedrockLLM:
"""Instantiate AWSBedrockLLM with a given model, all AWS calls mocked."""
config = AWSBedrockConfig(model=model, **kwargs)
return AWSBedrockLLM(config)
def _converse_response(text: str = "ok") -> dict:
"""Minimal Converse API response dict."""
return {"output": {"message": {"content": [{"text": text}]}}}
# ---------------------------------------------------------------------------
# extract_provider
# ---------------------------------------------------------------------------
class TestExtractProvider:
def test_standard_anthropic_model(self):
assert extract_provider("anthropic.claude-3-5-sonnet-20240620-v1:0") == "anthropic"
def test_inference_profile_us_prefix(self):
# Cross-region inference profile IDs look like us.anthropic.<model>
assert extract_provider("us.anthropic.claude-haiku-4-5-20251001-v1:0") == "anthropic"
def test_inference_profile_eu_prefix(self):
assert extract_provider("eu.anthropic.claude-sonnet-4-5-20250929-v1:0") == "anthropic"
def test_inference_profile_ap_prefix(self):
assert extract_provider("ap.anthropic.claude-3-opus-20240229-v1:0") == "anthropic"
def test_amazon_model(self):
assert extract_provider("amazon.nova-3-mini-20241119-v1:0") == "amazon"
def test_meta_model(self):
assert extract_provider("meta.llama3-8b-instruct-v1:0") == "meta"
def test_mistral_model(self):
assert extract_provider("mistral.mistral-7b-instruct-v0:2") == "mistral"
def test_unknown_model_raises(self):
with pytest.raises(ValueError, match="Unknown provider"):
extract_provider("unknown-vendor.some-model-v1:0")
# ---------------------------------------------------------------------------
# AWSBedrockConfig
# ---------------------------------------------------------------------------
class TestAWSBedrockConfig:
def test_top_p_defaults_to_none(self):
config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0")
assert config.top_p is None
def test_top_p_explicit_value_stored(self):
config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0", top_p=0.8)
assert config.top_p == 0.8
def test_get_model_config_excludes_top_p_by_default(self):
config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0", temperature=0.5)
model_cfg = config.get_model_config()
assert "top_p" not in model_cfg
def test_get_model_config_includes_top_p_when_set(self):
config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0", top_p=0.7)
model_cfg = config.get_model_config()
assert model_cfg["top_p"] == 0.7
def test_get_model_config_top_p_via_model_kwargs(self):
"""model_kwargs can supply top_p after merge; same semantics as explicit top_p."""
config = AWSBedrockConfig(
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
model_kwargs={"top_p": 0.88},
)
assert config.get_model_config()["top_p"] == 0.88
def test_get_model_config_always_includes_temperature(self):
config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0", temperature=0.3)
model_cfg = config.get_model_config()
assert model_cfg["temperature"] == 0.3
def test_aws_region_stored(self):
config = AWSBedrockConfig(
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
aws_region="us-east-2",
)
assert config.aws_region == "us-east-2"
# ---------------------------------------------------------------------------
# LlmFactory
# ---------------------------------------------------------------------------
class TestLlmFactory:
def test_aws_bedrock_uses_aws_bedrock_config(self):
_, config_class = LlmFactory.provider_to_class["aws_bedrock"]
assert config_class is AWSBedrockConfig
def test_factory_create_accepts_aws_region(self, mock_boto3):
"""LlmFactory.create must not crash when aws_region is in the config dict.
Before the fix, the factory mapped aws_bedrock to BaseLlmConfig which has
no aws_region parameter, causing: TypeError: __init__() got an unexpected
keyword argument 'aws_region'.
"""
user_dict = {
"model": "anthropic.claude-3-5-sonnet-20240620-v1:0",
"aws_region": "us-east-2",
"temperature": 0.1,
"max_tokens": 2000,
}
mock_boto3.converse.return_value = _converse_response()
# This must not raise TypeError
llm = LlmFactory.create("aws_bedrock", user_dict)
assert isinstance(llm.config, AWSBedrockConfig)
assert llm.config.aws_region == "us-east-2"
assert llm.config.top_p is None
# ---------------------------------------------------------------------------
# _build_inference_config
# ---------------------------------------------------------------------------
class TestBuildInferenceConfig:
"""
Unit tests for the _build_inference_config helper.
Validates the exact keys present in the returned dict.
"""
def test_anthropic_only_temperature_by_default(self, mock_boto3):
llm = _make_llm("anthropic.claude-3-5-sonnet-20240620-v1:0", mock_boto3, temperature=0.5)
cfg = llm._build_inference_config()
assert "temperature" in cfg
assert cfg["temperature"] == 0.5
assert "topP" not in cfg, "topP must be absent when top_p not configured"
def test_anthropic_top_p_explicitly_set_still_omits_top_p(self, mock_boto3):
# Anthropic rejects both; topP must still be omitted even when set
llm = _make_llm(
"anthropic.claude-3-5-sonnet-20240620-v1:0",
mock_boto3,
temperature=0.5,
top_p=0.9,
)
cfg = llm._build_inference_config()
assert "temperature" in cfg
assert "topP" not in cfg
def test_anthropic_inference_profile_omits_top_p(self, mock_boto3):
# Cross-region inference profiles (us.anthropic.*) follow the same rule
llm = _make_llm("us.anthropic.claude-haiku-4-5-20251001-v1:0", mock_boto3, temperature=0.1)
cfg = llm._build_inference_config()
assert "topP" not in cfg
def test_amazon_includes_top_p_when_set(self, mock_boto3):
llm = _make_llm(
"amazon.nova-3-mini-20241119-v1:0",
mock_boto3,
temperature=0.5,
top_p=0.85,
)
cfg = llm._build_inference_config()
assert cfg["temperature"] == 0.5
assert cfg["topP"] == 0.85
def test_amazon_omits_top_p_when_not_set(self, mock_boto3):
llm = _make_llm("amazon.nova-3-mini-20241119-v1:0", mock_boto3, temperature=0.5)
cfg = llm._build_inference_config()
assert "topP" not in cfg
def test_max_tokens_present(self, mock_boto3):
llm = _make_llm(
"anthropic.claude-3-5-sonnet-20240620-v1:0",
mock_boto3,
max_tokens=1024,
)
cfg = llm._build_inference_config()
assert cfg["maxTokens"] == 1024
def test_nova_fallback_max_tokens_when_absent(self, mock_boto3):
"""Legacy Nova Converse used 5000 when max_tokens was missing from the dict."""
llm = _make_llm("amazon.nova-3-mini-20241119-v1:0", mock_boto3)
llm.model_config.pop("max_tokens", None)
cfg = llm._build_inference_config()
assert cfg["maxTokens"] == 5000
def test_anthropic_fallback_max_tokens_when_absent(self, mock_boto3):
llm = _make_llm("anthropic.claude-3-5-sonnet-20240620-v1:0", mock_boto3)
llm.model_config.pop("max_tokens", None)
cfg = llm._build_inference_config()
assert cfg["maxTokens"] == 2000
# ---------------------------------------------------------------------------
# generate_response — Converse API call assertions
# ---------------------------------------------------------------------------
MESSAGES = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello"},
]
TOOLS = [
{
"type": "function",
"function": {
"name": "add_memory",
"description": "Store a memory",
"parameters": {
"type": "object",
"properties": {"data": {"type": "string"}},
"required": ["data"],
},
},
}
]
class TestGenerateResponseConverse:
def test_standard_anthropic_no_top_p_in_converse_call(self, mock_boto3):
mock_boto3.converse.return_value = _converse_response()
llm = _make_llm("anthropic.claude-3-5-sonnet-20240620-v1:0", mock_boto3, temperature=0.2)
llm.generate_response(MESSAGES)
_, kwargs = mock_boto3.converse.call_args
inference_cfg = kwargs["inferenceConfig"]
assert "topP" not in inference_cfg
assert inference_cfg["temperature"] == 0.2
def test_standard_anthropic_inference_profile_no_top_p(self, mock_boto3):
mock_boto3.converse.return_value = _converse_response()
llm = _make_llm("us.anthropic.claude-haiku-4-5-20251001-v1:0", mock_boto3, temperature=0.1)
llm.generate_response(MESSAGES)
_, kwargs = mock_boto3.converse.call_args
assert "topP" not in kwargs["inferenceConfig"]
def test_with_tools_anthropic_no_top_p(self, mock_boto3):
mock_boto3.converse.return_value = _converse_response()
llm = _make_llm("anthropic.claude-3-5-sonnet-20240620-v1:0", mock_boto3, temperature=0.2)
llm.generate_response(MESSAGES, tools=TOOLS)
_, kwargs = mock_boto3.converse.call_args
assert "topP" not in kwargs["inferenceConfig"]
def test_with_tools_anthropic_top_p_set_still_omitted(self, mock_boto3):
# Even if user explicitly sets top_p, Anthropic must not receive topP
mock_boto3.converse.return_value = _converse_response()
llm = _make_llm(
"anthropic.claude-3-5-sonnet-20240620-v1:0",
mock_boto3,
temperature=0.2,
top_p=0.9,
)
llm.generate_response(MESSAGES, tools=TOOLS)
_, kwargs = mock_boto3.converse.call_args
assert "topP" not in kwargs["inferenceConfig"]
def test_nova_includes_top_p_when_explicitly_set(self, mock_boto3):
mock_boto3.converse.return_value = _converse_response()
llm = _make_llm(
"amazon.nova-3-mini-20241119-v1:0",
mock_boto3,
temperature=0.5,
top_p=0.85,
)
llm.generate_response(MESSAGES)
_, kwargs = mock_boto3.converse.call_args
assert kwargs["inferenceConfig"]["topP"] == 0.85
def test_nova_omits_top_p_when_not_set(self, mock_boto3):
mock_boto3.converse.return_value = _converse_response()
llm = _make_llm("amazon.nova-3-mini-20241119-v1:0", mock_boto3, temperature=0.5)
llm.generate_response(MESSAGES)
_, kwargs = mock_boto3.converse.call_args
assert "topP" not in kwargs["inferenceConfig"]
def test_anthropic_model_kwargs_top_p_still_omits_top_p_in_converse(self, mock_boto3):
"""top_p injected via model_kwargs must not add topP for Anthropic Converse."""
mock_boto3.converse.return_value = _converse_response()
llm = _make_llm(
"anthropic.claude-3-5-sonnet-20240620-v1:0",
mock_boto3,
temperature=0.2,
model_kwargs={"top_p": 0.88},
)
llm.generate_response(MESSAGES)
_, kwargs = mock_boto3.converse.call_args
assert "topP" not in kwargs["inferenceConfig"]