diff --git a/src/strands_tools/code_interpreter/agent_core_code_interpreter.py b/src/strands_tools/code_interpreter/agent_core_code_interpreter.py index ab7771fe..b86e4328 100644 --- a/src/strands_tools/code_interpreter/agent_core_code_interpreter.py +++ b/src/strands_tools/code_interpreter/agent_core_code_interpreter.py @@ -21,8 +21,19 @@ logger = logging.getLogger(__name__) -# Module-level session cache - persists across object instances -_session_mapping: Dict[str, str] = {} # user_session_name -> aws_session_id +# Module-level session cache - lets a new object instance in the same process reconnect to +# an AWS code interpreter session a previous instance started, instead of paying to recreate +# it. Keys are partitioned per caller (see partition_key / _scoped_key) so that a session name +# supplied in one caller's partition can never address a session created in another's. +# +# The real isolation boundary is AWS-side, not this map: a session is reachable only by holding +# its server-issued, unguessable sessionId, and every operation is authorized against the +# interpreter's IAM execution role. This cache is a per-process reconnect convenience. The +# partition_key is what keeps callers apart within one process; it is bound at construction +# from operator-controlled input and is never taken from a tool action, so a model-supplied +# session name cannot reach across partitions. Multi-tenant deployments should pass a distinct +# partition_key per caller (e.g. the authenticated principal or agent session id). +_session_mapping: Dict[str, str] = {} # "partition_key\x00session_name" -> aws_session_id @dataclass @@ -56,6 +67,7 @@ def __init__( persist_sessions: bool = True, session_timeout_seconds: int = 900, boto_session: Optional[boto3.Session] = None, + partition_key: Optional[str] = None, ) -> None: """ Initialize the Bedrock AgentCore code interpreter with session persistence support. @@ -112,6 +124,18 @@ def __init__( enabling cross-account access or custom credential configurations. If None (default), the client uses the default credential chain. + partition_key (Optional[str]): Caller partition for the in-process reconnect cache. + Session names are namespaced under this key, so a session name used in one + partition can never address a session created in another. Bind it at construction + from operator-controlled input (it is never read from a tool action), passing a + distinct value per caller in multi-tenant deployments: + - The authenticated principal / tenant id, or + - The agent session id (e.g. context.session_id), when one identity may run + several isolated sessions. + Defaults to a shared "default" partition when not provided, which preserves the + prior single-caller reconnect behavior. The cache key is no longer derived from + AWS credentials. + Session Lifecycle: Invocation 1 (Instance #1): 1. Create new instance with session_name="user-abc-123" @@ -173,10 +197,18 @@ def invoke(payload, context): Notes: - Module-level cache persists in long-running Python processes (AgentCore) - Cache does NOT persist across container restarts (cold starts) - - Session names must be unique per user/conversation for isolation + - Cache keys are namespaced by partition_key, so a session name used in one + partition cannot address a session created in another. partition_key is bound at + construction and is never taken from a tool action - AWS session IDs are globally unique (ULID format) - Sessions can be manually stopped via AWS console/API if needed + The in-process cache is a reconnect convenience, NOT the isolation boundary. The + real boundary is AWS-side: a session is reachable only by holding its server-issued, + unguessable sessionId, and every operation is authorized against the interpreter's + IAM execution role. Within a single process, partition_key is what keeps callers + apart, so multi-tenant deployments must pass a distinct partition_key per caller. + Raises: ValueError: If auto_create=False and session doesn't exist, or if session_name is already in use by another instance @@ -200,14 +232,25 @@ def invoke(payload, context): else: self.default_session = session_name + # Partition for the in-process reconnect cache. Bound here from operator-controlled + # input and never read from a tool action, so a model-supplied session name cannot + # reach across partitions. Defaults to a shared "default" partition, preserving the + # prior single-caller reconnect behavior; multi-tenant callers must pass a distinct + # partition_key per principal to keep their sessions isolated. + self.partition_key = partition_key if partition_key is not None else "default" + self._sessions: Dict[str, SessionInfo] = {} logger.info( f"Initialized CodeInterpreter with session='{self.default_session}', " - f"identifier='{self.identifier}', auto_create={auto_create}, " - f"persist_sessions={persist_sessions}" + f"partition='{self.partition_key}', identifier='{self.identifier}', " + f"auto_create={auto_create}, persist_sessions={persist_sessions}" ) + def _scoped_key(self, session_name: str) -> str: + """Build the module-level cache key for a session name under this instance's partition.""" + return f"{self.partition_key}\x00{session_name}" + def start_platform(self) -> None: """Initialize the Bedrock AgentCoreplatform connection.""" pass @@ -262,8 +305,8 @@ def init_session(self, action: InitSessionAction) -> Dict[str, Any]: if session_name in self._sessions: return {"status": "error", "content": [{"text": f"Session '{session_name}' already exists"}]} - # Check if session name already in use (module-level cache) - if session_name in _session_mapping: + # Check if session name already in use (module-level cache, scoped to this partition) + if self._scoped_key(session_name) in _session_mapping: error_msg = ( f"Session '{session_name}' is already in use by another instance. " f"Use a unique session name or reconnect to the existing session " @@ -284,8 +327,8 @@ def init_session(self, action: InitSessionAction) -> Dict[str, Any]: aws_session_id = client.session_id - # Store mapping in module-level cache - _session_mapping[session_name] = aws_session_id + # Store mapping in module-level cache (scoped to this partition) + _session_mapping[self._scoped_key(session_name)] = aws_session_id # Store session info locally self._sessions[session_name] = SessionInfo( @@ -360,8 +403,9 @@ def _ensure_session(self, session_name: Optional[str]) -> tuple[str, Optional[Di logger.debug(f"Using cached session: {target_session}") return target_session, None - # Check module-level cache for AWS session ID - aws_session_id = _session_mapping.get(target_session) + # Check module-level cache for AWS session ID (scoped to this partition) + scoped_key = self._scoped_key(target_session) + aws_session_id = _session_mapping.get(scoped_key) if aws_session_id: # Found in module cache - try to reconnect @@ -387,13 +431,13 @@ def _ensure_session(self, session_name: Optional[str]) -> tuple[str, Optional[Di else: # Session exists but not ready - remove from cache logger.warning(f"Session {target_session} not READY, removing from cache") - del _session_mapping[target_session] + del _session_mapping[scoped_key] except Exception as e: # Session doesn't exist or error - remove from cache logger.debug(f"Session reconnection failed: {e}") - if target_session in _session_mapping: - del _session_mapping[target_session] + if scoped_key in _session_mapping: + del _session_mapping[scoped_key] # Session not found - create new if auto_create enabled if self.auto_create: diff --git a/tests/code_interpreter/test_agent_core_code_interpreter.py b/tests/code_interpreter/test_agent_core_code_interpreter.py index a9174f2d..1576c1d0 100644 --- a/tests/code_interpreter/test_agent_core_code_interpreter.py +++ b/tests/code_interpreter/test_agent_core_code_interpreter.py @@ -133,9 +133,6 @@ def test_ensure_session_reconnection_via_module_cache(mock_client): ) as mock_client_class: mock_resolve.return_value = "us-west-2" - # Setup module cache with existing session - _session_mapping["target-session"] = "found-session-id" - reconnect_client = MagicMock() reconnect_client.session_id = "found-session-id" reconnect_client.get_session.return_value = {"status": "READY"} @@ -143,6 +140,9 @@ def test_ensure_session_reconnection_via_module_cache(mock_client): interpreter = AgentCoreCodeInterpreter(region="us-west-2", persist_sessions=True) + # Setup module cache with existing session (keyed per caller identity) + _session_mapping[interpreter._scoped_key("target-session")] = "found-session-id" + session_name, error = interpreter._ensure_session("target-session") assert session_name == "target-session" @@ -164,9 +164,6 @@ def test_ensure_session_module_cache_session_not_ready(): ) as mock_client_class: mock_resolve.return_value = "us-west-2" - # Setup module cache - _session_mapping["stale-session"] = "stale-session-id" - reconnect_client = MagicMock() reconnect_client.session_id = "new-session-id" reconnect_client.get_session.return_value = {"status": "STOPPED"} @@ -174,6 +171,9 @@ def test_ensure_session_module_cache_session_not_ready(): interpreter = AgentCoreCodeInterpreter(region="us-west-2", persist_sessions=True, auto_create=True) + # Setup module cache (keyed per caller identity) + _session_mapping[interpreter._scoped_key("stale-session")] = "stale-session-id" + with patch.object(interpreter, "init_session") as mock_init: mock_init.return_value = {"status": "success", "content": [{"text": "Created"}]} @@ -182,7 +182,7 @@ def test_ensure_session_module_cache_session_not_ready(): assert session_name == "stale-session" assert error is None # Session should be removed from cache - assert "stale-session" not in _session_mapping + assert interpreter._scoped_key("stale-session") not in _session_mapping # New session should be created mock_init.assert_called_once() @@ -195,9 +195,6 @@ def test_ensure_session_module_cache_get_session_fails(): ) as mock_client_class: mock_resolve.return_value = "us-west-2" - # Setup module cache - _session_mapping["missing-session"] = "missing-session-id" - reconnect_client = MagicMock() reconnect_client.get_session.side_effect = Exception("Session not found") reconnect_client.session_id = "new-session-id" @@ -205,6 +202,9 @@ def test_ensure_session_module_cache_get_session_fails(): interpreter = AgentCoreCodeInterpreter(region="us-west-2", persist_sessions=True, auto_create=True) + # Setup module cache (keyed per caller identity) + _session_mapping[interpreter._scoped_key("missing-session")] = "missing-session-id" + with patch.object(interpreter, "init_session") as mock_init: mock_init.return_value = {"status": "success", "content": [{"text": "Created"}]} @@ -213,7 +213,7 @@ def test_ensure_session_module_cache_get_session_fails(): assert session_name == "missing-session" assert error is None # Session should be removed from cache after error - assert "missing-session" not in _session_mapping + assert interpreter._scoped_key("missing-session") not in _session_mapping mock_init.assert_called_once() @@ -425,8 +425,8 @@ def test_init_session_success(mock_client_class, interpreter, mock_client): assert session_info.description == "Test session" assert session_info.client == mock_client - # Check module-level cache - assert _session_mapping.get("my-session") == "test-session-id-123" + # Check module-level cache (keyed per caller identity) + assert _session_mapping.get(interpreter._scoped_key("my-session")) == "test-session-id-123" @patch("strands_tools.code_interpreter.agent_core_code_interpreter.BedrockAgentCoreCodeInterpreterClient") @@ -903,7 +903,7 @@ def test_module_level_session_mapping(): result = interpreter1.init_session(action) assert result["status"] == "success" - assert _session_mapping["shared-session"] == "aws-session-123" + assert _session_mapping[interpreter1._scoped_key("shared-session")] == "aws-session-123" # Second instance should find session in module cache mock_client2 = MagicMock() @@ -957,16 +957,133 @@ def test_ensure_session_passes_boto_session_on_reconnect(mock_client_class): mock_resolve.return_value = "us-west-2" mock_session = MagicMock() - _session_mapping["cached-session"] = "cached-session-id" - reconnect_client = MagicMock() reconnect_client.get_session.return_value = {"status": "READY"} mock_client_class.return_value = reconnect_client interpreter = AgentCoreCodeInterpreter(region="us-west-2", boto_session=mock_session) + _session_mapping[interpreter._scoped_key("cached-session")] = "cached-session-id" + session_name, error = interpreter._ensure_session("cached-session") assert session_name == "cached-session" assert error is None mock_client_class.assert_called_once_with(region="us-west-2", session=mock_session) + + +def test_partition_key_defaults_to_shared_default(): + """When no partition_key is given, the partition is the shared 'default' (single-caller).""" + with patch("strands_tools.code_interpreter.agent_core_code_interpreter.resolve_region") as mock_resolve: + mock_resolve.return_value = "us-west-2" + + interpreter = AgentCoreCodeInterpreter(region="us-west-2", session_name="my-session") + + assert interpreter.partition_key == "default" + assert interpreter._scoped_key("my-session") == "default\x00my-session" + + +def test_session_cache_scoped_per_partition_keys(): + """Instances in different partitions produce different module-cache keys for the same name.""" + with patch("strands_tools.code_interpreter.agent_core_code_interpreter.resolve_region") as mock_resolve: + mock_resolve.return_value = "us-west-2" + + interpreter_a = AgentCoreCodeInterpreter(region="us-west-2", partition_key="tenant-a") + interpreter_b = AgentCoreCodeInterpreter(region="us-west-2", partition_key="tenant-b") + + # Same user-facing session name resolves to different module-cache keys. + assert interpreter_a._scoped_key("shared-name") != interpreter_b._scoped_key("shared-name") + + +def test_partition_key_independent_of_credential_rotation(): + """The cache key depends only on partition_key, not on AWS credentials. + + A service runs under one IAM identity but serves many callers; two instances given the + same partition_key must produce the same scoped key (so reconnection works) regardless of + which boto session or credentials back them. The cache key is no longer derived from AWS + credentials, so rotating credentials cannot cause a reconnect miss. + """ + with patch("strands_tools.code_interpreter.agent_core_code_interpreter.resolve_region") as mock_resolve: + mock_resolve.return_value = "us-west-2" + + interpreter_before = AgentCoreCodeInterpreter( + region="us-west-2", partition_key="user-1", boto_session=MagicMock() + ) + interpreter_after = AgentCoreCodeInterpreter( + region="us-west-2", partition_key="user-1", boto_session=MagicMock() + ) + + assert interpreter_before._scoped_key("shared-name") == interpreter_after._scoped_key("shared-name") + + +@patch("strands_tools.code_interpreter.agent_core_code_interpreter.BedrockAgentCoreCodeInterpreterClient") +def test_session_not_shared_across_partitions(mock_client_class): + """A session cached in one partition is not reachable from an instance in another partition. + + This is the multi-tenant case the maintainer review called out: one IAM identity serves + many callers, so isolation must be keyed on a per-caller partition rather than on AWS + credentials. Instance A initializes a session under partition "tenant-a"; instance B in + partition "tenant-b" must not reconnect to A's session by reusing the same session name. + """ + with patch("strands_tools.code_interpreter.agent_core_code_interpreter.resolve_region") as mock_resolve: + mock_resolve.return_value = "us-west-2" + + # Instance A creates a session named "shared-name" in partition tenant-a. + client_a = MagicMock() + client_a.session_id = "aws-session-a" + mock_client_class.return_value = client_a + + interpreter_a = AgentCoreCodeInterpreter(region="us-west-2", partition_key="tenant-a") + result = interpreter_a.init_session( + InitSessionAction(type="initSession", description="A", session_name="shared-name") + ) + assert result["status"] == "success" + assert _session_mapping[interpreter_a._scoped_key("shared-name")] == "aws-session-a" + + # Instance B in a different partition tries to reach "shared-name". + # It must NOT reconnect to A's session; with auto_create it creates its own instead. + client_b = MagicMock() + client_b.session_id = "aws-session-b" + client_b.get_session.return_value = {"status": "READY"} + mock_client_class.return_value = client_b + + interpreter_b = AgentCoreCodeInterpreter(region="us-west-2", partition_key="tenant-b", auto_create=True) + + # B does not see A's cached entry under its own partition scope. + assert interpreter_b._scoped_key("shared-name") not in _session_mapping + + session_name, error = interpreter_b._ensure_session("shared-name") + assert session_name == "shared-name" + assert error is None + + # B never reconnected to A's AWS session id. + client_b.get_session.assert_not_called() + # A's cached mapping is unchanged by B's activity. + assert _session_mapping[interpreter_a._scoped_key("shared-name")] == "aws-session-a" + + +@patch("strands_tools.code_interpreter.agent_core_code_interpreter.BedrockAgentCoreCodeInterpreterClient") +def test_same_partition_reconnects_across_instances(mock_client_class): + """Two instances sharing a partition_key reconnect to the same cached AWS session.""" + with patch("strands_tools.code_interpreter.agent_core_code_interpreter.resolve_region") as mock_resolve: + mock_resolve.return_value = "us-west-2" + + client_a = MagicMock() + client_a.session_id = "aws-session-shared" + mock_client_class.return_value = client_a + + interpreter_a = AgentCoreCodeInterpreter(region="us-west-2", partition_key="user-1") + interpreter_a.init_session(InitSessionAction(type="initSession", description="A", session_name="shared-name")) + + # A second instance in the same partition reconnects to the existing AWS session. + client_b = MagicMock() + client_b.get_session.return_value = {"status": "READY"} + mock_client_class.return_value = client_b + + interpreter_b = AgentCoreCodeInterpreter(region="us-west-2", partition_key="user-1") + session_name, error = interpreter_b._ensure_session("shared-name") + + assert error is None + client_b.get_session.assert_called_once_with( + interpreter_id=interpreter_b.identifier, session_id="aws-session-shared" + )