Skip to content
Draft
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
72 changes: 58 additions & 14 deletions src/strands_tools/code_interpreter/agent_core_code_interpreter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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 "
Expand All @@ -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(
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand Down
149 changes: 133 additions & 16 deletions tests/code_interpreter/test_agent_core_code_interpreter.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,16 +133,16 @@ 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"}
mock_client_class.return_value = reconnect_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"
Expand All @@ -164,16 +164,16 @@ 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"}
mock_client_class.return_value = reconnect_client

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"}]}

Expand All @@ -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()

Expand All @@ -195,16 +195,16 @@ 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"
mock_client_class.return_value = reconnect_client

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"}]}

Expand All @@ -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()


Expand Down Expand Up @@ -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)
Comment thread
yonib05 marked this conversation as resolved.
assert _session_mapping.get(interpreter._scoped_key("my-session")) == "test-session-id-123"


@patch("strands_tools.code_interpreter.agent_core_code_interpreter.BedrockAgentCoreCodeInterpreterClient")
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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"
)
Loading