diff --git a/sagemaker-train/src/sagemaker/train/evaluate/multi_turn_rl_evaluator.py b/sagemaker-train/src/sagemaker/train/evaluate/multi_turn_rl_evaluator.py index cb5184cdc6..754c0944e4 100644 --- a/sagemaker-train/src/sagemaker/train/evaluate/multi_turn_rl_evaluator.py +++ b/sagemaker-train/src/sagemaker/train/evaluate/multi_turn_rl_evaluator.py @@ -253,10 +253,16 @@ def _resolve_trainer_defaults(self) -> None: or getattr(trainer, "_agent_config", None) or getattr(trainer, "agent_env", None) ) - # AgentRFTJob.agent_config returns a dict like {"AgentRuntimeArn": "..."} + # AgentRFTJob.agent_config returns a nested dict like + # {"BedrockAgentCoreConfig": {"AgentRuntimeArn": "arn:..."}} or + # {"CustomAgentLambdaConfig": {"LambdaArn": "arn:..."}} if isinstance(resolved_agent, dict): + bedrock_cfg = resolved_agent.get("BedrockAgentCoreConfig", {}) + lambda_cfg = resolved_agent.get("CustomAgentLambdaConfig", {}) self.agent_config = ( - resolved_agent.get("AgentRuntimeArn") + bedrock_cfg.get("AgentRuntimeArn") + or lambda_cfg.get("LambdaArn") + or resolved_agent.get("AgentRuntimeArn") or resolved_agent.get("AgentLambdaArn") or resolved_agent.get("LambdaArn") ) diff --git a/sagemaker-train/tests/integ/train/test_mtrl_evaluator.py b/sagemaker-train/tests/integ/train/test_mtrl_evaluator.py index 9d493a0508..264aa17eec 100644 --- a/sagemaker-train/tests/integ/train/test_mtrl_evaluator.py +++ b/sagemaker-train/tests/integ/train/test_mtrl_evaluator.py @@ -301,6 +301,47 @@ def test_evaluator_construction_with_base_model(self, test_config): assert evaluator is not None assert evaluator.model == test_config["base_model"] + def test_evaluator_infers_agent_config_from_trainer(self, mtrl_trainer, test_config): + """Test that agent_config is inferred from trainer's nested dict (no explicit agent_config).""" + # Simulate AgentRFTJob.agent_config returning a nested dict + mtrl_trainer.agent_config = { + "BedrockAgentCoreConfig": {"AgentRuntimeArn": test_config["agent_arn"]} + } + + evaluator = MultiTurnRLEvaluator( + model=mtrl_trainer, + dataset=test_config["dataset"], + s3_output_path=f'{test_config["s3_output_path"]}integ-infer-agent/', + mlflow_resource_arn=test_config["mlflow_resource_arn"], + region=test_config["region"], + ) + + evaluator._resolve_trainer_defaults() + assert evaluator.agent_config == test_config["agent_arn"] + + # Clean up: remove the attribute so other tests using this fixture aren't affected + del mtrl_trainer.agent_config + + def test_evaluator_infers_lambda_agent_config_from_trainer(self, mtrl_trainer, test_config): + """Test that agent_config is inferred from trainer's nested CustomAgentLambdaConfig dict.""" + lambda_arn = "arn:aws:lambda:us-west-2:123456789012:function:my-agent" + mtrl_trainer.agent_config = { + "CustomAgentLambdaConfig": {"LambdaArn": lambda_arn} + } + + evaluator = MultiTurnRLEvaluator( + model=mtrl_trainer, + dataset=test_config["dataset"], + s3_output_path=f'{test_config["s3_output_path"]}integ-infer-lambda/', + mlflow_resource_arn=test_config["mlflow_resource_arn"], + region=test_config["region"], + ) + + evaluator._resolve_trainer_defaults() + assert evaluator.agent_config == lambda_arn + + del mtrl_trainer.agent_config + def test_get_all_mtrl_evaluations(self, test_config): """Test listing all MTRL evaluation executions.""" all_execs = MultiTurnRLEvaluator.get_all(region=test_config["region"]) diff --git a/sagemaker-train/tests/unit/train/evaluate/test_mtrl_evaluator_agent_config.py b/sagemaker-train/tests/unit/train/evaluate/test_mtrl_evaluator_agent_config.py new file mode 100644 index 0000000000..a2dae57c8a --- /dev/null +++ b/sagemaker-train/tests/unit/train/evaluate/test_mtrl_evaluator_agent_config.py @@ -0,0 +1,95 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Unit tests for _resolve_trainer_defaults agent_config dict parsing.""" +from __future__ import absolute_import + +import pytest +from unittest.mock import MagicMock + +from sagemaker.train.evaluate.multi_turn_rl_evaluator import MultiTurnRLEvaluator + + +AGENT_ARN = "arn:aws:bedrock-agentcore:us-west-2:123456789012:runtime/test-agent-aBcDeFgHiJ" +LAMBDA_ARN = "arn:aws:lambda:us-west-2:123456789012:function:my-agent" +SOURCE_MP_ARN = "arn:aws:sagemaker:us-west-2:123456789012:model-package/test-mpg/1" +BASE_MODEL = "huggingface-reasoning-qwen3-32b" + + +def _make_evaluator_with_trainer(agent_config_dict): + """Create a mock evaluator + trainer to test _resolve_trainer_defaults directly.""" + trainer = MagicMock() + type(trainer).__name__ = "MultiTurnRLTrainer" + trainer.output_model_package_arn = SOURCE_MP_ARN + trainer.model_package_arn = SOURCE_MP_ARN + trainer.base_model_arn = "arn:aws:sagemaker:us-west-2:aws:hub-content/test" + trainer.base_model_name = BASE_MODEL + trainer.agent_config = agent_config_dict + trainer._agent_config = None + trainer.agent_env = None + trainer.agent_qualifier = None + trainer._agent_qualifier = None + + evaluator = MagicMock() + evaluator.model = trainer + evaluator.agent_config = None + evaluator.agent_qualifier = None + evaluator._source_model_package_arn_cache = None + evaluator._base_model_arn_cache = None + evaluator._base_model_name_cache = None + return evaluator + + +class TestResolveTrainerAgentConfig: + """Tests for _resolve_trainer_defaults agent_config dict parsing.""" + + def test_nested_bedrock_agent_core_config(self): + """Test that nested BedrockAgentCoreConfig dict is correctly parsed.""" + evaluator = _make_evaluator_with_trainer( + {"BedrockAgentCoreConfig": {"AgentRuntimeArn": AGENT_ARN}} + ) + MultiTurnRLEvaluator._resolve_trainer_defaults(evaluator) + assert evaluator.agent_config == AGENT_ARN + + def test_nested_custom_agent_lambda_config(self): + """Test that nested CustomAgentLambdaConfig dict is correctly parsed.""" + evaluator = _make_evaluator_with_trainer( + {"CustomAgentLambdaConfig": {"LambdaArn": LAMBDA_ARN}} + ) + MultiTurnRLEvaluator._resolve_trainer_defaults(evaluator) + assert evaluator.agent_config == LAMBDA_ARN + + def test_flat_dict_fallback_agent_runtime_arn(self): + """Test that flat AgentRuntimeArn key still works as a fallback.""" + evaluator = _make_evaluator_with_trainer( + {"AgentRuntimeArn": AGENT_ARN} + ) + MultiTurnRLEvaluator._resolve_trainer_defaults(evaluator) + assert evaluator.agent_config == AGENT_ARN + + def test_flat_dict_fallback_lambda_arn(self): + """Test that flat LambdaArn key still works as a fallback.""" + evaluator = _make_evaluator_with_trainer( + {"LambdaArn": LAMBDA_ARN} + ) + MultiTurnRLEvaluator._resolve_trainer_defaults(evaluator) + assert evaluator.agent_config == LAMBDA_ARN + + def test_customer_provided_agent_config_not_overwritten(self): + """Test that customer-provided agent_config is not overwritten by trainer.""" + customer_arn = "arn:aws:bedrock-agentcore:us-west-2:123456789012:runtime/customer-agent" + evaluator = _make_evaluator_with_trainer( + {"BedrockAgentCoreConfig": {"AgentRuntimeArn": AGENT_ARN}} + ) + evaluator.agent_config = customer_arn + MultiTurnRLEvaluator._resolve_trainer_defaults(evaluator) + assert evaluator.agent_config == customer_arn