Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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")
)
Expand Down
41 changes: 41 additions & 0 deletions sagemaker-train/tests/integ/train/test_mtrl_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"])
Expand Down
Original file line number Diff line number Diff line change
@@ -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
Loading