diff --git a/docs/components/llms/models/aws_bedrock.mdx b/docs/components/llms/models/aws_bedrock.mdx index b58fbd389..089ec3c05 100644 --- a/docs/components/llms/models/aws_bedrock.mdx +++ b/docs/components/llms/models/aws_bedrock.mdx @@ -78,6 +78,38 @@ await memory.add(messages, { userId: 'alice', metadata: { category: 'movies' } } The TypeScript provider calls the Bedrock [Converse API](https://docs.aws.amazon.com/bedrock/latest/userguide/conversation-inference.html), a single uniform interface across the current Bedrock model families. Streaming and `InvokeModel`-only models are not supported yet. +### Application inference profiles + +Bedrock resolves the model family from the model identifier. An application inference profile ARN ends in an opaque ID, so there is nothing to resolve from. Set `provider_override` (Python) / `providerOverride` (TypeScript) when your model is one: + + +```python Python +config = { + "llm": { + "provider": "aws_bedrock", + "config": { + "model": "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz", + "provider_override": "anthropic", + } + } +} +``` + +```typescript TypeScript +const config = { + llm: { + provider: 'aws_bedrock', + config: { + model: 'arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz', + providerOverride: 'anthropic', + }, + }, +}; +``` + + +Without it, initialization raises `Unknown provider in model` (Python: `ValueError`; TypeScript: `Error`). Plain model IDs and cross-region inference profiles such as `us.anthropic.claude-sonnet-4-20250514-v1:0` still resolve automatically and need no override. + ### Config All available parameters for the `aws_bedrock` config are present in [Master List of All Params in Config](../config). diff --git a/mem0-ts/src/oss/src/llms/aws_bedrock.ts b/mem0-ts/src/oss/src/llms/aws_bedrock.ts index fd1871472..39c8b2b56 100644 --- a/mem0-ts/src/oss/src/llms/aws_bedrock.ts +++ b/mem0-ts/src/oss/src/llms/aws_bedrock.ts @@ -29,7 +29,18 @@ const PROVIDERS = [ * Extract the model-family provider from a Bedrock model id * (e.g. `anthropic.claude-3-sonnet-...` -> `anthropic`). */ -export function extractProvider(model: string): string { +export function extractProvider( + model: string, + providerOverride?: string, +): string { + if (providerOverride) { + if (!PROVIDERS.includes(providerOverride)) { + throw new Error( + `Unknown providerOverride '${providerOverride}'. Valid providers: ${PROVIDERS.join(", ")}`, + ); + } + return providerOverride; + } for (const provider of PROVIDERS) { const re = new RegExp( `\\b${provider.replace(/[.*+?^${}()|[\]\\]/g, "\\$&")}\\b`, @@ -79,7 +90,7 @@ export class AWSBedrockLLM implements LLM { this.model = (typeof config.model === "string" && config.model) || "anthropic.claude-3-5-sonnet-20240620-v1:0"; - this.provider = extractProvider(this.model); + this.provider = extractProvider(this.model, config.providerOverride); this.temperature = config.temperature ?? 0.1; this.maxTokens = config.maxTokens ?? 2000; this.topP = config.topP; diff --git a/mem0-ts/src/oss/src/tests/aws_bedrock.test.ts b/mem0-ts/src/oss/src/tests/aws_bedrock.test.ts index b02ab880e..7ddb48caf 100644 --- a/mem0-ts/src/oss/src/tests/aws_bedrock.test.ts +++ b/mem0-ts/src/oss/src/tests/aws_bedrock.test.ts @@ -68,6 +68,29 @@ describe("extractProvider", () => { /Unknown provider/, ); }); + + const ARN = + "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz"; + + it("throws for an application inference profile ARN without providerOverride", () => { + expect(() => extractProvider(ARN)).toThrow(/Unknown provider in model/); + }); + + it("resolves an application inference profile ARN via providerOverride", () => { + expect(extractProvider(ARN, "anthropic")).toBe("anthropic"); + }); + + it("lets providerOverride take precedence over regex detection", () => { + expect( + extractProvider("anthropic.claude-3-5-sonnet-20240620-v1:0", "amazon"), + ).toBe("amazon"); + }); + + it("throws on a misspelled providerOverride", () => { + expect(() => extractProvider(ARN, "anthorpic")).toThrow( + /Unknown providerOverride 'anthorpic'/, + ); + }); }); describe("AWSBedrockLLM", () => { @@ -176,6 +199,28 @@ describe("AWSBedrockLLM", () => { expect(res).toEqual({ content: "hello from bedrock", role: "assistant" }); }); + it("uses providerOverride to resolve inference profile ARNs (topP omitted like anthropic)", async () => { + const arn = + "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz"; + const client = new FakeBedrockClient(textResponse); + const llm = makeLLM(client, { + model: arn, + providerOverride: "anthropic", + topP: 0.9, + }); + await llm.generateResponse([{ role: "user", content: "hi" }]); + expect(client.lastInput.modelId).toBe(arn); + expect(client.lastInput.inferenceConfig.topP).toBeUndefined(); + }); + + it("throws when constructed with an inference profile ARN and no providerOverride", () => { + const arn = + "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz"; + expect(() => new AWSBedrockLLM({ model: arn })).toThrow( + /Unknown provider in model/, + ); + }); + it("does not construct the Bedrock client until the first request", () => { const before = mockClientConstructions; // No injected client: the old constructor eagerly built a real one. diff --git a/mem0-ts/src/oss/src/types/index.ts b/mem0-ts/src/oss/src/types/index.ts index b2aa42efe..fb3e5b391 100644 --- a/mem0-ts/src/oss/src/types/index.ts +++ b/mem0-ts/src/oss/src/types/index.ts @@ -78,6 +78,7 @@ export interface LLMConfig { awsAccessKeyId?: string; awsSecretAccessKey?: string; awsSessionToken?: string; + providerOverride?: string; // Optional pre-constructed client (e.g. BedrockRuntimeClient) for DI/testing. client?: any; } diff --git a/mem0/configs/llms/aws_bedrock.py b/mem0/configs/llms/aws_bedrock.py index 9b0e65391..075d62b4e 100644 --- a/mem0/configs/llms/aws_bedrock.py +++ b/mem0/configs/llms/aws_bedrock.py @@ -24,6 +24,7 @@ class AWSBedrockConfig(BaseLlmConfig): aws_session_token: Optional[str] = None, aws_profile: Optional[str] = None, model_kwargs: Optional[Dict[str, Any]] = None, + provider_override: Optional[str] = None, **kwargs, ): """ @@ -42,6 +43,10 @@ class AWSBedrockConfig(BaseLlmConfig): aws_session_token: AWS session token for temporary credentials aws_profile: AWS profile name for credentials model_kwargs: Additional model-specific parameters + provider_override: Explicit provider name (e.g. "anthropic"), required when + model is an application inference profile ARN whose opaque ID has no + provider substring for automatic detection. Defaults to None (uses + automatic detection from the model identifier). **kwargs: Additional arguments passed to base class """ super().__init__( @@ -59,6 +64,7 @@ class AWSBedrockConfig(BaseLlmConfig): self.aws_session_token = aws_session_token self.aws_profile = aws_profile self.model_kwargs = model_kwargs or {} + self.provider_override = provider_override @property def provider(self) -> str: diff --git a/mem0/llms/aws_bedrock.py b/mem0/llms/aws_bedrock.py index d749ebb27..20b7f3883 100644 --- a/mem0/llms/aws_bedrock.py +++ b/mem0/llms/aws_bedrock.py @@ -9,8 +9,8 @@ try: except ImportError: raise ImportError("The 'boto3' library is required. Please install it using 'pip install boto3'.") -from mem0.configs.llms.base import BaseLlmConfig from mem0.configs.llms.aws_bedrock import AWSBedrockConfig +from mem0.configs.llms.base import BaseLlmConfig from mem0.llms.base import LLMBase from mem0.memory.utils import extract_json @@ -23,8 +23,14 @@ PROVIDERS = [ ] -def extract_provider(model: str) -> str: - """Extract provider from model identifier.""" +def extract_provider(model: str, explicit_provider: Optional[str] = None) -> str: + """Extract provider from model identifier, or return explicit_provider when set.""" + if explicit_provider: + if explicit_provider not in PROVIDERS: + raise ValueError( + f"Unknown provider_override '{explicit_provider}'. Valid providers: {', '.join(PROVIDERS)}" + ) + return explicit_provider for provider in PROVIDERS: if re.search(rf"\b{re.escape(provider)}\b", model): return provider @@ -69,7 +75,7 @@ class AWSBedrockLLM(LLMBase): # Get model configuration self.model_config = self.config.get_model_config() - self.provider = extract_provider(self.config.model) + self.provider = extract_provider(self.config.model, self.config.provider_override) # Initialize provider-specific settings self._initialize_provider_settings() diff --git a/tests/llms/test_aws_bedrock.py b/tests/llms/test_aws_bedrock.py index 059cce650..f683f8c3f 100644 --- a/tests/llms/test_aws_bedrock.py +++ b/tests/llms/test_aws_bedrock.py @@ -6,7 +6,6 @@ from mem0.configs.llms.aws_bedrock import AWSBedrockConfig from mem0.llms.aws_bedrock import AWSBedrockLLM, extract_provider from mem0.utils.factory import LlmFactory - # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- @@ -70,6 +69,23 @@ class TestExtractProvider: with pytest.raises(ValueError, match="Unknown provider"): extract_provider("unknown-vendor.some-model-v1:0") + def test_application_inference_profile_arn_without_override_raises(self): + arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz" + with pytest.raises(ValueError, match="Unknown provider"): + extract_provider(arn) + + def test_application_inference_profile_arn_with_explicit_provider(self): + arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz" + assert extract_provider(arn, "anthropic") == "anthropic" + + def test_explicit_provider_takes_precedence_over_regex(self): + assert extract_provider("anthropic.claude-3-5-sonnet-20240620-v1:0", "amazon") == "amazon" + + def test_explicit_provider_typo_raises(self): + arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz" + with pytest.raises(ValueError, match="Unknown provider_override 'anthorpic'"): + extract_provider(arn, "anthorpic") + # --------------------------------------------------------------------------- # AWSBedrockConfig @@ -114,6 +130,44 @@ class TestAWSBedrockConfig: ) assert config.aws_region == "us-east-2" + def test_provider_override_defaults_to_none(self): + config = AWSBedrockConfig(model="anthropic.claude-3-5-sonnet-20240620-v1:0") + assert config.provider_override is None + + def test_provider_override_stored(self): + arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz" + config = AWSBedrockConfig(model=arn, provider_override="anthropic") + assert config.provider_override == "anthropic" + + +# --------------------------------------------------------------------------- +# AWSBedrockLLM with application inference profile ARNs +# --------------------------------------------------------------------------- + +class TestApplicationInferenceProfileArn: + ARN = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz" + + def test_arn_without_provider_override_raises(self, mock_boto3): + with pytest.raises(ValueError, match="Unknown provider"): + _make_llm(self.ARN, mock_boto3) + + def test_arn_with_provider_override_resolves(self, mock_boto3): + llm = _make_llm(self.ARN, mock_boto3, provider_override="anthropic") + assert llm.provider == "anthropic" + assert llm.supports_tools is True + + def test_plain_model_id_unaffected(self, mock_boto3): + llm = _make_llm("anthropic.claude-3-5-sonnet-20240620-v1:0", mock_boto3) + assert llm.provider == "anthropic" + + def test_cross_region_inference_profile_unaffected(self, mock_boto3): + llm = _make_llm("us.anthropic.claude-haiku-4-5-20251001-v1:0", mock_boto3) + assert llm.provider == "anthropic" + + def test_arn_with_misspelled_provider_override_raises(self, mock_boto3): + with pytest.raises(ValueError, match="Unknown provider_override 'anthorpic'"): + _make_llm(self.ARN, mock_boto3, provider_override="anthorpic") + # --------------------------------------------------------------------------- # LlmFactory