From f0a0ac570372b9a73411927d4931169b3fb3b3a6 Mon Sep 17 00:00:00 2001 From: Phil Merrell Date: Sun, 27 Sep 2026 21:51:07 -0600 Subject: [PATCH 1/4] perf(backend): A/B the first-turn agent build, default off `agent_build.session_mgr` is 616-738ms of a 1.0-1.2s first-turn build on dev. The AgentCore SDK's session manager builds a MemoryClient (fresh boto3 session, two clients), then a second fresh session and two clients that replace the first pair, on every construction. And `create_agent` is synchronous on the event loop, so a Nova title reply that lands mid-build sits unprocessed until the build ends. Both fixes run behind one per-session experiment flag, AGENT_BUILD_EXPERIMENT (unset = control everywhere), so they ship only if a dev A/B says they help: - shared_clients: one process-wide boto3 session whose client() returns one client per configuration, handed to the SDK (and to its MemoryClient via the SDK module's name, the only seam). - shared_clients_off_loop: that, plus the build under asyncio.to_thread, one build at a time, with the route emitting a title that lands mid-build. The frozen loop used to serialize everything for free, so this also locks ExternalMCPIntegration's maps, creates its singleton under a lock (two instances would let a needs_approval tool run unapproved), and pins the external MCP clients a build is handed until its agent registers as their consumer. turn_prelude gains buildArm, processBuilds and two sub-stages (session_mgr_clients, strands_agent); scripts/experiment_agent_build_arms.py drives interleaved first turns per arm and joins them to it. Co-Authored-By: Claude Opus 5.5 --- .../scripts/experiment_agent_build_arms.py | 296 ++++++++++++++++++ backend/src/agents/main_agent/base_agent.py | 26 +- .../integrations/external_mcp_client.py | 142 ++++++--- .../main_agent/session/session_factory.py | 126 +++++++- .../session/turn_based_session_manager.py | 16 + backend/src/apis/inference_api/chat/routes.py | 98 +++--- .../src/apis/inference_api/chat/service.py | 66 +++- .../apis/inference_api/chat/turn_timing.py | 4 +- backend/src/apis/shared/feature_flags.py | 87 +++++ .../integrations/test_external_mcp_client.py | 81 +++++ .../test_session_factory_shared_clients.py | 135 ++++++++ .../test_base_agent_external_registration.py | 37 +++ .../apis/inference_api/test_chat_service.py | 149 +++++++++ .../inference_api/test_preparing_phase.py | 75 +++++ .../shared/test_agent_build_experiment_arm.py | 54 ++++ docs/specs/turn-latency-preamble.md | 69 +++- 16 files changed, 1361 insertions(+), 100 deletions(-) create mode 100644 backend/scripts/experiment_agent_build_arms.py create mode 100644 backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py create mode 100644 backend/tests/shared/test_agent_build_experiment_arm.py diff --git a/backend/scripts/experiment_agent_build_arms.py b/backend/scripts/experiment_agent_build_arms.py new file mode 100644 index 000000000..a41fdfae2 --- /dev/null +++ b/backend/scripts/experiment_agent_build_arms.py @@ -0,0 +1,296 @@ +"""Agent-build A/B: does sharing Memory clients / building off the loop help a first turn? + +Drives real first turns through the deployed AgentCore Runtime and compares the +three arms of ``agent_build_experiment_arm`` (``apis/shared/feature_flags.py``): + + control today's build + shared_clients AgentCore Memory session managers share boto3 clients + shared_clients_off_loop that, plus the synchronous build runs in a worker thread + +**Precondition:** the Runtime must have ``AGENT_BUILD_EXPERIMENT=ab``. Arms are +assigned server-side by hashing the session id, so this script generates session +ids until each lands in the arm it wants, then runs the arms interleaved (one of +each per round) so network drift over the run hits every arm alike. Every +conversation runs in its own Runtime process, which is exactly the first-turn, +cold-process build under test (``processBuilds`` = 1 on ``turn_prelude``). + +Two sources per turn, joined on session id: + +- **Client side** (this script, via ``on_event``): when ``preparing``, + ``prepared``, ``session_title``, the first ``content_block_delta`` and + ``done`` arrived, measured from the request. +- **Server side** (``turn_prelude`` in the Runtime log group): the build's + sub-stages, ``buildArm`` and ``processBuilds``. + +What it establishes: per-arm medians for the build and its sub-stages, time to +first token, and whether the title beats ``prepared``. What it cannot: fleet +magnitude, or behaviour under concurrent load. Report which claim you make. + +Each turn is a real conversation owned by ``--user-id``: it draws on that user's +quota and appears in their sidebar. Use ``--cleanup`` to soft-delete the +experiment's sessions afterwards. + +Usage (an authenticated dev-ai profile and an active headless grant for +--user-id; see apis/shared/harness/grants.py): + + cd backend + AWS_PROFILE=dev-ai uv run python scripts/experiment_agent_build_arms.py \\ + --user-id --per-arm 15 +""" + +from __future__ import annotations + +import argparse +import asyncio +import json +import logging +import os +import re +import statistics +import sys +import time +import uuid +from dataclasses import asdict, dataclass, field +from typing import Any, Dict, List, Optional + +logging.basicConfig(level=logging.INFO, format="%(levelname)s %(name)s: %(message)s") +logging.getLogger("httpx").setLevel(logging.WARNING) +logger = logging.getLogger("experiment") + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) +sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "src")) + +ARMS = ("control", "shared_clients", "shared_clients_off_loop") + +# Frames whose first arrival is recorded, keyed by how the table names them. +_TIMED = ("preparing", "prepared", "session_title", "first_token", "done") + + +@dataclass +class Turn: + arm: str + session_id: str + ok: bool = False + error: Optional[str] = None + client_ms: Dict[str, int] = field(default_factory=dict) + prelude: Dict[str, Any] = field(default_factory=dict) + + +def session_for(arm: str) -> str: + """A fresh session id that the server's ``ab`` hash puts in ``arm``.""" + from apis.shared.feature_flags import agent_build_experiment_arm + + previous = os.environ.get("AGENT_BUILD_EXPERIMENT") + os.environ["AGENT_BUILD_EXPERIMENT"] = "ab" + try: + while True: + candidate = str(uuid.uuid4()) + if agent_build_experiment_arm(candidate) == arm: + return candidate + finally: + if previous is None: + os.environ.pop("AGENT_BUILD_EXPERIMENT", None) + else: + os.environ["AGENT_BUILD_EXPERIMENT"] = previous + + +async def run_turn(*, arm: str, user_id: str, prompt: str, auth: Any, model_id: Optional[str]) -> Turn: + from apis.shared.harness import run_agent_headless + + turn = Turn(arm=arm, session_id=session_for(arm)) + started = time.monotonic() + + def stamp(key: str) -> None: + turn.client_ms.setdefault(key, int((time.monotonic() - started) * 1000)) + + async def on_event(name: str, data: Dict[str, Any]) -> None: + if name == "agent_status" and data.get("phase") in ("preparing", "prepared"): + stamp(data["phase"]) + elif name == "session_title": + stamp("session_title") + elif name == "content_block_delta": + stamp("first_token") + elif name == "done": + stamp("done") + + try: + run = await run_agent_headless( + user_id=user_id, + prompt=prompt, + auth=auth, + session_id=turn.session_id, + model_id=model_id, + trigger="experiment", + on_event=on_event, + ) + turn.ok = run.status == "completed" + turn.error = None if turn.ok else (run.error or run.status) + except Exception as exc: # noqa: BLE001 - one bad turn must not end the run + turn.error = str(exc)[:200] + logger.info( + " %-24s %s %s", arm, turn.session_id[:8], + "ok" if turn.ok else f"FAILED: {turn.error}", + ) + return turn + + +def attach_preludes(turns: List[Turn], log_group: str, region: str, since: float) -> None: + """Join each turn to its ``turn_prelude`` line by session id.""" + import boto3 + + logs = boto3.client("logs", region_name=region) + by_session = {t.session_id: t for t in turns} + query = logs.start_query( + logGroupName=log_group, + startTime=int(since) - 60, + endTime=int(time.time()) + 60, + queryString="fields body | filter body like /turn_prelude/ | limit 10000", + )["queryId"] + while True: + result = logs.get_query_results(queryId=query) + if result["status"] in ("Complete", "Failed", "Cancelled", "Timeout"): + break + time.sleep(2) + for row in result.get("results", []): + body = next((f["value"] for f in row if f["field"] == "body"), "") + match = re.search(r"turn_prelude (\{.*\})", body) + if not match: + continue + prelude = json.loads(match.group(1)) + turn = by_session.get(prelude.get("sessionId")) + if turn is not None: + turn.prelude = prelude + + +def _median(values: List[float]) -> str: + return f"{statistics.median(values):.0f}" if values else "-" + + +def _p75(values: List[float]) -> str: + if len(values) < 4: + return "-" + return f"{statistics.quantiles(values, n=4)[2]:.0f}" + + +def report(turns: List[Turn]) -> None: + def stage(t: Turn, name: str) -> Optional[float]: + return t.prelude.get("stages", {}).get(name) + + rows = [ + ("agent_build (group)", lambda t: t.prelude.get("groups", {}).get("agent_build")), + (" session_mgr_clients", lambda t: stage(t, "agent_build.session_mgr_clients")), + (" session_mgr (network)", lambda t: stage(t, "agent_build.session_mgr")), + (" strands_agent", lambda t: stage(t, "agent_build.strands_agent")), + (" finalize (restore)", lambda t: stage(t, "agent_build.finalize")), + ("prelude total", lambda t: t.prelude.get("totalMs")), + ("client: prepared", lambda t: t.client_ms.get("prepared")), + ("client: first token", lambda t: t.client_ms.get("first_token")), + ("client: session_title", lambda t: t.client_ms.get("session_title")), + ( + "title minus prepared", + lambda t: (t.client_ms["session_title"] - t.client_ms["prepared"]) + if "session_title" in t.client_ms and "prepared" in t.client_ms else None, + ), + ] + + print() + header = f"{'metric (ms)':26s}" + "".join(f"{arm:>28s}" for arm in ARMS) + print(header) + print(f"{'':26s}" + "".join(f"{'median / p75':>28s}" for _ in ARMS)) + for label, fn in rows: + cells = [] + for arm in ARMS: + values = [v for t in turns if t.arm == arm and t.ok and (v := fn(t)) is not None] + cells.append(f"{_median(values)} / {_p75(values)} (n={len(values)})") + print(f"{label:26s}" + "".join(f"{c:>28s}" for c in cells)) + + print() + for arm in ARMS: + mine = [t for t in turns if t.arm == arm] + matched = [t for t in mine if t.prelude] + mislabeled = [t for t in matched if t.prelude.get("buildArm") != arm] + warm = [t for t in matched if t.prelude.get("processBuilds") not in (1, None)] + early = [ + t for t in mine + if "session_title" in t.client_ms and "prepared" in t.client_ms + and t.client_ms["session_title"] < t.client_ms["prepared"] + ] + print( + f"{arm:24s} turns={len(mine)} ok={sum(t.ok for t in mine)} " + f"prelude-matched={len(matched)} arm-mismatch={len(mislabeled)} " + f"not-first-build={len(warm)} title-before-prepared={len(early)}" + ) + if any(t.prelude and t.prelude.get("buildArm") == "control" for t in turns if t.arm != "control"): + print("\n⚠️ Non-control turns reported buildArm=control: is AGENT_BUILD_EXPERIMENT=ab set on the Runtime?") + + +def cleanup(turns: List[Turn], user_id: str) -> None: + """Soft-delete the experiment's conversations from the user's sidebar.""" + import boto3 + + table = boto3.resource("dynamodb", region_name=os.environ["AWS_REGION"]).Table( + os.environ["DYNAMODB_SESSIONS_METADATA_TABLE_NAME"] + ) + removed = 0 + for turn in turns: + try: + table.update_item( + Key={"PK": f"USER#{user_id}", "SK": f"S#{turn.session_id}"}, + UpdateExpression="SET #s = :deleted REMOVE GSI4_PK, GSI4_SK", + ConditionExpression="attribute_exists(PK)", + ExpressionAttributeNames={"#s": "status"}, + ExpressionAttributeValues={":deleted": "deleted"}, + ) + removed += 1 + except Exception as exc: # noqa: BLE001 - best effort + logger.warning("cleanup %s: %s", turn.session_id, exc) + logger.info("soft-deleted %d/%d experiment sessions", removed, len(turns)) + + +async def main() -> None: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--user-id", required=True) + parser.add_argument("--per-arm", type=int, default=15) + parser.add_argument("--gap-seconds", type=float, default=3.0) + parser.add_argument("--model-id", default=None) + parser.add_argument("--prefix", default="dev-boisestateai-v2") + parser.add_argument("--region", default="us-west-2") + parser.add_argument("--out", default=None, help="write raw per-turn JSON here") + parser.add_argument("--cleanup", action="store_true", help="soft-delete the sessions afterwards") + args = parser.parse_args() + + from spike_headless_run import resolve_environment + from apis.shared.harness.auth import CognitoRefreshBearerAuth + + env = resolve_environment(args.prefix, args.region) + runtime_id = env["runtime_arn"].rsplit("/", 1)[1] + log_group = f"/aws/bedrock-agentcore/runtimes/{runtime_id}-DEFAULT" + auth = CognitoRefreshBearerAuth() + + since = time.time() + turns: List[Turn] = [] + for round_index in range(args.per_arm): + order = list(ARMS) + # Rotate the order each round so no arm always runs first. + order = order[round_index % len(order):] + order[: round_index % len(order)] + logger.info("── round %d/%d", round_index + 1, args.per_arm) + for arm in order: + prompt = f"In one short sentence, name a river in country number {round_index + 1} of Africa, alphabetically." + turns.append(await run_turn(arm=arm, user_id=args.user_id, prompt=prompt, auth=auth, model_id=args.model_id)) + await asyncio.sleep(args.gap_seconds) + + logger.info("waiting 45s for turn_prelude lines to reach CloudWatch…") + await asyncio.sleep(45) + attach_preludes(turns, log_group, args.region, since) + report(turns) + + if args.out: + with open(args.out, "w") as fh: + json.dump([asdict(t) for t in turns], fh, indent=2) + logger.info("raw results: %s", args.out) + if args.cleanup: + cleanup(turns, args.user_id) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/backend/src/agents/main_agent/base_agent.py b/backend/src/agents/main_agent/base_agent.py index 5e59d1abf..aabe7f9f5 100644 --- a/backend/src/agents/main_agent/base_agent.py +++ b/backend/src/agents/main_agent/base_agent.py @@ -201,10 +201,29 @@ def __init__( # Initialize streaming coordinator self.stream_coordinator = StreamCoordinator() - # Create the agent (subclass-specific) - self._create_agent() + # Create the agent (subclass-specific). External MCP clients handed to + # it are pinned for the build (see `load_external_tools`'s + # `consumer_pin`); once the agent has registered as their consumer, or + # failed to, the pin goes. + self._mcp_build_pin = object() + self._mcp_pinned_clients: List[Any] = [] + try: + self._create_agent() + finally: + self._release_mcp_build_pins() mark_stage("finalize") + def _release_mcp_build_pins(self) -> None: + for client in self._mcp_pinned_clients: + try: + client.remove_consumer(self._mcp_build_pin) + except Exception: # noqa: BLE001 - a failed stop must not fail the build + logger.warning("Releasing an MCP build pin failed", exc_info=True) + self._mcp_pinned_clients = [] + # Scoped to the constructor's build: a later `_create_agent` (the + # `stream_async` fallback) takes no pin, so it can never leak one. + self._mcp_build_pin = None + @abstractmethod def _create_agent(self) -> None: """Create the specific agent type. Subclasses must implement.""" @@ -550,6 +569,7 @@ async def _load_with_context(): external_mcp_tool_ids, user_id=self.user_id, auth_token=self.auth_token, + consumer_pin=getattr(self, "_mcp_build_pin", None), ) # Probe with ``get_running_loop`` rather than ``get_event_loop``: @@ -572,6 +592,8 @@ async def _load_with_context(): future = executor.submit(asyncio.run, _load_with_context()) external_clients = future.result() + if getattr(self, "_mcp_build_pin", None) is not None: + self._mcp_pinned_clients = list(external_clients) for client in external_clients: if client not in local_tools: local_tools.append(client) diff --git a/backend/src/agents/main_agent/integrations/external_mcp_client.py b/backend/src/agents/main_agent/integrations/external_mcp_client.py index bdce9ff21..84ef09be2 100644 --- a/backend/src/agents/main_agent/integrations/external_mcp_client.py +++ b/backend/src/agents/main_agent/integrations/external_mcp_client.py @@ -15,6 +15,7 @@ import logging import re +import threading from typing import Any, Callable, Iterator, Optional, List, Set from urllib.parse import urlparse @@ -371,6 +372,13 @@ def __init__(self): # tool silently. Keyed by user because this integration is a # process-wide singleton shared across concurrent sessions. self._pending_consents: dict[str, dict[str, str]] = {} + # Guards every dict above. Agent builds run in a worker thread + # (``agent_build_off_loop_enabled``) while the event loop reads these + # maps (``get_client``, ``take_pending_consents``); iterating one while + # a build inserts raises. Held only around dict access, never across + # an await: a lock held through an MCP pre-flight would stall the loop + # for the length of a network round trip. + self._lock = threading.RLock() def take_pending_consents(self, user_id: str) -> dict[str, str]: """Pop and return {provider_id: authorization_url} for `user_id`. @@ -381,7 +389,8 @@ def take_pending_consents(self, user_id: str) -> dict[str, str]: agent-cache hit no loading happens, nothing is recorded, and nothing is emitted — correct, because the prompt already went out once. """ - return self._pending_consents.pop(user_id, {}) + with self._lock: + return self._pending_consents.pop(user_id, {}) async def _recover_oauth_preflight( self, @@ -485,7 +494,8 @@ async def _recover_oauth_preflight( ) return False - self._pending_consents.setdefault(user_id, {})[provider_id] = authorization_url + with self._lock: + self._pending_consents.setdefault(user_id, {})[provider_id] = authorization_url logger.info( f"External MCP tool {tool_id} needs {provider_id} consent; " "surfacing oauth_required instead of dropping it silently" @@ -510,6 +520,7 @@ async def load_external_tools( enabled_tool_ids: List[str], user_id: Optional[str] = None, auth_token: Optional[str] = None, + consumer_pin: Any = None, ) -> List[MCPClient]: """ Load external MCP clients for enabled tools. @@ -524,6 +535,15 @@ async def load_external_tools( enabled_tool_ids: List of enabled tool IDs user_id: User ID (required for OAuth-gated and OIDC-forwarded tools) auth_token: Raw OIDC token for forwarding + consumer_pin: When given, registered as a consumer of every + returned client, so no client can be stopped between this + hand-out and the new agent registering itself as a consumer. + A cached client is shared across agents, and Strands stops one + the moment its last consumer goes (``remove_consumer``). With + the build in a worker thread, a turn ending on the event loop + could drop that last consumer mid-build and leave the new agent + holding tools from a stopped client. The caller removes the pin + once the agent exists (``BaseAgent.__init__``). Returns: List of MCPClient instances to add to the agent's tools @@ -569,21 +589,22 @@ async def load_external_tools( to_iso(tool.updated_at) if tool.updated_at else "" ) - if ( - cache_key in self.clients - and self._client_versions.get(cache_key) == tool_version - ): - clients.append(self.clients[cache_key]) - continue + with self._lock: + cached = self.clients.get(cache_key) + if cached is not None and self._client_versions.get(cache_key) == tool_version: + if consumer_pin is not None: + cached.add_consumer(consumer_pin) + clients.append(cached) + continue - # Stale entry — admin edited this tool since the client - # was built. Drop it so the block below creates a fresh - # client with the current config. - if cache_key in self.clients: - stale = self.clients.pop(cache_key) - self._client_versions.pop(cache_key, None) - self._provider_for_client_id.pop(id(stale), None) - self._approval_names_for_client_id.pop(id(stale), None) + # Stale entry — admin edited this tool since the client + # was built. Drop it so the block below creates a fresh + # client with the current config. + if cached is not None: + self.clients.pop(cache_key, None) + self._client_versions.pop(cache_key, None) + self._provider_for_client_id.pop(id(cached), None) + self._approval_names_for_client_id.pop(id(cached), None) static_token: Optional[str] = None token_provider: Optional[Callable[[], Optional[str]]] = None @@ -698,13 +719,16 @@ async def _exchange( if not recovered: continue - self.clients[cache_key] = client - self._client_versions[cache_key] = tool_version - if provider_id: - self._provider_for_client_id[id(client)] = provider_id approval_names = tool.mcp_config.approval_required_names() - if approval_names: - self._approval_names_for_client_id[id(client)] = approval_names + with self._lock: + self.clients[cache_key] = client + self._client_versions[cache_key] = tool_version + if provider_id: + self._provider_for_client_id[id(client)] = provider_id + if approval_names: + self._approval_names_for_client_id[id(client)] = approval_names + if consumer_pin is not None: + client.add_consumer(consumer_pin) clients.append(client) auth_label = ( " (with OIDC forwarding)" if forward_auth and static_token @@ -739,19 +763,22 @@ def get_client(self, tool_id: str, user_id: Optional[str] = None) -> Optional[MC base = base_tool_id(tool_id) # Exact keys win — a whole-server binding has no "|allow:" suffix. exact_keys = [f"{user_id}:{base}", base] if user_id else [base] - for key in exact_keys: - if key in self.clients: - return self.clients[key] - # Subset-scoped fallback: cache key is "|allow:". - for key in exact_keys: - prefix = f"{key}|allow:" - for cache_key, client in self.clients.items(): - if cache_key.startswith(prefix): - return client + with self._lock: + for key in exact_keys: + if key in self.clients: + return self.clients[key] + # Subset-scoped fallback: cache key is "|allow:". + for key in exact_keys: + prefix = f"{key}|allow:" + for cache_key, client in self.clients.items(): + if cache_key.startswith(prefix): + return client return None def add_to_tool_list(self, tools: List[Any]) -> List[Any]: - for client in self.clients.values(): + with self._lock: + cached = list(self.clients.values()) + for client in cached: if client not in tools: tools.append(client) return tools @@ -764,15 +791,16 @@ def clear_user_clients(self, user_id: str) -> None: agent build creates fresh clients (and the token cache miss forces a new consent flow). """ - keys_to_remove = [ - key for key in self.clients.keys() - if key.startswith(f"{user_id}:") - ] - for key in keys_to_remove: - client = self.clients.pop(key) - self._client_versions.pop(key, None) - self._provider_for_client_id.pop(id(client), None) - self._approval_names_for_client_id.pop(id(client), None) + with self._lock: + keys_to_remove = [ + key for key in self.clients.keys() + if key.startswith(f"{user_id}:") + ] + for key in keys_to_remove: + client = self.clients.pop(key) + self._client_versions.pop(key, None) + self._provider_for_client_id.pop(id(client), None) + self._approval_names_for_client_id.pop(id(client), None) if keys_to_remove: logger.info(f"Cleared {len(keys_to_remove)} cached MCP clients for user {user_id}") @@ -786,26 +814,36 @@ def clear_tool_clients(self, tool_id: str) -> None: config. Without this, clients cached at process start continue to point at the old URL for the lifetime of the process. """ - keys_to_remove = [ - key for key in self.clients.keys() - if key == tool_id or key.endswith(f":{tool_id}") - ] - for key in keys_to_remove: - client = self.clients.pop(key) - self._client_versions.pop(key, None) - self._provider_for_client_id.pop(id(client), None) - self._approval_names_for_client_id.pop(id(client), None) + with self._lock: + keys_to_remove = [ + key for key in self.clients.keys() + if key == tool_id or key.endswith(f":{tool_id}") + ] + for key in keys_to_remove: + client = self.clients.pop(key) + self._client_versions.pop(key, None) + self._provider_for_client_id.pop(id(client), None) + self._approval_names_for_client_id.pop(id(client), None) if keys_to_remove: logger.info(f"Cleared {len(keys_to_remove)} cached MCP clients for tool {tool_id}") _external_mcp_integration: Optional[ExternalMCPIntegration] = None +_external_mcp_integration_lock = threading.Lock() def get_external_mcp_integration() -> ExternalMCPIntegration: - """Get or create the global ExternalMCPIntegration instance.""" + """Get or create the global ExternalMCPIntegration instance. + + Locked because two first builds can now run concurrently in worker + threads. Two instances would split the maps: the approval hook built + against one would read an empty ``approval_names_for_client`` from the + other, and a ``needs_approval`` tool would run without asking. + """ global _external_mcp_integration if _external_mcp_integration is None: - _external_mcp_integration = ExternalMCPIntegration() + with _external_mcp_integration_lock: + if _external_mcp_integration is None: + _external_mcp_integration = ExternalMCPIntegration() return _external_mcp_integration diff --git a/backend/src/agents/main_agent/session/session_factory.py b/backend/src/agents/main_agent/session/session_factory.py index 6bb12856e..1e69f9140 100644 --- a/backend/src/agents/main_agent/session/session_factory.py +++ b/backend/src/agents/main_agent/session/session_factory.py @@ -1,8 +1,10 @@ """ Session manager factory for creating AgentCore Memory session managers """ +import contextvars import os import logging +import threading from typing import Optional, Any, Dict, Tuple from functools import lru_cache @@ -68,11 +70,112 @@ def session_async_persistence_enabled() -> bool: try: from bedrock_agentcore.memory.integrations.strands.config import AgentCoreMemoryConfig, RetrievalConfig from bedrock_agentcore.memory import MemoryClient + from bedrock_agentcore.memory.integrations.strands import session_manager as _sdk_session_manager AGENTCORE_MEMORY_AVAILABLE = True except ImportError: AGENTCORE_MEMORY_AVAILABLE = False +# --------------------------------------------------------------------------- +# One set of AgentCore Memory clients per process +# --------------------------------------------------------------------------- +# +# The SDK's session manager builds its clients from scratch on every +# construction, twice (see ``memory_shared_clients_enabled``). Everything +# below exists to hand it one process-wide session instead. + +# Set by the factory for the duration of one session manager's construction. +# The SDK builds its ``MemoryClient`` with no session and no session id, so this +# is how the per-session decision reaches ``_SharedSessionMemoryClient``. +_use_shared_memory_clients: contextvars.ContextVar[bool] = contextvars.ContextVar( + "use_shared_memory_clients", default=False +) + +# Sized for concurrent sessions sharing one pool. Async persistence writes each +# message through ``asyncio.to_thread``, so a busy container can have dozens of +# ``CreateEvent`` calls in flight on these clients at once, where each session +# used to have a pool of its own. Past the pool size urllib3 still serves the +# request, but discards the extra connection afterwards (and warns). +_SHARED_MEMORY_MAX_POOL_CONNECTIONS = 50 + + +def _client_cache_key(service_name: str, region_name: Optional[str], config: Any) -> Tuple[str, Optional[str], str]: + # ``Config`` is not hashable; its user-provided options are what makes two + # configs different (the SDK and ``MemoryClient`` differ only in user agent). + options = getattr(config, "_user_provided_options", None) or {} + return service_name, region_name, repr(sorted(options.items())) + + +if AGENTCORE_MEMORY_AVAILABLE: + import boto3 + from botocore.config import Config as _BotocoreConfig + + class _ClientReusingSession(boto3.Session): + """A boto3 session whose ``client()`` returns one client per configuration. + + boto3 clients are thread-safe; building them is not, and building one + is the expensive part, so creation happens once, under a lock. Calls + carrying anything beyond region and config (explicit credentials, an + endpoint override) are not ours to share and pass straight through. + """ + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + self._shared_clients: Dict[Tuple[str, Optional[str], str], Any] = {} + self._shared_clients_lock = threading.Lock() + + def client(self, service_name: str, region_name: Optional[str] = None, config: Any = None, **kwargs: Any) -> Any: # type: ignore[override] + if kwargs: + return super().client(service_name, region_name=region_name, config=config, **kwargs) + key = _client_cache_key(service_name, region_name, config) + with self._shared_clients_lock: + client = self._shared_clients.get(key) + if client is None: + pooled = _BotocoreConfig(max_pool_connections=_SHARED_MEMORY_MAX_POOL_CONNECTIONS) + merged = pooled.merge(config) if config is not None else pooled + client = super().client(service_name, region_name=region_name, config=merged) + self._shared_clients[key] = client + return client + + _shared_memory_session: Optional[_ClientReusingSession] = None + _shared_memory_session_lock = threading.Lock() + + def shared_memory_boto_session() -> _ClientReusingSession: + global _shared_memory_session + if _shared_memory_session is None: + with _shared_memory_session_lock: + if _shared_memory_session is None: + _shared_memory_session = _ClientReusingSession() + return _shared_memory_session + + class _SharedSessionMemoryClient(MemoryClient): + """``MemoryClient`` built from the shared session when the flag is on. + + The SDK session manager constructs ``MemoryClient(region_name=...)`` + with no session, so it cannot be handed one; its clients are then + replaced a few lines later, so they were pure cost. Rebinding the name + in the SDK module is the only seam. Pinned by + ``test_session_factory_shared_clients.py``, which fails if an SDK + upgrade stops going through it. + """ + + def __init__( + self, + region_name: Optional[str] = None, + integration_source: Optional[str] = None, + boto3_session: Any = None, + ) -> None: + if boto3_session is None and _use_shared_memory_clients.get(): + boto3_session = shared_memory_boto_session() + super().__init__( + region_name=region_name, + integration_source=integration_source, + boto3_session=boto3_session, + ) + + _sdk_session_manager.MemoryClient = _SharedSessionMemoryClient + + @lru_cache(maxsize=1) def _discover_strategy_ids(memory_id: str, region: str) -> Tuple[Optional[str], Optional[str], Optional[str]]: """ @@ -282,13 +385,21 @@ def _create_cloud_session_manager( compaction_config.token_threshold = compaction_threshold # Create session manager with compaction built-in - session_manager = TurnBasedSessionManager( - agentcore_memory_config=agentcore_memory_config, - region_name=aws_region, - compaction_config=compaction_config if compaction_config.enabled else None, - user_id=user_id, - summarization_strategy_id=summary_id, - ) + from apis.shared.feature_flags import memory_shared_clients_enabled + + shared_clients = memory_shared_clients_enabled(session_id) + token = _use_shared_memory_clients.set(shared_clients) + try: + session_manager = TurnBasedSessionManager( + agentcore_memory_config=agentcore_memory_config, + region_name=aws_region, + compaction_config=compaction_config if compaction_config.enabled else None, + user_id=user_id, + summarization_strategy_id=summary_id, + boto_session=shared_memory_boto_session() if shared_clients else None, + ) + finally: + _use_shared_memory_clients.reset(token) logger.info("✅ AgentCore Memory initialized") logger.info(" • Storage: AWS-managed DynamoDB") @@ -299,6 +410,7 @@ def _create_cloud_session_manager( else: logger.info(" • Compaction: Disabled") logger.info(" • Persistence: %s", "Async (off the event loop)" if async_persistence else "Sync (blocking)") + logger.info(" • Clients: %s", "Shared (process-wide)" if shared_clients else "Per session manager") return session_manager diff --git a/backend/src/agents/main_agent/session/turn_based_session_manager.py b/backend/src/agents/main_agent/session/turn_based_session_manager.py index 450437542..eac12a28c 100644 --- a/backend/src/agents/main_agent/session/turn_based_session_manager.py +++ b/backend/src/agents/main_agent/session/turn_based_session_manager.py @@ -38,6 +38,7 @@ from typing import Optional, Dict, Any, List, Tuple, TYPE_CHECKING from agents.main_agent.config.constants import Defaults, EnvVars +from apis.shared.observability.build_stages import mark_stage from bedrock_agentcore.memory.integrations.strands.session_manager import AgentCoreMemorySessionManager from bedrock_agentcore.memory.integrations.strands.config import AgentCoreMemoryConfig @@ -419,6 +420,17 @@ def retrieve_for_namespace(namespace: str, cfg: Any) -> List[str]: logger.error("Failed to retrieve customer context: %s", e) return None + def read_session(self, session_id: str, **kwargs: Any) -> Any: + """The SDK's session read, preceded by a build-stage mark. + + The SDK constructor builds its boto3 clients and then calls this, so the + mark splits ``agent_build.session_mgr`` into client setup + (``session_mgr_clients``) and the network calls that follow. A no-op + outside a turn's build. + """ + mark_stage("session_mgr_clients") + return super().read_session(session_id, **kwargs) + def initialize(self, agent: "Agent", **kwargs: Any) -> None: """ Initialize agent with two-feature compaction. @@ -428,6 +440,10 @@ def initialize(self, agent: "Agent", **kwargs: Any) -> None: 2. Let the SDK restore agent state and load messages from AgentCore Memory 3. Apply compaction (checkpoint + truncation) on the loaded messages """ + # Splits `agent_build.finalize`: everything before this is Strands' + # own Agent construction (tool registration, including MCP + # `load_tools`); everything after is the session restore. + mark_stage("strands_agent") logger.info(f"TurnBasedSessionManager.initialize() called for agent_id={agent.agent_id}") # Let the SDK handle all session restore logic: diff --git a/backend/src/apis/inference_api/chat/routes.py b/backend/src/apis/inference_api/chat/routes.py index 44a414f42..25ad85a0c 100644 --- a/backend/src/apis/inference_api/chat/routes.py +++ b/backend/src/apis/inference_api/chat/routes.py @@ -27,6 +27,7 @@ ) from apis.inference_api.runtime_health import ping_payload from apis.shared.feature_flags import ( + agent_build_experiment_arm, agent_preparing_phase_enabled, agents_enabled, attachment_turn_guard_enabled, @@ -90,7 +91,7 @@ from .app_tool_dispatch import AppToolCallError, dispatch_app_tool_call from .agent_binding_policy import binds_conversation from .models import FileContent, InvocationRequest -from .service import generate_conversation_title, get_agent +from .service import generate_conversation_title, get_agent, process_build_count from .turn_timing import TurnPrelude from .system_prompt_resolver import ( append_active_prompt, @@ -3802,43 +3803,45 @@ async def _build_main_agent(): } ) + # One-shot `session_title` SSE: once the concurrent title task + # (kicked off before the quota check on first turns) finishes, + # push the title to the client so the sidebar/header rename in + # parallel with the pending response instead of at stream end. + # Never awaited, so it adds no latency. Polled from two places: + # the coordinator's live status merge (every 100ms, so a title + # that lands during the model's time-to-first-token or a long + # tool call goes out right away) and between agent events below + # (the only route while that merge is switched off). It is also + # checked while a deferred build is in flight, which only a build + # off the event loop can reach (`agent_build_off_loop_enabled`). A + # stream that outruns Nova Micro simply never emits and the SPA's + # post-close metadata refresh covers it. + title_emitted = False + + def _session_title_sse() -> Optional[str]: + nonlocal title_emitted + if title_emitted or title_task is None or not title_task.done(): + return None + title_emitted = True + try: + generated_title = title_task.result() + except Exception as title_err: # noqa: BLE001 - cancelled/failed task must not break the stream + logger.warning("Title task unavailable for SSE emit: %s", title_err) + return None + # Generation failures return the "New Conversation" + # placeholder — nothing worth pushing over the wire. + if not generated_title or generated_title == "New Conversation": + return None + payload = { + "type": "session_title", + "sessionId": input_data.session_id, + "title": generated_title, + } + return f"event: session_title\ndata: {json.dumps(payload)}\n\n" + # Create stream with optional quota warning injection async def stream_with_quota_warning() -> AsyncGenerator[str, None]: """Wrap agent stream to inject quota warning at start if needed""" - # One-shot `session_title` SSE: once the concurrent title task - # (kicked off before the quota check on first turns) finishes, - # push the title to the client so the sidebar/header rename in - # parallel with the pending response instead of at stream end. - # Never awaited, so it adds no latency. Polled from two places: - # the coordinator's live status merge (every 100ms, so a title - # that lands during the model's time-to-first-token or a long - # tool call goes out right away) and between agent events below - # (the only route while that merge is switched off). A stream - # that outruns Nova Micro simply never emits and the SPA's - # post-close metadata refresh covers it. - title_emitted = False - - def _session_title_sse() -> Optional[str]: - nonlocal title_emitted - if title_emitted or title_task is None or not title_task.done(): - return None - title_emitted = True - try: - generated_title = title_task.result() - except Exception as title_err: # noqa: BLE001 - cancelled/failed task must not break the stream - logger.warning("Title task unavailable for SSE emit: %s", title_err) - return None - # Generation failures return the "New Conversation" - # placeholder — nothing worth pushing over the wire. - if not generated_title or generated_title == "New Conversation": - return None - payload = { - "type": "session_title", - "sessionId": input_data.session_id, - "title": generated_title, - } - return f"event: session_title\ndata: {json.dumps(payload)}\n\n" - # Yield quota warning event first if applicable if quota_warning_event: yield quota_warning_event.to_sse_format() @@ -4098,7 +4101,24 @@ async def _guarded_stream() -> AsyncGenerator[str, None]: + "\n\n" ) try: - agent = await _build_main_agent() + # Raced against the title so a title that lands + # mid-build goes out mid-build. With the build on the + # event loop the title task cannot finish first, so + # this reduces to the plain await it replaced. + build = asyncio.ensure_future(_build_main_agent()) + try: + if title_task is not None and not title_task.done(): + await asyncio.wait( + {build, title_task}, + return_when=asyncio.FIRST_COMPLETED, + ) + title_sse = _session_title_sse() + if title_sse: + yield title_sse + agent = await build + finally: + if not build.done(): + build.cancel() except Exception as build_error: # The handler has already returned, so the two `except` # arms below cannot see this — a build that fails here @@ -4164,6 +4184,12 @@ async def _guarded_stream() -> AsyncGenerator[str, None]: "isResume": is_resume, "hasAssistant": bool(input_data.rag_assistant_id), "deferredBuild": deferred_build, + # The agent-build A/B (`agent_build_experiment_arm`), + # and how many builds this process has run: 1 on a + # conversation's first turn, which is always a fresh + # Runtime process. + "buildArm": agent_build_experiment_arm(input_data.session_id), + "processBuilds": process_build_count(), }, ) diff --git a/backend/src/apis/inference_api/chat/service.py b/backend/src/apis/inference_api/chat/service.py index 8d90145e4..8f1da3d24 100644 --- a/backend/src/apis/inference_api/chat/service.py +++ b/backend/src/apis/inference_api/chat/service.py @@ -5,6 +5,7 @@ import asyncio import json +import weakref import logging import hashlib import os @@ -295,6 +296,59 @@ def _adopt_session_conversation(agent: BaseAgent, session_id: str) -> None: logger.debug("Session %s: could not sync compaction live offset", scrub_log(session_id), exc_info=True) +# One agent build at a time per process. Builds used to be serialized for +# free, because each one froze the event loop, and the external MCP layer +# relies on that: a process-wide client cache and Strands' `MCPClient` +# start/stop are not safe against two builds at once. Keyed by loop because an +# asyncio.Lock binds to the loop that first waits on it (one loop in +# production, one per test). +_agent_build_locks: "weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, asyncio.Lock]" = weakref.WeakKeyDictionary() + + +def _agent_build_lock() -> asyncio.Lock: + loop = asyncio.get_running_loop() + lock = _agent_build_locks.get(loop) + if lock is None: + lock = _agent_build_locks[loop] = asyncio.Lock() + return lock + + +# Agent builds (cache misses) this process has run. Every conversation's first +# turn runs in a fresh Runtime process, so 1 marks exactly the cold-process +# build the agent-build experiment is about (stamped on `turn_prelude`). +_process_build_count = 0 + + +def process_build_count() -> int: + return _process_build_count + + +async def _build_agent_off_loop(create_kwargs: Dict[str, Any]) -> BaseAgent: + """Run the synchronous build in a worker thread, one build at a time. + + The lock is released when the THREAD finishes, not when this coroutine + does: a client disconnect cancels the await but cannot stop the thread, + and releasing early would let the next build overlap the orphaned one. + """ + lock = _agent_build_lock() + await lock.acquire() + try: + build = asyncio.ensure_future(asyncio.to_thread(create_agent, **create_kwargs)) + except BaseException: + lock.release() + raise + + def _on_build_done(finished: "asyncio.Future[BaseAgent]") -> None: + lock.release() + # Retrieve the outcome so an orphaned build's failure is not + # reported as "exception was never retrieved". + if not finished.cancelled(): + finished.exception() + + build.add_done_callback(_on_build_done) + return await asyncio.shield(build) + + async def get_agent( session_id: str, user_id: Optional[str] = None, @@ -494,9 +548,19 @@ async def get_agent( set_stage_recorder, ) + from apis.shared.feature_flags import agent_build_off_loop_enabled + + global _process_build_count + _process_build_count += 1 _stage_token = set_stage_recorder(build_stage_recorder) try: - agent = create_agent(**create_kwargs) + if agent_build_off_loop_enabled(session_id): + # The build is synchronous and ~1s on a first turn; on the loop it + # froze every other coroutine in the container for that long. The + # worker gets a copy of this context, recorder included. + agent = await _build_agent_off_loop(create_kwargs) + else: + agent = create_agent(**create_kwargs) finally: reset_stage_recorder(_stage_token) diff --git a/backend/src/apis/inference_api/chat/turn_timing.py b/backend/src/apis/inference_api/chat/turn_timing.py index bd3295df1..daad56dc0 100644 --- a/backend/src/apis/inference_api/chat/turn_timing.py +++ b/backend/src/apis/inference_api/chat/turn_timing.py @@ -280,7 +280,9 @@ def _emit_metrics( metrics.setdefault(_metric_name(prefix), total) properties: Dict[str, Any] = {"streamKind": stream_kind, "sessionId": session_id} - for key in ("isResume", "deferredBuild", "hasAssistant"): + # Properties, never dimensions: `buildArm` is the agent-build A/B and + # `processBuilds` is unbounded. + for key in ("isResume", "deferredBuild", "hasAssistant", "buildArm", "processBuilds"): if extra and key in extra: properties[key] = extra[key] diff --git a/backend/src/apis/shared/feature_flags.py b/backend/src/apis/shared/feature_flags.py index a4b4b03fe..333378850 100644 --- a/backend/src/apis/shared/feature_flags.py +++ b/backend/src/apis/shared/feature_flags.py @@ -15,7 +15,9 @@ module reload (import-time paths) without a process restart. """ +import hashlib import os +from typing import Optional def skills_enabled() -> bool: @@ -653,3 +655,88 @@ def compaction_summary_extract_enabled() -> bool: every planted fact on the quality harness, Nova Micro 88%. """ return os.environ.get("COMPACTION_SUMMARY_EXTRACT_ENABLED", "").strip().lower() != "false" + + + +AGENT_BUILD_ARMS = ("control", "shared_clients", "shared_clients_off_loop") + + +def agent_build_experiment_arm(session_id: Optional[str]) -> str: + """Which agent-build variant this session runs (an A/B experiment, default OFF). + + Two changes to the first-turn agent build, measured before either ships: + + - ``shared_clients``: AgentCore Memory session managers share one set of + boto3 clients (``memory_shared_clients_enabled``). + - ``shared_clients_off_loop``: that, plus the synchronous build runs in a + worker thread (``agent_build_off_loop_enabled``). + + ``AGENT_BUILD_EXPERIMENT`` selects the mode: + + - unset / empty / anything unrecognised: ``control`` for every session. + This is the default everywhere, so other deployments see no change. + - ``ab``: each session is hashed into one of the three arms. Every + conversation runs in its own Runtime process, so arms never share + process state, and they run interleaved in time, which cancels network + drift between arms. + - an arm name: every session runs that arm. + + The arm is stamped on ``turn_prelude`` (``buildArm``) so Logs Insights can + compare stage timings per arm. There is no CDK entry: the Runtime's + environment is capped at 50 variables and this is a temporary experiment, + so it is set out of band on the Runtime (``update-agent-runtime``), which + ``backend.yml`` deploys preserve and a ``platform.yml`` deploy resets. + """ + mode = os.environ.get("AGENT_BUILD_EXPERIMENT", "").strip().lower() + if mode in AGENT_BUILD_ARMS: + return mode + if mode != "ab" or not session_id: + return "control" + digest = hashlib.sha256(session_id.encode("utf-8")).digest() + return AGENT_BUILD_ARMS[digest[0] % len(AGENT_BUILD_ARMS)] + + +def memory_shared_clients_enabled(session_id: Optional[str]) -> bool: + """Whether this session's AgentCore Memory session manager uses shared clients. + + The SDK's ``AgentCoreMemorySessionManager.__init__`` builds a + ``MemoryClient`` (a fresh ``boto3.Session`` plus two clients) and then a + second fresh session plus two more clients that replace the first pair. + A fresh session re-loads botocore's service models, so every session + manager paid ~360ms of CPU (measured locally, before any network call), + and each one opened its own connection pool, so its first ``list_events`` + also paid a TLS handshake. With this on, the factory hands the SDK one + process-wide session whose ``client()`` returns the same client per + configuration. In a fresh process that halves the model loading; in a + warm one (a later cache-miss build in the same conversation) the clients + cost nothing and their connections are already open. + + Arm of ``agent_build_experiment_arm``; off by default. Nothing reaches + the prompt. + """ + return agent_build_experiment_arm(session_id) in ("shared_clients", "shared_clients_off_loop") + + +def agent_build_off_loop_enabled(session_id: Optional[str]) -> bool: + """Whether ``get_agent`` runs this session's synchronous build in a worker thread. + + ``create_agent`` is synchronous: prompt assembly, tool catalog lookups, + external MCP ``tools/list``, session manager construction and the + AgentCore Memory restore all run on the calling thread. Called from + ``get_agent`` on the event loop, a first-turn build (1.0-1.2s on dev) + freezes everything else in the process, including the concurrent + session-title task, whose Nova reply sits unprocessed until the build + ends. Off the loop, the stream can emit that title during the build. + + ``asyncio.to_thread`` copies contextvars (the build-stage recorder, the + AgentCore request context) into the worker. The build's sync-to-async + bridges then take their no-loop branch and ``asyncio.run`` their + coroutine directly. Builds stay one at a time (``_build_agent_off_loop``), + because the external MCP layer relied on the frozen loop for that, and + ``ExternalMCPIntegration`` locks its maps and pins handed-out clients for + the build (see ``load_external_tools``'s ``consumer_pin``). + + Arm of ``agent_build_experiment_arm``; off by default. Nothing reaches + the prompt. + """ + return agent_build_experiment_arm(session_id) == "shared_clients_off_loop" diff --git a/backend/tests/agents/main_agent/integrations/test_external_mcp_client.py b/backend/tests/agents/main_agent/integrations/test_external_mcp_client.py index 360bd3001..cc8426db9 100644 --- a/backend/tests/agents/main_agent/integrations/test_external_mcp_client.py +++ b/backend/tests/agents/main_agent/integrations/test_external_mcp_client.py @@ -902,3 +902,84 @@ def test_take_pending_consents_drains_and_is_per_user(self): self.PROVIDER: "https://consent.example/b" } assert integration.take_pending_consents("carol") == {} + + +class _ConsumerTrackingClient: + """Stand-in for a Strands `MCPClient`'s consumer bookkeeping.""" + + def __init__(self) -> None: + self.consumers: set = set() + self.load_tools = AsyncMock(return_value=[]) + + def add_consumer(self, consumer_id, **kwargs) -> None: + self.consumers.add(consumer_id) + + def remove_consumer(self, consumer_id, **kwargs) -> None: + self.consumers.discard(consumer_id) + + +class TestBuildConsumerPin: + """With the agent build in a worker thread, a turn ending on the event loop + can drop a shared client's last consumer mid-build, and Strands stops a + client the moment that happens. The build's pin keeps the count above zero + until the new agent has registered itself.""" + + @pytest.mark.asyncio + async def test_new_and_cached_clients_are_both_pinned(self): + integration = ExternalMCPIntegration() + tool = _fake_tool(datetime(2025, 1, 1, tzinfo=timezone.utc)) + repo = SimpleNamespace(get_tool=AsyncMock(return_value=tool)) + client = _ConsumerTrackingClient() + first_pin, second_pin = object(), object() + + with patch( + "apis.shared.tools.repository.get_tool_catalog_repository", + return_value=repo, + ), patch( + "agents.main_agent.integrations.external_mcp_client.create_external_mcp_client", + return_value=client, + ): + await integration.load_external_tools(["gmail"], consumer_pin=first_pin) + await integration.load_external_tools(["gmail"], consumer_pin=second_pin) + + assert client.consumers == {first_pin, second_pin} + + @pytest.mark.asyncio + async def test_no_pin_registers_nothing(self): + integration = ExternalMCPIntegration() + tool = _fake_tool(datetime(2025, 1, 1, tzinfo=timezone.utc)) + repo = SimpleNamespace(get_tool=AsyncMock(return_value=tool)) + client = _ConsumerTrackingClient() + + with patch( + "apis.shared.tools.repository.get_tool_catalog_repository", + return_value=repo, + ), patch( + "agents.main_agent.integrations.external_mcp_client.create_external_mcp_client", + return_value=client, + ): + await integration.load_external_tools(["gmail"]) + + assert client.consumers == set() + + +class TestSingletonUnderConcurrentBuilds: + def test_concurrent_first_calls_share_one_instance(self, monkeypatch): + """Two instances would split the approval map: the approval hook built + against one would find no `needs_approval` names in the other.""" + import threading + from concurrent.futures import ThreadPoolExecutor + + from agents.main_agent.integrations import external_mcp_client as module + + monkeypatch.setattr(module, "_external_mcp_integration", None) + barrier = threading.Barrier(8) + + def first_call(): + barrier.wait() + return module.get_external_mcp_integration() + + with ThreadPoolExecutor(max_workers=8) as pool: + instances = list(pool.map(lambda _: first_call(), range(8))) + + assert len({id(i) for i in instances}) == 1 diff --git a/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py b/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py new file mode 100644 index 000000000..62f8f2a5e --- /dev/null +++ b/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py @@ -0,0 +1,135 @@ +"""AgentCore Memory session managers share one set of boto3 clients. + +The SDK's ``AgentCoreMemorySessionManager.__init__`` built a ``MemoryClient`` +(fresh session, two clients) and then a second fresh session and two more +clients, on every construction. That was ~360ms of CPU per session manager on +the event loop, plus a cold connection pool. These tests build REAL +``TurnBasedSessionManager`` instances through the factory, with only the SDK's +network calls stubbed, so they fail if an SDK upgrade stops going through the +seams the factory relies on. +""" + +import threading +from concurrent.futures import ThreadPoolExecutor +from typing import Any, List + +import boto3 +import pytest + +from agents.main_agent.session import session_factory as factory +from bedrock_agentcore.memory.integrations.strands import session_manager as sdk + + +@pytest.fixture +def memory_env(monkeypatch): + """Real construction, no network: the SDK's session read/create are stubbed.""" + monkeypatch.setenv("AGENTCORE_MEMORY_ID", "mem-test") + monkeypatch.setenv("AWS_REGION", "us-west-2") + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients") + monkeypatch.setattr(factory, "_discover_strategy_ids", lambda memory_id, region: (None, None, None)) + monkeypatch.setattr(sdk.AgentCoreMemorySessionManager, "read_session", lambda self, session_id, **k: None) + monkeypatch.setattr(sdk.AgentCoreMemorySessionManager, "create_session", lambda self, session, **k: session) + # A fresh shared session per test, so client counts start from zero. + monkeypatch.setattr(factory, "_shared_memory_session", None) + + +@pytest.fixture +def client_builds(monkeypatch) -> List[str]: + """Every boto3 client actually constructed, by service name.""" + built: List[str] = [] + original = boto3.Session.client + + def counting_client(self, service_name, *args, **kwargs): + built.append(service_name) + return original(self, service_name, *args, **kwargs) + + monkeypatch.setattr(boto3.Session, "client", counting_client) + return built + + +def _build(session_id: str) -> Any: + return factory.SessionFactory.create_session_manager(session_id=session_id, user_id="user-1") + + +class TestSharedClients: + def test_the_sdk_builds_its_memory_client_through_the_shared_seam(self): + assert sdk.MemoryClient is factory._SharedSessionMemoryClient + + def test_later_session_managers_build_no_clients(self, memory_env, client_builds): + first = _build("session-a") + built_by_first = len(client_builds) + + second = _build("session-b") + + assert built_by_first > 0 + assert len(client_builds) == built_by_first + assert second.memory_client.gmdp_client is first.memory_client.gmdp_client + assert second.memory_client.gmcp_client is first.memory_client.gmcp_client + + def test_each_manager_keeps_its_own_session_state(self, memory_env, client_builds): + """Only the clients are shared; the session identity is per manager.""" + first = _build("session-a") + second = _build("session-b") + + assert first.session_id == "session-a" + assert second.session_id == "session-b" + assert first.config is not second.config + + def test_shared_clients_get_a_bigger_pool_and_keep_the_sdk_user_agent(self, memory_env): + manager = _build("session-a") + config = manager.memory_client.gmdp_client.meta.config + + assert config.max_pool_connections == factory._SHARED_MEMORY_MAX_POOL_CONNECTIONS + assert "strands-agents" in (config.user_agent_extra or "") + + def test_concurrent_first_builds_share_one_client(self, memory_env, client_builds): + """Builds run in worker threads; the first ones race to create the clients.""" + barrier = threading.Barrier(4) + + def build(i: int) -> Any: + barrier.wait() + return _build(f"session-{i}") + + with ThreadPoolExecutor(max_workers=4) as pool: + managers = list(pool.map(build, range(4))) + + assert len({id(m.memory_client.gmdp_client) for m in managers}) == 1 + + +class TestControlArm: + def test_control_keeps_the_sdks_per_manager_clients(self, memory_env, client_builds, monkeypatch): + """The default: no experiment set means every session is control.""" + monkeypatch.delenv("AGENT_BUILD_EXPERIMENT", raising=False) + + first = _build("session-a") + built_by_first = len(client_builds) + second = _build("session-b") + + assert len(client_builds) == 2 * built_by_first + assert second.memory_client.gmdp_client is not first.memory_client.gmdp_client + + def test_the_memory_client_seam_is_inert_outside_a_shared_construction(self, memory_env): + """The rebinding in the SDK module must not change a MemoryClient built + anywhere else (e.g. `_discover_strategy_ids`).""" + client = sdk.MemoryClient(region_name="us-west-2") + assert client.gmdp_client is not factory.shared_memory_boto_session().client( + "bedrock-agentcore", region_name="us-west-2" + ) + + +class TestClientReusingSession: + def test_explicit_credentials_are_never_shared(self): + session = factory._ClientReusingSession() + kwargs = dict(region_name="us-west-2", aws_access_key_id="a", aws_secret_access_key="b") + + assert session.client("sts", **kwargs) is not session.client("sts", **kwargs) + + def test_different_configs_get_different_clients(self): + from botocore.config import Config + + session = factory._ClientReusingSession() + a = session.client("sts", region_name="us-west-2", config=Config(user_agent_extra="a")) + b = session.client("sts", region_name="us-west-2", config=Config(user_agent_extra="b")) + + assert a is not b + assert session.client("sts", region_name="us-west-2", config=Config(user_agent_extra="a")) is a diff --git a/backend/tests/agents/main_agent/test_base_agent_external_registration.py b/backend/tests/agents/main_agent/test_base_agent_external_registration.py index 739d3caab..c07216f27 100644 --- a/backend/tests/agents/main_agent/test_base_agent_external_registration.py +++ b/backend/tests/agents/main_agent/test_base_agent_external_registration.py @@ -73,3 +73,40 @@ def test_bare_id_still_registers(self): BaseAgent._register_external_mcp_tools(agent) assert tool_filter._external_mcp_tools == {"canvas"} + + +class _PinnedClient: + def __init__(self, pin, fail: bool = False) -> None: + self.consumers = {pin} + self._fail = fail + + def remove_consumer(self, consumer_id, **kwargs) -> None: + if self._fail: + raise RuntimeError("stop failed") + self.consumers.discard(consumer_id) + + +class TestMcpBuildPinRelease: + """The build pins external MCP clients (see `load_external_tools`'s + `consumer_pin`) until the agent has registered as their consumer.""" + + def test_release_drops_every_pin_and_disarms(self): + pin = object() + clients = [_PinnedClient(pin), _PinnedClient(pin)] + agent = SimpleNamespace(_mcp_build_pin=pin, _mcp_pinned_clients=list(clients)) + + BaseAgent._release_mcp_build_pins(agent) + + assert all(c.consumers == set() for c in clients) + assert agent._mcp_pinned_clients == [] + # A later `_create_agent` (the stream_async fallback) takes no pin. + assert agent._mcp_build_pin is None + + def test_a_failing_release_does_not_fail_the_build(self): + pin = object() + broken, healthy = _PinnedClient(pin, fail=True), _PinnedClient(pin) + agent = SimpleNamespace(_mcp_build_pin=pin, _mcp_pinned_clients=[broken, healthy]) + + BaseAgent._release_mcp_build_pins(agent) + + assert healthy.consumers == set() diff --git a/backend/tests/apis/inference_api/test_chat_service.py b/backend/tests/apis/inference_api/test_chat_service.py index ebba2e2b3..1f36f1a8c 100644 --- a/backend/tests/apis/inference_api/test_chat_service.py +++ b/backend/tests/apis/inference_api/test_chat_service.py @@ -841,3 +841,152 @@ async def test_resume_with_the_snapshot_memory_hits_the_paused_agent( session_id="s1", user_id="u1", system_prompt="P", is_resume=True, cache_write=False, ) assert stale is not first + + +# --------------------------------------------------------------------------- +# The build runs off the event loop +# --------------------------------------------------------------------------- + + +class TestBuildOffLoop: + """``create_agent`` is synchronous and ~1s on a first turn. On the loop it + froze every coroutine in the container, including other users' streams.""" + + @pytest.mark.asyncio + async def test_the_build_runs_in_a_worker_thread(self, mock_freshness_hash, monkeypatch): + import threading + + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") + loop_thread = threading.get_ident() + build_threads = [] + + def fake_create_agent(**kwargs): + build_threads.append(threading.get_ident()) + return _fake_agent() + + with patch.object(service, "create_agent", side_effect=fake_create_agent): + await service.get_agent(session_id="s", user_id="u") + + assert build_threads and build_threads[0] != loop_thread + + @pytest.mark.asyncio + async def test_the_loop_keeps_serving_while_a_build_runs(self, mock_freshness_hash, monkeypatch): + import asyncio + import time + + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") + ticks = 0 + + async def other_stream(): + nonlocal ticks + while True: + await asyncio.sleep(0.01) + ticks += 1 + + def slow_create_agent(**kwargs): + time.sleep(0.2) + return _fake_agent() + + ticker = asyncio.create_task(other_stream()) + try: + with patch.object(service, "create_agent", side_effect=slow_create_agent): + await service.get_agent(session_id="s", user_id="u") + finally: + ticker.cancel() + + assert ticks >= 5 + + @pytest.mark.asyncio + async def test_build_stages_are_recorded_from_the_worker(self, mock_freshness_hash, monkeypatch): + from apis.shared.observability.build_stages import mark_stage + + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") + recorded = [] + + def fake_create_agent(**kwargs): + mark_stage("session_mgr") + return _fake_agent() + + with patch.object(service, "create_agent", side_effect=fake_create_agent): + await service.get_agent(session_id="s", user_id="u", build_stage_recorder=recorded.append) + + assert recorded == ["session_mgr"] + + @pytest.mark.asyncio + async def test_the_control_arm_builds_on_the_loop(self, mock_freshness_hash, monkeypatch): + """The default: no experiment set means the build stays where it was.""" + import threading + + monkeypatch.delenv("AGENT_BUILD_EXPERIMENT", raising=False) + loop_thread = threading.get_ident() + build_threads = [] + + def fake_create_agent(**kwargs): + build_threads.append(threading.get_ident()) + return _fake_agent() + + with patch.object(service, "create_agent", side_effect=fake_create_agent): + await service.get_agent(session_id="s", user_id="u") + + assert build_threads == [loop_thread] + + @pytest.mark.asyncio + async def test_builds_never_overlap(self, mock_freshness_hash, monkeypatch): + """The external MCP layer is only safe one build at a time, which the + frozen loop used to guarantee for free.""" + import asyncio + import threading + import time + + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") + active, peak = 0, 0 + guard = threading.Lock() + + def slow_create_agent(**kwargs): + nonlocal active, peak + with guard: + active += 1 + peak = max(peak, active) + time.sleep(0.05) + with guard: + active -= 1 + return _fake_agent() + + with patch.object(service, "create_agent", side_effect=slow_create_agent): + await asyncio.gather(*(service.get_agent(session_id=f"s{i}", user_id="u") for i in range(4))) + + assert peak == 1 + + @pytest.mark.asyncio + async def test_a_cancelled_build_still_holds_the_lock_until_its_thread_ends( + self, mock_freshness_hash, monkeypatch + ): + """A client disconnect cancels the await, not the thread.""" + import asyncio + import threading + + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") + first_started, release_first = threading.Event(), threading.Event() + order = [] + + def create_agent(**kwargs): + if kwargs["session_id"] == "first": + first_started.set() + release_first.wait(5) + order.append("first-finished") + else: + order.append("second-started") + return _fake_agent() + + with patch.object(service, "create_agent", side_effect=create_agent): + first = asyncio.create_task(service.get_agent(session_id="first", user_id="u")) + await asyncio.to_thread(first_started.wait, 5) + first.cancel() + second = asyncio.create_task(service.get_agent(session_id="second", user_id="u")) + await asyncio.sleep(0.05) + assert order == [] + release_first.set() + await second + + assert order == ["first-finished", "second-started"] + assert first.cancelled() diff --git a/backend/tests/apis/inference_api/test_preparing_phase.py b/backend/tests/apis/inference_api/test_preparing_phase.py index d2817f157..0f01fc4f1 100644 --- a/backend/tests/apis/inference_api/test_preparing_phase.py +++ b/backend/tests/apis/inference_api/test_preparing_phase.py @@ -251,3 +251,78 @@ def test_the_route_no_longer_races_its_own_build(self): source = Path(routes_module.__file__).read_text() assert "_PREPARING_NOTICE_SECONDS" not in source + + +class TestTitleDuringTheBuild: + """A first turn's title can land while the agent is still being built. + + The route races the deferred build against the title task and emits a + finished title before `prepared`. That only helps when the build runs off + the event loop (`agent_build_off_loop_enabled`): a synchronous build on the + loop stops the title task from finishing first, so the race reduces to the + plain await it replaced, which is the control arm's behaviour. + """ + + @staticmethod + async def _stream(build, title_task) -> List[str]: + """Mirrors the route: preparing, race, prepared.""" + frames = ['event: agent_status\ndata: {"phase": "preparing"}\n\n'] + emitted = False + + def title_sse() -> Optional[str]: + nonlocal emitted + if emitted or not title_task.done(): + return None + emitted = True + return f'event: session_title\ndata: {{"title": "{title_task.result()}"}}\n\n' + + task = asyncio.ensure_future(build()) + if not title_task.done(): + await asyncio.wait({task, title_task}, return_when=asyncio.FIRST_COMPLETED) + frame = title_sse() + if frame: + frames.append(frame) + await task + frames.append('event: agent_status\ndata: {"phase": "prepared"}\n\n') + return frames + + @staticmethod + def _kinds(frames: List[str]) -> List[str]: + return [f.split("\n", 1)[0].removeprefix("event: ") for f in frames] + + @pytest.mark.asyncio + async def test_a_build_off_the_loop_lets_the_title_out_first(self): + import time + + async def title() -> str: + await asyncio.sleep(0.02) + return "Biology Syllabus" + + async def build_in_a_thread() -> object: + return await asyncio.to_thread(time.sleep, 0.2) + + frames = await self._stream(build_in_a_thread, asyncio.ensure_future(title())) + + assert self._kinds(frames) == ["agent_status", "session_title", "agent_status"] + + @pytest.mark.asyncio + async def test_a_build_on_the_loop_cannot(self): + import time + + async def title() -> str: + await asyncio.sleep(0.02) + return "Biology Syllabus" + + async def build_on_the_loop() -> object: + time.sleep(0.2) # synchronous, as `create_agent` is + return object() + + frames = await self._stream(build_on_the_loop, asyncio.ensure_future(title())) + + assert "session_title" not in self._kinds(frames) + + def test_route_still_races_the_build_against_the_title(self): + source = Path(routes_module.__file__).read_text() + + assert "{build, title_task}" in source + assert "return_when=asyncio.FIRST_COMPLETED" in source diff --git a/backend/tests/shared/test_agent_build_experiment_arm.py b/backend/tests/shared/test_agent_build_experiment_arm.py new file mode 100644 index 000000000..adb1f64e3 --- /dev/null +++ b/backend/tests/shared/test_agent_build_experiment_arm.py @@ -0,0 +1,54 @@ +"""`agent_build_experiment_arm`: default OFF, and a stable per-session split.""" + +import uuid + +import pytest + +from apis.shared.feature_flags import ( + AGENT_BUILD_ARMS, + agent_build_experiment_arm, + agent_build_off_loop_enabled, + memory_shared_clients_enabled, +) + + +@pytest.mark.parametrize("value", [None, "", "off", "false", "true", "nonsense"]) +def test_everything_but_ab_or_an_arm_name_is_control(monkeypatch, value): + if value is None: + monkeypatch.delenv("AGENT_BUILD_EXPERIMENT", raising=False) + else: + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", value) + + assert agent_build_experiment_arm("s1") == "control" + assert not memory_shared_clients_enabled("s1") + assert not agent_build_off_loop_enabled("s1") + + +@pytest.mark.parametrize("arm", AGENT_BUILD_ARMS) +def test_an_arm_name_forces_that_arm(monkeypatch, arm): + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", f" {arm.upper()} ") + assert agent_build_experiment_arm("s1") == arm + + +def test_the_arms_imply_their_changes(monkeypatch): + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients") + assert memory_shared_clients_enabled("s") and not agent_build_off_loop_enabled("s") + + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") + assert memory_shared_clients_enabled("s") and agent_build_off_loop_enabled("s") + + +def test_ab_is_stable_per_session_and_uses_every_arm(monkeypatch): + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "ab") + sessions = [str(uuid.UUID(int=i)) for i in range(300)] + + arms = [agent_build_experiment_arm(s) for s in sessions] + + assert arms == [agent_build_experiment_arm(s) for s in sessions] + counts = {arm: arms.count(arm) for arm in AGENT_BUILD_ARMS} + assert all(60 <= n <= 140 for n in counts.values()), counts + + +def test_ab_without_a_session_is_control(monkeypatch): + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "ab") + assert agent_build_experiment_arm(None) == "control" diff --git a/docs/specs/turn-latency-preamble.md b/docs/specs/turn-latency-preamble.md index aac0de11e..fb259a9a9 100644 --- a/docs/specs/turn-latency-preamble.md +++ b/docs/specs/turn-latency-preamble.md @@ -746,7 +746,74 @@ synchronous SDK handshake, TLS setup, or tool-registry work nobody has looked at. Second target after that: `agent_build.session_mgr` at 830ms (AgentCore Memory -restore, never timed). +restore, never timed). Now PR-6. + +## PR-6 — agent-build A/B: shared Memory clients, and the build off the loop (IN PROGRESS) + +The second target named above, `agent_build.session_mgr`, measured **616-738ms** +on dev first turns (2026-09-28) — more than half the 1.0-1.2s build. + +**What it does.** The AgentCore SDK's `AgentCoreMemorySessionManager.__init__` +builds a `MemoryClient` (a fresh `boto3.Session` plus two clients), then a second +fresh session plus two more clients that **replace** the first pair, then calls +`read_session`. For a new session that is two sequential `list_events` (the +second is a legacy-format fallback) and a `create_event`, each on a cold +connection pool. Laptop timing: ~360ms of client construction per session +manager. PR-3 above is the warning about reading that number: construction was +~30x slower on the container than on a laptop. + +**Every first turn is a fresh process.** Each conversation's turns run in their +own Runtime process (`service.instance.id` differs per session in the runtime +logs). So a first turn always pays cold construction, and "a build freezes other +users' streams" does not happen: a process serves one conversation. Process-wide +caching helps a first turn only by doing less cold work (one session loads the +service models once instead of twice); it helps later cache-miss builds in the +same conversation (an `@`-mention, a changed toolset) fully. + +**Two changes, behind one per-session experiment flag, default off** +(`agent_build_experiment_arm` in `apis/shared/feature_flags.py`, +`AGENT_BUILD_EXPERIMENT`): + +| Arm | Change | +|---|---| +| `control` | today's build | +| `shared_clients` | the factory hands the SDK one process-wide boto3 session whose `client()` returns one client per configuration, and rebinds the SDK module's `MemoryClient` so the discarded pair is built from it too | +| `shared_clients_off_loop` | that, plus `create_agent` runs under `asyncio.to_thread`, one build at a time, and the route emits a title that lands mid-build | + +The off-loop arm needed hardening that the frozen loop used to provide for free: +builds stay serialized (the lock is released when the *thread* ends, so a +cancelled request cannot let a second build overlap an orphan), +`ExternalMCPIntegration` locks its maps and creates its singleton under a lock +(two instances would split the approval map, so a `needs_approval` tool could run +unapproved), and each build pins the external MCP clients it is handed until its +agent registers as their consumer (Strands stops a client whose last consumer +goes). + +**New sub-stages.** `agent_build.session_mgr_clients` (everything before the SDK's +`read_session`, i.e. client setup) now precedes `agent_build.session_mgr` (the +session read/create network). `agent_build.strands_agent` (Strands' own `Agent` +construction, including MCP `load_tools`) now precedes `agent_build.finalize` +(the session restore). `turn_prelude` carries `buildArm` and `processBuilds` +(1 on a first turn). + +**How to run it.** Set `AGENT_BUILD_EXPERIMENT=ab` on the dev Runtime out of band +(`update-agent-runtime`; the Runtime is at 48 of its 50 environment variables, +and a temporary experiment should not take a CDK slot). `backend.yml` deploys +preserve it; a `platform.yml` deploy resets it. Then: + + cd backend + AWS_PROFILE=dev-ai uv run python scripts/experiment_agent_build_arms.py \ + --user-id --per-arm 15 --cleanup + +Arms are assigned by hashing the session id, and the script runs one turn per +arm per round in rotating order, so time-of-day drift hits every arm alike. + +**Decision rule, fixed before the data.** Ship `shared_clients` (default on, with +a kill switch) if its median `agent_build.session_mgr_clients` + +`agent_build.session_mgr` beats control's by more than the run-to-run spread +and no stage regresses. Ship the off-loop build only if the title demonstrably +lands before `prepared` and time to first token does not regress; otherwise +remove it and the MCP hardening that exists only for it. ## Declined / overtaken — `asyncio.to_thread` for the DynamoDB calls From 8113f3f8c5df489a26737edea0f0bdd032192b2e Mon Sep 17 00:00:00 2001 From: Phil Merrell Date: Wed, 30 Sep 2026 08:17:08 -0600 Subject: [PATCH 2/4] refactor(backend): withdraw the off-loop agent-build arm and its MCP hardening `AGENT_BUILD_ARMS` is now `("control", "shared_clients")`. The `shared_clients_off_loop` arm ran `create_agent` under `asyncio.to_thread` so a first-turn title could go out mid-build, and needed hardening the frozen loop used to provide for free: process-wide build serialization, a lock on `ExternalMCPIntegration`'s maps and singleton, and a consumer pin on every external MCP client handed to a build. A thread does not make a synchronous build faster, so the arm could not reduce time to first token, and the overlap it would have enabled is reachable inside the constructor without any of that (see docs/specs/turn-path-ttft.md, sections 4 and 5 P3). All of it goes, with its tests; the route's title race reverts to the plain awaited build. `process_build_count` stays: it stamps `processBuilds` on `turn_prelude`. Co-Authored-By: Claude Fable 5.1 --- .../scripts/experiment_agent_build_arms.py | 21 +-- backend/src/agents/main_agent/base_agent.py | 26 +-- .../integrations/external_mcp_client.py | 142 ++++++----------- backend/src/apis/inference_api/chat/routes.py | 89 ++++------- .../src/apis/inference_api/chat/service.py | 54 +------ backend/src/apis/shared/feature_flags.py | 46 ++---- .../integrations/test_external_mcp_client.py | 81 ---------- .../test_base_agent_external_registration.py | 37 ----- .../apis/inference_api/test_chat_service.py | 149 ------------------ .../inference_api/test_preparing_phase.py | 75 --------- .../shared/test_agent_build_experiment_arm.py | 18 ++- docs/specs/turn-latency-preamble.md | 55 ++++--- 12 files changed, 161 insertions(+), 632 deletions(-) diff --git a/backend/scripts/experiment_agent_build_arms.py b/backend/scripts/experiment_agent_build_arms.py index a41fdfae2..001262361 100644 --- a/backend/scripts/experiment_agent_build_arms.py +++ b/backend/scripts/experiment_agent_build_arms.py @@ -1,11 +1,12 @@ -"""Agent-build A/B: does sharing Memory clients / building off the loop help a first turn? +"""Agent-build A/B: does sharing boto3 clients help a first turn? Drives real first turns through the deployed AgentCore Runtime and compares the -three arms of ``agent_build_experiment_arm`` (``apis/shared/feature_flags.py``): +two arms of ``agent_build_experiment_arm`` (``apis/shared/feature_flags.py``): - control today's build - shared_clients AgentCore Memory session managers share boto3 clients - shared_clients_off_loop that, plus the synchronous build runs in a worker thread + control today's build + shared_clients the session manager, strategy-id discovery and Bedrock + model client share one process-wide boto3 session, built + at container warm-up **Precondition:** the Runtime must have ``AGENT_BUILD_EXPERIMENT=ab``. Arms are assigned server-side by hashing the session id, so this script generates session @@ -23,8 +24,10 @@ sub-stages, ``buildArm`` and ``processBuilds``. What it establishes: per-arm medians for the build and its sub-stages, time to -first token, and whether the title beats ``prepared``. What it cannot: fleet -magnitude, or behaviour under concurrent load. Report which claim you make. +first token, and whether the title beats ``prepared`` (it cannot while the build +is synchronous on the loop; the column is kept so a regression there is visible). +What it cannot: fleet magnitude, or behaviour under concurrent load. Report which +claim you make. Each turn is a real conversation owned by ``--user-id``: it draws on that user's quota and appears in their sidebar. Use ``--cleanup`` to soft-delete the @@ -35,7 +38,7 @@ cd backend AWS_PROFILE=dev-ai uv run python scripts/experiment_agent_build_arms.py \\ - --user-id --per-arm 15 + --user-id --per-arm 15 --cleanup """ from __future__ import annotations @@ -60,7 +63,7 @@ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "src")) -ARMS = ("control", "shared_clients", "shared_clients_off_loop") +ARMS = ("control", "shared_clients") # Frames whose first arrival is recorded, keyed by how the table names them. _TIMED = ("preparing", "prepared", "session_title", "first_token", "done") diff --git a/backend/src/agents/main_agent/base_agent.py b/backend/src/agents/main_agent/base_agent.py index aabe7f9f5..5e59d1abf 100644 --- a/backend/src/agents/main_agent/base_agent.py +++ b/backend/src/agents/main_agent/base_agent.py @@ -201,29 +201,10 @@ def __init__( # Initialize streaming coordinator self.stream_coordinator = StreamCoordinator() - # Create the agent (subclass-specific). External MCP clients handed to - # it are pinned for the build (see `load_external_tools`'s - # `consumer_pin`); once the agent has registered as their consumer, or - # failed to, the pin goes. - self._mcp_build_pin = object() - self._mcp_pinned_clients: List[Any] = [] - try: - self._create_agent() - finally: - self._release_mcp_build_pins() + # Create the agent (subclass-specific) + self._create_agent() mark_stage("finalize") - def _release_mcp_build_pins(self) -> None: - for client in self._mcp_pinned_clients: - try: - client.remove_consumer(self._mcp_build_pin) - except Exception: # noqa: BLE001 - a failed stop must not fail the build - logger.warning("Releasing an MCP build pin failed", exc_info=True) - self._mcp_pinned_clients = [] - # Scoped to the constructor's build: a later `_create_agent` (the - # `stream_async` fallback) takes no pin, so it can never leak one. - self._mcp_build_pin = None - @abstractmethod def _create_agent(self) -> None: """Create the specific agent type. Subclasses must implement.""" @@ -569,7 +550,6 @@ async def _load_with_context(): external_mcp_tool_ids, user_id=self.user_id, auth_token=self.auth_token, - consumer_pin=getattr(self, "_mcp_build_pin", None), ) # Probe with ``get_running_loop`` rather than ``get_event_loop``: @@ -592,8 +572,6 @@ async def _load_with_context(): future = executor.submit(asyncio.run, _load_with_context()) external_clients = future.result() - if getattr(self, "_mcp_build_pin", None) is not None: - self._mcp_pinned_clients = list(external_clients) for client in external_clients: if client not in local_tools: local_tools.append(client) diff --git a/backend/src/agents/main_agent/integrations/external_mcp_client.py b/backend/src/agents/main_agent/integrations/external_mcp_client.py index 84ef09be2..bdce9ff21 100644 --- a/backend/src/agents/main_agent/integrations/external_mcp_client.py +++ b/backend/src/agents/main_agent/integrations/external_mcp_client.py @@ -15,7 +15,6 @@ import logging import re -import threading from typing import Any, Callable, Iterator, Optional, List, Set from urllib.parse import urlparse @@ -372,13 +371,6 @@ def __init__(self): # tool silently. Keyed by user because this integration is a # process-wide singleton shared across concurrent sessions. self._pending_consents: dict[str, dict[str, str]] = {} - # Guards every dict above. Agent builds run in a worker thread - # (``agent_build_off_loop_enabled``) while the event loop reads these - # maps (``get_client``, ``take_pending_consents``); iterating one while - # a build inserts raises. Held only around dict access, never across - # an await: a lock held through an MCP pre-flight would stall the loop - # for the length of a network round trip. - self._lock = threading.RLock() def take_pending_consents(self, user_id: str) -> dict[str, str]: """Pop and return {provider_id: authorization_url} for `user_id`. @@ -389,8 +381,7 @@ def take_pending_consents(self, user_id: str) -> dict[str, str]: agent-cache hit no loading happens, nothing is recorded, and nothing is emitted — correct, because the prompt already went out once. """ - with self._lock: - return self._pending_consents.pop(user_id, {}) + return self._pending_consents.pop(user_id, {}) async def _recover_oauth_preflight( self, @@ -494,8 +485,7 @@ async def _recover_oauth_preflight( ) return False - with self._lock: - self._pending_consents.setdefault(user_id, {})[provider_id] = authorization_url + self._pending_consents.setdefault(user_id, {})[provider_id] = authorization_url logger.info( f"External MCP tool {tool_id} needs {provider_id} consent; " "surfacing oauth_required instead of dropping it silently" @@ -520,7 +510,6 @@ async def load_external_tools( enabled_tool_ids: List[str], user_id: Optional[str] = None, auth_token: Optional[str] = None, - consumer_pin: Any = None, ) -> List[MCPClient]: """ Load external MCP clients for enabled tools. @@ -535,15 +524,6 @@ async def load_external_tools( enabled_tool_ids: List of enabled tool IDs user_id: User ID (required for OAuth-gated and OIDC-forwarded tools) auth_token: Raw OIDC token for forwarding - consumer_pin: When given, registered as a consumer of every - returned client, so no client can be stopped between this - hand-out and the new agent registering itself as a consumer. - A cached client is shared across agents, and Strands stops one - the moment its last consumer goes (``remove_consumer``). With - the build in a worker thread, a turn ending on the event loop - could drop that last consumer mid-build and leave the new agent - holding tools from a stopped client. The caller removes the pin - once the agent exists (``BaseAgent.__init__``). Returns: List of MCPClient instances to add to the agent's tools @@ -589,22 +569,21 @@ async def load_external_tools( to_iso(tool.updated_at) if tool.updated_at else "" ) - with self._lock: - cached = self.clients.get(cache_key) - if cached is not None and self._client_versions.get(cache_key) == tool_version: - if consumer_pin is not None: - cached.add_consumer(consumer_pin) - clients.append(cached) - continue + if ( + cache_key in self.clients + and self._client_versions.get(cache_key) == tool_version + ): + clients.append(self.clients[cache_key]) + continue - # Stale entry — admin edited this tool since the client - # was built. Drop it so the block below creates a fresh - # client with the current config. - if cached is not None: - self.clients.pop(cache_key, None) - self._client_versions.pop(cache_key, None) - self._provider_for_client_id.pop(id(cached), None) - self._approval_names_for_client_id.pop(id(cached), None) + # Stale entry — admin edited this tool since the client + # was built. Drop it so the block below creates a fresh + # client with the current config. + if cache_key in self.clients: + stale = self.clients.pop(cache_key) + self._client_versions.pop(cache_key, None) + self._provider_for_client_id.pop(id(stale), None) + self._approval_names_for_client_id.pop(id(stale), None) static_token: Optional[str] = None token_provider: Optional[Callable[[], Optional[str]]] = None @@ -719,16 +698,13 @@ async def _exchange( if not recovered: continue + self.clients[cache_key] = client + self._client_versions[cache_key] = tool_version + if provider_id: + self._provider_for_client_id[id(client)] = provider_id approval_names = tool.mcp_config.approval_required_names() - with self._lock: - self.clients[cache_key] = client - self._client_versions[cache_key] = tool_version - if provider_id: - self._provider_for_client_id[id(client)] = provider_id - if approval_names: - self._approval_names_for_client_id[id(client)] = approval_names - if consumer_pin is not None: - client.add_consumer(consumer_pin) + if approval_names: + self._approval_names_for_client_id[id(client)] = approval_names clients.append(client) auth_label = ( " (with OIDC forwarding)" if forward_auth and static_token @@ -763,22 +739,19 @@ def get_client(self, tool_id: str, user_id: Optional[str] = None) -> Optional[MC base = base_tool_id(tool_id) # Exact keys win — a whole-server binding has no "|allow:" suffix. exact_keys = [f"{user_id}:{base}", base] if user_id else [base] - with self._lock: - for key in exact_keys: - if key in self.clients: - return self.clients[key] - # Subset-scoped fallback: cache key is "|allow:". - for key in exact_keys: - prefix = f"{key}|allow:" - for cache_key, client in self.clients.items(): - if cache_key.startswith(prefix): - return client + for key in exact_keys: + if key in self.clients: + return self.clients[key] + # Subset-scoped fallback: cache key is "|allow:". + for key in exact_keys: + prefix = f"{key}|allow:" + for cache_key, client in self.clients.items(): + if cache_key.startswith(prefix): + return client return None def add_to_tool_list(self, tools: List[Any]) -> List[Any]: - with self._lock: - cached = list(self.clients.values()) - for client in cached: + for client in self.clients.values(): if client not in tools: tools.append(client) return tools @@ -791,16 +764,15 @@ def clear_user_clients(self, user_id: str) -> None: agent build creates fresh clients (and the token cache miss forces a new consent flow). """ - with self._lock: - keys_to_remove = [ - key for key in self.clients.keys() - if key.startswith(f"{user_id}:") - ] - for key in keys_to_remove: - client = self.clients.pop(key) - self._client_versions.pop(key, None) - self._provider_for_client_id.pop(id(client), None) - self._approval_names_for_client_id.pop(id(client), None) + keys_to_remove = [ + key for key in self.clients.keys() + if key.startswith(f"{user_id}:") + ] + for key in keys_to_remove: + client = self.clients.pop(key) + self._client_versions.pop(key, None) + self._provider_for_client_id.pop(id(client), None) + self._approval_names_for_client_id.pop(id(client), None) if keys_to_remove: logger.info(f"Cleared {len(keys_to_remove)} cached MCP clients for user {user_id}") @@ -814,36 +786,26 @@ def clear_tool_clients(self, tool_id: str) -> None: config. Without this, clients cached at process start continue to point at the old URL for the lifetime of the process. """ - with self._lock: - keys_to_remove = [ - key for key in self.clients.keys() - if key == tool_id or key.endswith(f":{tool_id}") - ] - for key in keys_to_remove: - client = self.clients.pop(key) - self._client_versions.pop(key, None) - self._provider_for_client_id.pop(id(client), None) - self._approval_names_for_client_id.pop(id(client), None) + keys_to_remove = [ + key for key in self.clients.keys() + if key == tool_id or key.endswith(f":{tool_id}") + ] + for key in keys_to_remove: + client = self.clients.pop(key) + self._client_versions.pop(key, None) + self._provider_for_client_id.pop(id(client), None) + self._approval_names_for_client_id.pop(id(client), None) if keys_to_remove: logger.info(f"Cleared {len(keys_to_remove)} cached MCP clients for tool {tool_id}") _external_mcp_integration: Optional[ExternalMCPIntegration] = None -_external_mcp_integration_lock = threading.Lock() def get_external_mcp_integration() -> ExternalMCPIntegration: - """Get or create the global ExternalMCPIntegration instance. - - Locked because two first builds can now run concurrently in worker - threads. Two instances would split the maps: the approval hook built - against one would read an empty ``approval_names_for_client`` from the - other, and a ``needs_approval`` tool would run without asking. - """ + """Get or create the global ExternalMCPIntegration instance.""" global _external_mcp_integration if _external_mcp_integration is None: - with _external_mcp_integration_lock: - if _external_mcp_integration is None: - _external_mcp_integration = ExternalMCPIntegration() + _external_mcp_integration = ExternalMCPIntegration() return _external_mcp_integration diff --git a/backend/src/apis/inference_api/chat/routes.py b/backend/src/apis/inference_api/chat/routes.py index 25ad85a0c..3630bd13b 100644 --- a/backend/src/apis/inference_api/chat/routes.py +++ b/backend/src/apis/inference_api/chat/routes.py @@ -3803,45 +3803,43 @@ async def _build_main_agent(): } ) - # One-shot `session_title` SSE: once the concurrent title task - # (kicked off before the quota check on first turns) finishes, - # push the title to the client so the sidebar/header rename in - # parallel with the pending response instead of at stream end. - # Never awaited, so it adds no latency. Polled from two places: - # the coordinator's live status merge (every 100ms, so a title - # that lands during the model's time-to-first-token or a long - # tool call goes out right away) and between agent events below - # (the only route while that merge is switched off). It is also - # checked while a deferred build is in flight, which only a build - # off the event loop can reach (`agent_build_off_loop_enabled`). A - # stream that outruns Nova Micro simply never emits and the SPA's - # post-close metadata refresh covers it. - title_emitted = False - - def _session_title_sse() -> Optional[str]: - nonlocal title_emitted - if title_emitted or title_task is None or not title_task.done(): - return None - title_emitted = True - try: - generated_title = title_task.result() - except Exception as title_err: # noqa: BLE001 - cancelled/failed task must not break the stream - logger.warning("Title task unavailable for SSE emit: %s", title_err) - return None - # Generation failures return the "New Conversation" - # placeholder — nothing worth pushing over the wire. - if not generated_title or generated_title == "New Conversation": - return None - payload = { - "type": "session_title", - "sessionId": input_data.session_id, - "title": generated_title, - } - return f"event: session_title\ndata: {json.dumps(payload)}\n\n" - # Create stream with optional quota warning injection async def stream_with_quota_warning() -> AsyncGenerator[str, None]: """Wrap agent stream to inject quota warning at start if needed""" + # One-shot `session_title` SSE: once the concurrent title task + # (kicked off before the quota check on first turns) finishes, + # push the title to the client so the sidebar/header rename in + # parallel with the pending response instead of at stream end. + # Never awaited, so it adds no latency. Polled from two places: + # the coordinator's live status merge (every 100ms, so a title + # that lands during the model's time-to-first-token or a long + # tool call goes out right away) and between agent events below + # (the only route while that merge is switched off). A stream + # that outruns Nova Micro simply never emits and the SPA's + # post-close metadata refresh covers it. + title_emitted = False + + def _session_title_sse() -> Optional[str]: + nonlocal title_emitted + if title_emitted or title_task is None or not title_task.done(): + return None + title_emitted = True + try: + generated_title = title_task.result() + except Exception as title_err: # noqa: BLE001 - cancelled/failed task must not break the stream + logger.warning("Title task unavailable for SSE emit: %s", title_err) + return None + # Generation failures return the "New Conversation" + # placeholder — nothing worth pushing over the wire. + if not generated_title or generated_title == "New Conversation": + return None + payload = { + "type": "session_title", + "sessionId": input_data.session_id, + "title": generated_title, + } + return f"event: session_title\ndata: {json.dumps(payload)}\n\n" + # Yield quota warning event first if applicable if quota_warning_event: yield quota_warning_event.to_sse_format() @@ -4101,24 +4099,7 @@ async def _guarded_stream() -> AsyncGenerator[str, None]: + "\n\n" ) try: - # Raced against the title so a title that lands - # mid-build goes out mid-build. With the build on the - # event loop the title task cannot finish first, so - # this reduces to the plain await it replaced. - build = asyncio.ensure_future(_build_main_agent()) - try: - if title_task is not None and not title_task.done(): - await asyncio.wait( - {build, title_task}, - return_when=asyncio.FIRST_COMPLETED, - ) - title_sse = _session_title_sse() - if title_sse: - yield title_sse - agent = await build - finally: - if not build.done(): - build.cancel() + agent = await _build_main_agent() except Exception as build_error: # The handler has already returned, so the two `except` # arms below cannot see this — a build that fails here diff --git a/backend/src/apis/inference_api/chat/service.py b/backend/src/apis/inference_api/chat/service.py index 8f1da3d24..b9222fece 100644 --- a/backend/src/apis/inference_api/chat/service.py +++ b/backend/src/apis/inference_api/chat/service.py @@ -5,7 +5,6 @@ import asyncio import json -import weakref import logging import hashlib import os @@ -296,23 +295,6 @@ def _adopt_session_conversation(agent: BaseAgent, session_id: str) -> None: logger.debug("Session %s: could not sync compaction live offset", scrub_log(session_id), exc_info=True) -# One agent build at a time per process. Builds used to be serialized for -# free, because each one froze the event loop, and the external MCP layer -# relies on that: a process-wide client cache and Strands' `MCPClient` -# start/stop are not safe against two builds at once. Keyed by loop because an -# asyncio.Lock binds to the loop that first waits on it (one loop in -# production, one per test). -_agent_build_locks: "weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, asyncio.Lock]" = weakref.WeakKeyDictionary() - - -def _agent_build_lock() -> asyncio.Lock: - loop = asyncio.get_running_loop() - lock = _agent_build_locks.get(loop) - if lock is None: - lock = _agent_build_locks[loop] = asyncio.Lock() - return lock - - # Agent builds (cache misses) this process has run. Every conversation's first # turn runs in a fresh Runtime process, so 1 marks exactly the cold-process # build the agent-build experiment is about (stamped on `turn_prelude`). @@ -323,32 +305,6 @@ def process_build_count() -> int: return _process_build_count -async def _build_agent_off_loop(create_kwargs: Dict[str, Any]) -> BaseAgent: - """Run the synchronous build in a worker thread, one build at a time. - - The lock is released when the THREAD finishes, not when this coroutine - does: a client disconnect cancels the await but cannot stop the thread, - and releasing early would let the next build overlap the orphaned one. - """ - lock = _agent_build_lock() - await lock.acquire() - try: - build = asyncio.ensure_future(asyncio.to_thread(create_agent, **create_kwargs)) - except BaseException: - lock.release() - raise - - def _on_build_done(finished: "asyncio.Future[BaseAgent]") -> None: - lock.release() - # Retrieve the outcome so an orphaned build's failure is not - # reported as "exception was never retrieved". - if not finished.cancelled(): - finished.exception() - - build.add_done_callback(_on_build_done) - return await asyncio.shield(build) - - async def get_agent( session_id: str, user_id: Optional[str] = None, @@ -548,19 +504,11 @@ async def get_agent( set_stage_recorder, ) - from apis.shared.feature_flags import agent_build_off_loop_enabled - global _process_build_count _process_build_count += 1 _stage_token = set_stage_recorder(build_stage_recorder) try: - if agent_build_off_loop_enabled(session_id): - # The build is synchronous and ~1s on a first turn; on the loop it - # froze every other coroutine in the container for that long. The - # worker gets a copy of this context, recorder included. - agent = await _build_agent_off_loop(create_kwargs) - else: - agent = create_agent(**create_kwargs) + agent = create_agent(**create_kwargs) finally: reset_stage_recorder(_stage_token) diff --git a/backend/src/apis/shared/feature_flags.py b/backend/src/apis/shared/feature_flags.py index 333378850..36c4d42aa 100644 --- a/backend/src/apis/shared/feature_flags.py +++ b/backend/src/apis/shared/feature_flags.py @@ -658,24 +658,29 @@ def compaction_summary_extract_enabled() -> bool: -AGENT_BUILD_ARMS = ("control", "shared_clients", "shared_clients_off_loop") +AGENT_BUILD_ARMS = ("control", "shared_clients") def agent_build_experiment_arm(session_id: Optional[str]) -> str: """Which agent-build variant this session runs (an A/B experiment, default OFF). - Two changes to the first-turn agent build, measured before either ships: + One change to the first-turn agent build, measured before it ships: - - ``shared_clients``: AgentCore Memory session managers share one set of - boto3 clients (``memory_shared_clients_enabled``). - - ``shared_clients_off_loop``: that, plus the synchronous build runs in a - worker thread (``agent_build_off_loop_enabled``). + - ``shared_clients``: AgentCore Memory session managers, the strategy-id + discovery and the Bedrock model client share one process-wide boto3 + session (``memory_shared_clients_enabled``). + + A second arm, ``shared_clients_off_loop`` (the synchronous build on a + worker thread), was withdrawn before the A/B ran: a thread does not make + a synchronous build faster, and the overlap it would have enabled is + reachable inside the constructor without the MCP hardening it needed + (docs/specs/turn-path-ttft.md, sections 4 and 5 P3). ``AGENT_BUILD_EXPERIMENT`` selects the mode: - unset / empty / anything unrecognised: ``control`` for every session. This is the default everywhere, so other deployments see no change. - - ``ab``: each session is hashed into one of the three arms. Every + - ``ab``: each session is hashed into one of the two arms. Every conversation runs in its own Runtime process, so arms never share process state, and they run interleaved in time, which cancels network drift between arms. @@ -714,29 +719,4 @@ def memory_shared_clients_enabled(session_id: Optional[str]) -> bool: Arm of ``agent_build_experiment_arm``; off by default. Nothing reaches the prompt. """ - return agent_build_experiment_arm(session_id) in ("shared_clients", "shared_clients_off_loop") - - -def agent_build_off_loop_enabled(session_id: Optional[str]) -> bool: - """Whether ``get_agent`` runs this session's synchronous build in a worker thread. - - ``create_agent`` is synchronous: prompt assembly, tool catalog lookups, - external MCP ``tools/list``, session manager construction and the - AgentCore Memory restore all run on the calling thread. Called from - ``get_agent`` on the event loop, a first-turn build (1.0-1.2s on dev) - freezes everything else in the process, including the concurrent - session-title task, whose Nova reply sits unprocessed until the build - ends. Off the loop, the stream can emit that title during the build. - - ``asyncio.to_thread`` copies contextvars (the build-stage recorder, the - AgentCore request context) into the worker. The build's sync-to-async - bridges then take their no-loop branch and ``asyncio.run`` their - coroutine directly. Builds stay one at a time (``_build_agent_off_loop``), - because the external MCP layer relied on the frozen loop for that, and - ``ExternalMCPIntegration`` locks its maps and pins handed-out clients for - the build (see ``load_external_tools``'s ``consumer_pin``). - - Arm of ``agent_build_experiment_arm``; off by default. Nothing reaches - the prompt. - """ - return agent_build_experiment_arm(session_id) == "shared_clients_off_loop" + return agent_build_experiment_arm(session_id) == "shared_clients" diff --git a/backend/tests/agents/main_agent/integrations/test_external_mcp_client.py b/backend/tests/agents/main_agent/integrations/test_external_mcp_client.py index cc8426db9..360bd3001 100644 --- a/backend/tests/agents/main_agent/integrations/test_external_mcp_client.py +++ b/backend/tests/agents/main_agent/integrations/test_external_mcp_client.py @@ -902,84 +902,3 @@ def test_take_pending_consents_drains_and_is_per_user(self): self.PROVIDER: "https://consent.example/b" } assert integration.take_pending_consents("carol") == {} - - -class _ConsumerTrackingClient: - """Stand-in for a Strands `MCPClient`'s consumer bookkeeping.""" - - def __init__(self) -> None: - self.consumers: set = set() - self.load_tools = AsyncMock(return_value=[]) - - def add_consumer(self, consumer_id, **kwargs) -> None: - self.consumers.add(consumer_id) - - def remove_consumer(self, consumer_id, **kwargs) -> None: - self.consumers.discard(consumer_id) - - -class TestBuildConsumerPin: - """With the agent build in a worker thread, a turn ending on the event loop - can drop a shared client's last consumer mid-build, and Strands stops a - client the moment that happens. The build's pin keeps the count above zero - until the new agent has registered itself.""" - - @pytest.mark.asyncio - async def test_new_and_cached_clients_are_both_pinned(self): - integration = ExternalMCPIntegration() - tool = _fake_tool(datetime(2025, 1, 1, tzinfo=timezone.utc)) - repo = SimpleNamespace(get_tool=AsyncMock(return_value=tool)) - client = _ConsumerTrackingClient() - first_pin, second_pin = object(), object() - - with patch( - "apis.shared.tools.repository.get_tool_catalog_repository", - return_value=repo, - ), patch( - "agents.main_agent.integrations.external_mcp_client.create_external_mcp_client", - return_value=client, - ): - await integration.load_external_tools(["gmail"], consumer_pin=first_pin) - await integration.load_external_tools(["gmail"], consumer_pin=second_pin) - - assert client.consumers == {first_pin, second_pin} - - @pytest.mark.asyncio - async def test_no_pin_registers_nothing(self): - integration = ExternalMCPIntegration() - tool = _fake_tool(datetime(2025, 1, 1, tzinfo=timezone.utc)) - repo = SimpleNamespace(get_tool=AsyncMock(return_value=tool)) - client = _ConsumerTrackingClient() - - with patch( - "apis.shared.tools.repository.get_tool_catalog_repository", - return_value=repo, - ), patch( - "agents.main_agent.integrations.external_mcp_client.create_external_mcp_client", - return_value=client, - ): - await integration.load_external_tools(["gmail"]) - - assert client.consumers == set() - - -class TestSingletonUnderConcurrentBuilds: - def test_concurrent_first_calls_share_one_instance(self, monkeypatch): - """Two instances would split the approval map: the approval hook built - against one would find no `needs_approval` names in the other.""" - import threading - from concurrent.futures import ThreadPoolExecutor - - from agents.main_agent.integrations import external_mcp_client as module - - monkeypatch.setattr(module, "_external_mcp_integration", None) - barrier = threading.Barrier(8) - - def first_call(): - barrier.wait() - return module.get_external_mcp_integration() - - with ThreadPoolExecutor(max_workers=8) as pool: - instances = list(pool.map(lambda _: first_call(), range(8))) - - assert len({id(i) for i in instances}) == 1 diff --git a/backend/tests/agents/main_agent/test_base_agent_external_registration.py b/backend/tests/agents/main_agent/test_base_agent_external_registration.py index c07216f27..739d3caab 100644 --- a/backend/tests/agents/main_agent/test_base_agent_external_registration.py +++ b/backend/tests/agents/main_agent/test_base_agent_external_registration.py @@ -73,40 +73,3 @@ def test_bare_id_still_registers(self): BaseAgent._register_external_mcp_tools(agent) assert tool_filter._external_mcp_tools == {"canvas"} - - -class _PinnedClient: - def __init__(self, pin, fail: bool = False) -> None: - self.consumers = {pin} - self._fail = fail - - def remove_consumer(self, consumer_id, **kwargs) -> None: - if self._fail: - raise RuntimeError("stop failed") - self.consumers.discard(consumer_id) - - -class TestMcpBuildPinRelease: - """The build pins external MCP clients (see `load_external_tools`'s - `consumer_pin`) until the agent has registered as their consumer.""" - - def test_release_drops_every_pin_and_disarms(self): - pin = object() - clients = [_PinnedClient(pin), _PinnedClient(pin)] - agent = SimpleNamespace(_mcp_build_pin=pin, _mcp_pinned_clients=list(clients)) - - BaseAgent._release_mcp_build_pins(agent) - - assert all(c.consumers == set() for c in clients) - assert agent._mcp_pinned_clients == [] - # A later `_create_agent` (the stream_async fallback) takes no pin. - assert agent._mcp_build_pin is None - - def test_a_failing_release_does_not_fail_the_build(self): - pin = object() - broken, healthy = _PinnedClient(pin, fail=True), _PinnedClient(pin) - agent = SimpleNamespace(_mcp_build_pin=pin, _mcp_pinned_clients=[broken, healthy]) - - BaseAgent._release_mcp_build_pins(agent) - - assert healthy.consumers == set() diff --git a/backend/tests/apis/inference_api/test_chat_service.py b/backend/tests/apis/inference_api/test_chat_service.py index 1f36f1a8c..ebba2e2b3 100644 --- a/backend/tests/apis/inference_api/test_chat_service.py +++ b/backend/tests/apis/inference_api/test_chat_service.py @@ -841,152 +841,3 @@ async def test_resume_with_the_snapshot_memory_hits_the_paused_agent( session_id="s1", user_id="u1", system_prompt="P", is_resume=True, cache_write=False, ) assert stale is not first - - -# --------------------------------------------------------------------------- -# The build runs off the event loop -# --------------------------------------------------------------------------- - - -class TestBuildOffLoop: - """``create_agent`` is synchronous and ~1s on a first turn. On the loop it - froze every coroutine in the container, including other users' streams.""" - - @pytest.mark.asyncio - async def test_the_build_runs_in_a_worker_thread(self, mock_freshness_hash, monkeypatch): - import threading - - monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") - loop_thread = threading.get_ident() - build_threads = [] - - def fake_create_agent(**kwargs): - build_threads.append(threading.get_ident()) - return _fake_agent() - - with patch.object(service, "create_agent", side_effect=fake_create_agent): - await service.get_agent(session_id="s", user_id="u") - - assert build_threads and build_threads[0] != loop_thread - - @pytest.mark.asyncio - async def test_the_loop_keeps_serving_while_a_build_runs(self, mock_freshness_hash, monkeypatch): - import asyncio - import time - - monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") - ticks = 0 - - async def other_stream(): - nonlocal ticks - while True: - await asyncio.sleep(0.01) - ticks += 1 - - def slow_create_agent(**kwargs): - time.sleep(0.2) - return _fake_agent() - - ticker = asyncio.create_task(other_stream()) - try: - with patch.object(service, "create_agent", side_effect=slow_create_agent): - await service.get_agent(session_id="s", user_id="u") - finally: - ticker.cancel() - - assert ticks >= 5 - - @pytest.mark.asyncio - async def test_build_stages_are_recorded_from_the_worker(self, mock_freshness_hash, monkeypatch): - from apis.shared.observability.build_stages import mark_stage - - monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") - recorded = [] - - def fake_create_agent(**kwargs): - mark_stage("session_mgr") - return _fake_agent() - - with patch.object(service, "create_agent", side_effect=fake_create_agent): - await service.get_agent(session_id="s", user_id="u", build_stage_recorder=recorded.append) - - assert recorded == ["session_mgr"] - - @pytest.mark.asyncio - async def test_the_control_arm_builds_on_the_loop(self, mock_freshness_hash, monkeypatch): - """The default: no experiment set means the build stays where it was.""" - import threading - - monkeypatch.delenv("AGENT_BUILD_EXPERIMENT", raising=False) - loop_thread = threading.get_ident() - build_threads = [] - - def fake_create_agent(**kwargs): - build_threads.append(threading.get_ident()) - return _fake_agent() - - with patch.object(service, "create_agent", side_effect=fake_create_agent): - await service.get_agent(session_id="s", user_id="u") - - assert build_threads == [loop_thread] - - @pytest.mark.asyncio - async def test_builds_never_overlap(self, mock_freshness_hash, monkeypatch): - """The external MCP layer is only safe one build at a time, which the - frozen loop used to guarantee for free.""" - import asyncio - import threading - import time - - monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") - active, peak = 0, 0 - guard = threading.Lock() - - def slow_create_agent(**kwargs): - nonlocal active, peak - with guard: - active += 1 - peak = max(peak, active) - time.sleep(0.05) - with guard: - active -= 1 - return _fake_agent() - - with patch.object(service, "create_agent", side_effect=slow_create_agent): - await asyncio.gather(*(service.get_agent(session_id=f"s{i}", user_id="u") for i in range(4))) - - assert peak == 1 - - @pytest.mark.asyncio - async def test_a_cancelled_build_still_holds_the_lock_until_its_thread_ends( - self, mock_freshness_hash, monkeypatch - ): - """A client disconnect cancels the await, not the thread.""" - import asyncio - import threading - - monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") - first_started, release_first = threading.Event(), threading.Event() - order = [] - - def create_agent(**kwargs): - if kwargs["session_id"] == "first": - first_started.set() - release_first.wait(5) - order.append("first-finished") - else: - order.append("second-started") - return _fake_agent() - - with patch.object(service, "create_agent", side_effect=create_agent): - first = asyncio.create_task(service.get_agent(session_id="first", user_id="u")) - await asyncio.to_thread(first_started.wait, 5) - first.cancel() - second = asyncio.create_task(service.get_agent(session_id="second", user_id="u")) - await asyncio.sleep(0.05) - assert order == [] - release_first.set() - await second - - assert order == ["first-finished", "second-started"] - assert first.cancelled() diff --git a/backend/tests/apis/inference_api/test_preparing_phase.py b/backend/tests/apis/inference_api/test_preparing_phase.py index 0f01fc4f1..d2817f157 100644 --- a/backend/tests/apis/inference_api/test_preparing_phase.py +++ b/backend/tests/apis/inference_api/test_preparing_phase.py @@ -251,78 +251,3 @@ def test_the_route_no_longer_races_its_own_build(self): source = Path(routes_module.__file__).read_text() assert "_PREPARING_NOTICE_SECONDS" not in source - - -class TestTitleDuringTheBuild: - """A first turn's title can land while the agent is still being built. - - The route races the deferred build against the title task and emits a - finished title before `prepared`. That only helps when the build runs off - the event loop (`agent_build_off_loop_enabled`): a synchronous build on the - loop stops the title task from finishing first, so the race reduces to the - plain await it replaced, which is the control arm's behaviour. - """ - - @staticmethod - async def _stream(build, title_task) -> List[str]: - """Mirrors the route: preparing, race, prepared.""" - frames = ['event: agent_status\ndata: {"phase": "preparing"}\n\n'] - emitted = False - - def title_sse() -> Optional[str]: - nonlocal emitted - if emitted or not title_task.done(): - return None - emitted = True - return f'event: session_title\ndata: {{"title": "{title_task.result()}"}}\n\n' - - task = asyncio.ensure_future(build()) - if not title_task.done(): - await asyncio.wait({task, title_task}, return_when=asyncio.FIRST_COMPLETED) - frame = title_sse() - if frame: - frames.append(frame) - await task - frames.append('event: agent_status\ndata: {"phase": "prepared"}\n\n') - return frames - - @staticmethod - def _kinds(frames: List[str]) -> List[str]: - return [f.split("\n", 1)[0].removeprefix("event: ") for f in frames] - - @pytest.mark.asyncio - async def test_a_build_off_the_loop_lets_the_title_out_first(self): - import time - - async def title() -> str: - await asyncio.sleep(0.02) - return "Biology Syllabus" - - async def build_in_a_thread() -> object: - return await asyncio.to_thread(time.sleep, 0.2) - - frames = await self._stream(build_in_a_thread, asyncio.ensure_future(title())) - - assert self._kinds(frames) == ["agent_status", "session_title", "agent_status"] - - @pytest.mark.asyncio - async def test_a_build_on_the_loop_cannot(self): - import time - - async def title() -> str: - await asyncio.sleep(0.02) - return "Biology Syllabus" - - async def build_on_the_loop() -> object: - time.sleep(0.2) # synchronous, as `create_agent` is - return object() - - frames = await self._stream(build_on_the_loop, asyncio.ensure_future(title())) - - assert "session_title" not in self._kinds(frames) - - def test_route_still_races_the_build_against_the_title(self): - source = Path(routes_module.__file__).read_text() - - assert "{build, title_task}" in source - assert "return_when=asyncio.FIRST_COMPLETED" in source diff --git a/backend/tests/shared/test_agent_build_experiment_arm.py b/backend/tests/shared/test_agent_build_experiment_arm.py index adb1f64e3..4fa382bca 100644 --- a/backend/tests/shared/test_agent_build_experiment_arm.py +++ b/backend/tests/shared/test_agent_build_experiment_arm.py @@ -7,7 +7,6 @@ from apis.shared.feature_flags import ( AGENT_BUILD_ARMS, agent_build_experiment_arm, - agent_build_off_loop_enabled, memory_shared_clients_enabled, ) @@ -21,7 +20,6 @@ def test_everything_but_ab_or_an_arm_name_is_control(monkeypatch, value): assert agent_build_experiment_arm("s1") == "control" assert not memory_shared_clients_enabled("s1") - assert not agent_build_off_loop_enabled("s1") @pytest.mark.parametrize("arm", AGENT_BUILD_ARMS) @@ -30,12 +28,20 @@ def test_an_arm_name_forces_that_arm(monkeypatch, arm): assert agent_build_experiment_arm("s1") == arm -def test_the_arms_imply_their_changes(monkeypatch): +def test_the_shared_arm_implies_its_change(monkeypatch): monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients") - assert memory_shared_clients_enabled("s") and not agent_build_off_loop_enabled("s") + assert memory_shared_clients_enabled("s") + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "control") + assert not memory_shared_clients_enabled("s") + + +def test_the_withdrawn_off_loop_arm_is_not_an_arm(monkeypatch): + """`shared_clients_off_loop` was withdrawn before the A/B; naming it now + means control, like any other unrecognised value.""" + assert AGENT_BUILD_ARMS == ("control", "shared_clients") monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients_off_loop") - assert memory_shared_clients_enabled("s") and agent_build_off_loop_enabled("s") + assert agent_build_experiment_arm("s") == "control" def test_ab_is_stable_per_session_and_uses_every_arm(monkeypatch): @@ -46,7 +52,7 @@ def test_ab_is_stable_per_session_and_uses_every_arm(monkeypatch): assert arms == [agent_build_experiment_arm(s) for s in sessions] counts = {arm: arms.count(arm) for arm in AGENT_BUILD_ARMS} - assert all(60 <= n <= 140 for n in counts.values()), counts + assert all(110 <= n <= 190 for n in counts.values()), counts def test_ab_without_a_session_is_control(monkeypatch): diff --git a/docs/specs/turn-latency-preamble.md b/docs/specs/turn-latency-preamble.md index fb259a9a9..b0c18f1aa 100644 --- a/docs/specs/turn-latency-preamble.md +++ b/docs/specs/turn-latency-preamble.md @@ -748,7 +748,7 @@ at. Second target after that: `agent_build.session_mgr` at 830ms (AgentCore Memory restore, never timed). Now PR-6. -## PR-6 — agent-build A/B: shared Memory clients, and the build off the loop (IN PROGRESS) +## PR-6 — agent-build A/B: shared boto3 clients, built at warm-up (IN PROGRESS) The second target named above, `agent_build.session_mgr`, measured **616-738ms** on dev first turns (2026-09-28) — more than half the 1.0-1.2s build. @@ -760,34 +760,49 @@ fresh session plus two more clients that **replace** the first pair, then calls second is a legacy-format fallback) and a `create_event`, each on a cold connection pool. Laptop timing: ~360ms of client construction per session manager. PR-3 above is the warning about reading that number: construction was -~30x slower on the container than on a laptop. +~30x slower on the container than on a laptop. Two more fresh sessions sit next +to it on the same first turn: `_discover_strategy_ids` builds its own +`MemoryClient`, and Strands' `BedrockModel` builds its own `boto3.Session`. **Every first turn is a fresh process.** Each conversation's turns run in their own Runtime process (`service.instance.id` differs per session in the runtime logs). So a first turn always pays cold construction, and "a build freezes other users' streams" does not happen: a process serves one conversation. Process-wide -caching helps a first turn only by doing less cold work (one session loads the -service models once instead of twice); it helps later cache-miss builds in the -same conversation (an `@`-mention, a changed toolset) fully. - -**Two changes, behind one per-session experiment flag, default off** +caching helps a first turn only if the shared work is done **before** the turn +arrives, which is why the shared session is built at container warm-up +(`apis/inference_api/warmup.py`, on the startup daemon thread): its +`bedrock-agentcore`, `bedrock-agentcore-control` and `bedrock-runtime` clients +are constructed there, and `_discover_strategy_ids` is called once so the +strategy ids are cached before any turn. That discovery is a control-plane read +of static configuration, and it opens the one connection warm-up otherwise +avoids; `docs/specs/turn-path-ttft.md` §5 P2 accepts that (botocore retries a +connection error on this idempotent call) and names it as the thing to watch on +the Runtime V2 restore. + +**One change, behind one per-session experiment flag, default off** (`agent_build_experiment_arm` in `apis/shared/feature_flags.py`, `AGENT_BUILD_EXPERIMENT`): | Arm | Change | |---|---| | `control` | today's build | -| `shared_clients` | the factory hands the SDK one process-wide boto3 session whose `client()` returns one client per configuration, and rebinds the SDK module's `MemoryClient` so the discarded pair is built from it too | -| `shared_clients_off_loop` | that, plus `create_agent` runs under `asyncio.to_thread`, one build at a time, and the route emits a title that lands mid-build | - -The off-loop arm needed hardening that the frozen loop used to provide for free: -builds stay serialized (the lock is released when the *thread* ends, so a -cancelled request cannot let a second build overlap an orphan), -`ExternalMCPIntegration` locks its maps and creates its singleton under a lock -(two instances would split the approval map, so a `needs_approval` tool could run -unapproved), and each build pins the external MCP clients it is handed until its -agent registers as their consumer (Strands stops a client whose last consumer -goes). +| `shared_clients` | the factory hands the SDK the process-wide boto3 session (whose `client()` returns one client per configuration) and rebinds the SDK module's `MemoryClient` so the discarded pair is built from it too; `_discover_strategy_ids` builds its `MemoryClient` on it; `ModelConfig.to_bedrock_config` passes it to `BedrockModel` as `boto_session` (and then no `region_name`, which Strands rejects alongside a session) | + +**The off-loop arm was withdrawn before the A/B.** The PR as opened had a third +arm, `shared_clients_off_loop`, which ran `create_agent` under +`asyncio.to_thread` so the route could emit a first-turn title mid-build. It +came with hardening the frozen loop used to provide for free: process-wide +build serialization, locks on `ExternalMCPIntegration`'s maps and singleton, +and a consumer pin on every external MCP client handed to a build. It was +dropped, and the hardening with it, on the assessment in +`docs/specs/turn-path-ttft.md` §4: moving a synchronous build to a thread does +not make it faster (CPU-bound parts contend on the GIL, network-bound parts take +the same time), so the arm could not reduce time to first token and its A/B +would have read as noise; and the concurrency it would have bought — running +the build's independent IO at the same time — is reachable inside the +constructor itself (§5 P3a, two threads of one executor for `session_mgr` and +`tools`) without any of the hardening. A title that lands before `prepared` is +not worth shipping locks for. **New sub-stages.** `agent_build.session_mgr_clients` (everything before the SDK's `read_session`, i.e. client setup) now precedes `agent_build.session_mgr` (the @@ -811,9 +826,7 @@ arm per round in rotating order, so time-of-day drift hits every arm alike. **Decision rule, fixed before the data.** Ship `shared_clients` (default on, with a kill switch) if its median `agent_build.session_mgr_clients` + `agent_build.session_mgr` beats control's by more than the run-to-run spread -and no stage regresses. Ship the off-loop build only if the title demonstrably -lands before `prepared` and time to first token does not regress; otherwise -remove it and the MCP hardening that exists only for it. +and no stage regresses. Otherwise remove the arm and keep the instrumentation. ## Declined / overtaken — `asyncio.to_thread` for the DynamoDB calls From 628da6b6e0b19ff2882c9cdd5fc5a3adf5e04704 Mon Sep 17 00:00:00 2001 From: Phil Merrell Date: Wed, 30 Sep 2026 08:23:03 -0600 Subject: [PATCH 3/4] perf(backend): build the shared boto3 session at warm-up and hand it to every SDK client The shared session was created lazily by the first build, so the first turn of a conversation, the only turn the experiment is about, still parsed two service models and made the strategy-id call. Now: - `ClientReusingSession` and `shared_boto_session()` live in `apis.shared.aws_clients` beside the process client cache; `reset_cached_clients` drops the session too (the moto trap). A keyword passed as None (Strands passes `endpoint_url=None`) no longer defeats the sharing. - `warmup.warm_shared_session` builds `bedrock-agentcore`, `bedrock-agentcore-control` and `bedrock-runtime` clients on it and calls `warm_strategy_ids` once, on the startup daemon thread. That discovery is the one connection warm-up otherwise avoids; accepted per docs/specs/turn-path-ttft.md section 5 P2, with the V2 restore named as the thing to watch. - `_discover_strategy_ids` takes `shared_session=` (part of its cache key, so warm-up primes the shared entry and the control arm's first turn still does exactly what it did before). - `ModelConfig.to_bedrock_config(session_id=)` passes `boto_session` to `BedrockModel` on the shared arm, and never `region_name` beside it. Nothing here runs on the control arm's request path, and nothing reaches the prompt, `toolConfig` or restored history. Co-Authored-By: Claude Fable 5.1 --- backend/src/agents/main_agent/chat_agent.py | 1 + .../agents/main_agent/core/agent_factory.py | 11 +- .../agents/main_agent/core/model_config.py | 22 +++- .../main_agent/session/session_factory.py | 118 +++++++----------- backend/src/apis/inference_api/warmup.py | 42 +++++++ backend/src/apis/shared/aws_clients.py | 94 +++++++++++++- .../test_session_factory_shared_clients.py | 20 +-- 7 files changed, 224 insertions(+), 84 deletions(-) diff --git a/backend/src/agents/main_agent/chat_agent.py b/backend/src/agents/main_agent/chat_agent.py index 4b71e045b..08d651aae 100644 --- a/backend/src/agents/main_agent/chat_agent.py +++ b/backend/src/agents/main_agent/chat_agent.py @@ -105,6 +105,7 @@ def _create_agent(self) -> None: hooks=hooks, plugins=plugins, memory_context=getattr(self, "memory_context", None), + session_id=getattr(self, "session_id", None), ) except Exception as e: diff --git a/backend/src/agents/main_agent/core/agent_factory.py b/backend/src/agents/main_agent/core/agent_factory.py index b2cbf2dea..39f8abf46 100644 --- a/backend/src/agents/main_agent/core/agent_factory.py +++ b/backend/src/agents/main_agent/core/agent_factory.py @@ -24,18 +24,20 @@ class AgentFactory: """Factory for creating configured Strands Agent instances with multi-provider support""" @staticmethod - def _create_bedrock_model(model_config: ModelConfig) -> BedrockModel: + def _create_bedrock_model(model_config: ModelConfig, session_id: Optional[str] = None) -> BedrockModel: """ Create a BedrockModel instance Args: model_config: Model configuration + session_id: The conversation, which decides the agent-build A/B + arm (see ``ModelConfig.to_bedrock_config``). Returns: BedrockModel: Configured Bedrock model (a ``CountTokensBedrockModel`` so native CountTokens works for inference-profile model ids). """ - bedrock_config = model_config.to_bedrock_config() + bedrock_config = model_config.to_bedrock_config(session_id=session_id) # Strands awaits count_tokens before every model call; keep that local. # Native counts are taken off the critical path by the # context-attribution hook (native_count_tokens in a background task). @@ -190,6 +192,7 @@ def create_agent( hooks: Optional[List[Any]] = None, plugins: Optional[List[Any]] = None, memory_context: Optional[str] = None, + session_id: Optional[str] = None, ) -> Agent: """ Create a Strands Agent instance with the appropriate model provider @@ -206,6 +209,8 @@ def create_agent( memory_context: Optional rendered Memory-Space block. Sent after the system prompt, behind a cache point of its own when the model supports cache points (see below). + session_id: The conversation this agent serves. Only the Bedrock + provider reads it, to pick the agent-build A/B arm. Returns: Agent: Configured Strands Agent instance @@ -219,7 +224,7 @@ def create_agent( # Create appropriate model based on provider if provider == ModelProvider.BEDROCK: - model = AgentFactory._create_bedrock_model(model_config) + model = AgentFactory._create_bedrock_model(model_config, session_id=session_id) elif provider == ModelProvider.OPENAI: model = AgentFactory._create_openai_model(model_config) elif provider == ModelProvider.MANTLE: diff --git a/backend/src/agents/main_agent/core/model_config.py b/backend/src/agents/main_agent/core/model_config.py index 01c9815e2..7770c52d7 100644 --- a/backend/src/agents/main_agent/core/model_config.py +++ b/backend/src/agents/main_agent/core/model_config.py @@ -414,9 +414,27 @@ def bedrock_cache_points_supported(self) -> bool: and ("claude" in model_lower or "anthropic" in model_lower) ) - def to_bedrock_config(self) -> Dict[str, Any]: - """Convert to BedrockModel kwargs, translating canonical inference params.""" + def to_bedrock_config(self, session_id: Optional[str] = None) -> Dict[str, Any]: + """Convert to BedrockModel kwargs, translating canonical inference params. + + ``session_id`` decides the agent-build A/B arm + (``memory_shared_clients_enabled``): on the shared arm the model is + built on the process-wide boto3 session instead of the fresh + ``boto3.Session()`` Strands would otherwise construct (and re-parse + the bedrock-runtime model on). Nothing here reaches the prompt. + """ config: Dict[str, Any] = {"model_id": self.model_id} + + from apis.shared.feature_flags import memory_shared_clients_enabled + + if memory_shared_clients_enabled(session_id): + from apis.shared.aws_clients import shared_boto_session + + # Never alongside `region_name`: BedrockModel.__init__ raises when + # both are given (strands-agents 1.55.0). This config sets no + # region on either arm; the session resolves it from the + # environment exactly as Strands' own fresh session would. + config["boto_session"] = shared_boto_session() _apply_canonical_params( config, self.inference_params, _BEDROCK_PARAM_MAP, "bedrock", self.model_id ) diff --git a/backend/src/agents/main_agent/session/session_factory.py b/backend/src/agents/main_agent/session/session_factory.py index 1e69f9140..6176affc9 100644 --- a/backend/src/agents/main_agent/session/session_factory.py +++ b/backend/src/agents/main_agent/session/session_factory.py @@ -4,7 +4,6 @@ import contextvars import os import logging -import threading from typing import Optional, Any, Dict, Tuple from functools import lru_cache @@ -82,7 +81,9 @@ def session_async_persistence_enabled() -> bool: # # The SDK's session manager builds its clients from scratch on every # construction, twice (see ``memory_shared_clients_enabled``). Everything -# below exists to hand it one process-wide session instead. +# below exists to hand it the process-wide session from +# ``apis.shared.aws_clients`` instead, which warm-up builds at container +# start (``apis/inference_api/warmup.py``). # Set by the factory for the duration of one session manager's construction. # The SDK builds its ``MemoryClient`` with no session and no session id, so this @@ -91,62 +92,9 @@ def session_async_persistence_enabled() -> bool: "use_shared_memory_clients", default=False ) -# Sized for concurrent sessions sharing one pool. Async persistence writes each -# message through ``asyncio.to_thread``, so a busy container can have dozens of -# ``CreateEvent`` calls in flight on these clients at once, where each session -# used to have a pool of its own. Past the pool size urllib3 still serves the -# request, but discards the extra connection afterwards (and warns). -_SHARED_MEMORY_MAX_POOL_CONNECTIONS = 50 - - -def _client_cache_key(service_name: str, region_name: Optional[str], config: Any) -> Tuple[str, Optional[str], str]: - # ``Config`` is not hashable; its user-provided options are what makes two - # configs different (the SDK and ``MemoryClient`` differ only in user agent). - options = getattr(config, "_user_provided_options", None) or {} - return service_name, region_name, repr(sorted(options.items())) - if AGENTCORE_MEMORY_AVAILABLE: - import boto3 - from botocore.config import Config as _BotocoreConfig - - class _ClientReusingSession(boto3.Session): - """A boto3 session whose ``client()`` returns one client per configuration. - - boto3 clients are thread-safe; building them is not, and building one - is the expensive part, so creation happens once, under a lock. Calls - carrying anything beyond region and config (explicit credentials, an - endpoint override) are not ours to share and pass straight through. - """ - - def __init__(self, **kwargs: Any) -> None: - super().__init__(**kwargs) - self._shared_clients: Dict[Tuple[str, Optional[str], str], Any] = {} - self._shared_clients_lock = threading.Lock() - - def client(self, service_name: str, region_name: Optional[str] = None, config: Any = None, **kwargs: Any) -> Any: # type: ignore[override] - if kwargs: - return super().client(service_name, region_name=region_name, config=config, **kwargs) - key = _client_cache_key(service_name, region_name, config) - with self._shared_clients_lock: - client = self._shared_clients.get(key) - if client is None: - pooled = _BotocoreConfig(max_pool_connections=_SHARED_MEMORY_MAX_POOL_CONNECTIONS) - merged = pooled.merge(config) if config is not None else pooled - client = super().client(service_name, region_name=region_name, config=merged) - self._shared_clients[key] = client - return client - - _shared_memory_session: Optional[_ClientReusingSession] = None - _shared_memory_session_lock = threading.Lock() - - def shared_memory_boto_session() -> _ClientReusingSession: - global _shared_memory_session - if _shared_memory_session is None: - with _shared_memory_session_lock: - if _shared_memory_session is None: - _shared_memory_session = _ClientReusingSession() - return _shared_memory_session + from apis.shared.aws_clients import shared_boto_session class _SharedSessionMemoryClient(MemoryClient): """``MemoryClient`` built from the shared session when the flag is on. @@ -154,9 +102,9 @@ class _SharedSessionMemoryClient(MemoryClient): The SDK session manager constructs ``MemoryClient(region_name=...)`` with no session, so it cannot be handed one; its clients are then replaced a few lines later, so they were pure cost. Rebinding the name - in the SDK module is the only seam. Pinned by - ``test_session_factory_shared_clients.py``, which fails if an SDK - upgrade stops going through it. + in the SDK module is the only seam the pinned bedrock-agentcore offers. + Pinned by ``test_session_factory_shared_clients.py``, which fails if an + SDK upgrade stops going through it. """ def __init__( @@ -166,7 +114,7 @@ def __init__( boto3_session: Any = None, ) -> None: if boto3_session is None and _use_shared_memory_clients.get(): - boto3_session = shared_memory_boto_session() + boto3_session = shared_boto_session() super().__init__( region_name=region_name, integration_source=integration_source, @@ -176,19 +124,28 @@ def __init__( _sdk_session_manager.MemoryClient = _SharedSessionMemoryClient -@lru_cache(maxsize=1) -def _discover_strategy_ids(memory_id: str, region: str) -> Tuple[Optional[str], Optional[str], Optional[str]]: +@lru_cache(maxsize=2) +def _discover_strategy_ids( + memory_id: str, region: str, *, shared_session: bool = False +) -> Tuple[Optional[str], Optional[str], Optional[str]]: """ Discover the actual strategy IDs from the configured memory strategies. AgentCore Memory stores memories in strategy-specific namespaces: /strategies/{strategyId}/actors/{actorId} - This function queries the memory to find the actual strategy IDs. + This function queries the memory to find the actual strategy IDs. It is + a control-plane read of static configuration, cached for the life of the + process. Args: memory_id: AgentCore Memory ID region: AWS region + shared_session: Build the ``MemoryClient`` on the process-wide session + (``memory_shared_clients_enabled``) instead of a fresh one. Part + of the cache key on purpose: warm-up primes the shared entry at + container start, and the control arm's first turn must still do + exactly what it did before the experiment. Returns: Tuple of (semantic_strategy_id, preference_strategy_id, summary_strategy_id) @@ -197,7 +154,10 @@ def _discover_strategy_ids(memory_id: str, region: str) -> Tuple[Optional[str], return None, None, None try: - client = MemoryClient(region_name=region) + client = MemoryClient( + region_name=region, + boto3_session=shared_boto_session() if shared_session else None, + ) strategies = client.get_memory_strategies(memory_id=memory_id) semantic_id = None @@ -225,6 +185,20 @@ def _discover_strategy_ids(memory_id: str, region: str) -> Tuple[Optional[str], return None, None, None +def warm_strategy_ids() -> Tuple[Optional[str], Optional[str], Optional[str]]: + """Discover the memory's strategy ids on the shared session, once, at container start. + + Called from ``apis/inference_api/warmup.py`` on the startup daemon thread + so the shared arm's first turn finds the ids cached and its clients + built. Raises when no memory is configured (``load_memory_config``); the + warm-up step logs that and moves on. + """ + if not AGENTCORE_MEMORY_AVAILABLE: + return None, None, None + config = load_memory_config() + return _discover_strategy_ids(config.memory_id, config.region, shared_session=True) + + class SessionFactory: """Factory for creating appropriate session manager based on environment""" @@ -308,8 +282,15 @@ def _create_cloud_session_manager( logger.info(f" • Memory ID: {memory_id}") logger.info(f" • Region: {aws_region}") - # Discover actual strategy IDs from the memory configuration - semantic_id, preference_id, summary_id = _discover_strategy_ids(memory_id, aws_region) + # Discover actual strategy IDs from the memory configuration. On the + # shared arm this is a cache hit: warm-up discovered them at + # container start (`warm_strategy_ids`). + from apis.shared.feature_flags import memory_shared_clients_enabled + + shared_clients = memory_shared_clients_enabled(session_id) + semantic_id, preference_id, summary_id = _discover_strategy_ids( + memory_id, aws_region, shared_session=shared_clients + ) # Load retrieval thresholds from environment (configurable per deployment) relevance_score = float(os.environ.get(EnvVars.MEMORY_RELEVANCE_SCORE, str(Defaults.MEMORY_RELEVANCE_SCORE))) @@ -385,9 +366,6 @@ def _create_cloud_session_manager( compaction_config.token_threshold = compaction_threshold # Create session manager with compaction built-in - from apis.shared.feature_flags import memory_shared_clients_enabled - - shared_clients = memory_shared_clients_enabled(session_id) token = _use_shared_memory_clients.set(shared_clients) try: session_manager = TurnBasedSessionManager( @@ -396,7 +374,7 @@ def _create_cloud_session_manager( compaction_config=compaction_config if compaction_config.enabled else None, user_id=user_id, summarization_strategy_id=summary_id, - boto_session=shared_memory_boto_session() if shared_clients else None, + boto_session=shared_boto_session() if shared_clients else None, ) finally: _use_shared_memory_clients.reset(token) diff --git a/backend/src/apis/inference_api/warmup.py b/backend/src/apis/inference_api/warmup.py index 7e530a387..3b21bcdf2 100644 --- a/backend/src/apis/inference_api/warmup.py +++ b/backend/src/apis/inference_api/warmup.py @@ -53,6 +53,18 @@ "s3", ) +# The parse above is per *session*, and the two SDKs on the agent build (the +# AgentCore Memory session manager, Strands' BedrockModel) build their clients +# on the process-wide session from `apis.shared.aws_clients` when the +# `shared_clients` arm is on. Build those clients on it here, so the arm's +# first turn finds them parsed; on the control arm the session sits unused, +# which costs nothing on the request path. +WARM_SHARED_SESSION_SERVICES: tuple[str, ...] = ( + "bedrock-agentcore", + "bedrock-agentcore-control", + "bedrock-runtime", +) + def warmup_enabled() -> bool: """Whether startup warm-up runs. Empty/unset means on.""" @@ -87,12 +99,42 @@ def warm_boto_clients(services: Iterable[str] = WARM_BOTO_SERVICES) -> None: _timed(f"boto:{service}", lambda service=service: boto3.client(service, region_name=region)) +def warm_shared_session(services: Iterable[str] = WARM_SHARED_SESSION_SERVICES) -> None: + """Build the agent build's shared boto3 session, its clients, and the memory strategy ids. + + Everything here is client construction — no sockets — except the last + step. ``warm_strategy_ids`` is a control-plane read of the memory's + strategy ids (static configuration, cached for the life of the process), + and it opens the one connection warm-up otherwise avoids. That is + accepted (docs/specs/turn-path-ttft.md §5 P2): the alternative is paying + the read on the first turn, which is the turn this exists to shorten; the + call is idempotent, so botocore retries a connection error; and on + Runtime V2, where a warmed process is snapshotted and restored, the + restore's first turn is the thing to watch for a pool holding a socket + that did not survive. + """ + region = os.environ.get("AWS_REGION") or os.environ.get("AWS_DEFAULT_REGION") + if not region: + logger.info("warmup step=shared outcome=skipped error=no region configured") + return + from apis.shared.aws_clients import shared_boto_session + + session = shared_boto_session() + for service in services: + _timed(f"shared:{service}", lambda service=service: session.client(service, region_name=region)) + + from agents.main_agent.session.session_factory import warm_strategy_ids + + _timed("shared:strategy_ids", warm_strategy_ids) + + def run_warmup() -> None: """The whole warm-up, synchronously. Exposed for tests and for callers that want it inline.""" started = time.perf_counter() warm_modules() warm_boto_clients() + warm_shared_session() logger.info("warmup complete ms=%d", int((time.perf_counter() - started) * 1000)) diff --git a/backend/src/apis/shared/aws_clients.py b/backend/src/apis/shared/aws_clients.py index e77bd2870..7e4f4f9ae 100644 --- a/backend/src/apis/shared/aws_clients.py +++ b/backend/src/apis/shared/aws_clients.py @@ -19,6 +19,23 @@ calls, and credential refresh is handled inside the shared session. What is NOT safe is sharing one across a *moto* boundary — see below. +THE SHARED SESSION (`shared_boto_session`) +------------------------------------------ +The cache above serves callers that ask *this module* for a client. Two SDKs +on the first-turn agent build do not: the AgentCore Memory session manager +and Strands' `BedrockModel` each construct a fresh `boto3.Session` and build +their clients on it. A fresh session re-parses every service model it +touches (the parse is per session, not per process), so a cold first turn +paid that parse several times over — see `docs/specs/turn-path-ttft.md` +§5 P2. Both SDKs accept a session, so `shared_boto_session()` is the one +process-wide session handed to them (behind `memory_shared_clients_enabled`, +an A/B arm) and to `apis/inference_api/warmup.py`, which builds its clients +at container start so a first turn finds them already parsed. Its `client()` +returns one client per configuration, so every caller that asks for the same +service with the same config shares one client and one connection pool. +Reset with everything else: a session built under one moto backend is as +stale as a client built under it. + THE MOTO TRAP, AND WHY `reset_cached_clients` EXISTS ---------------------------------------------------- `moto.mock_aws()` is entered per test (`tests/*/conftest.py`). A client built @@ -39,6 +56,9 @@ import threading from typing import Any, Dict, Optional, Tuple +import boto3 +from botocore.config import Config as _BotocoreConfig + logger = logging.getLogger(__name__) # Keyed by (service_name, region). `region=None` means "whatever the ambient @@ -112,13 +132,85 @@ def get_dynamodb_table(table_name: str, region_name: Optional[str] = None) -> An return get_resource("dynamodb", region_name).Table(table_name) +# --------------------------------------------------------------------------- +# The shared session +# --------------------------------------------------------------------------- + +# Sized for concurrent sessions sharing one pool. Async persistence writes each +# message through ``asyncio.to_thread``, so a busy container can have dozens of +# ``CreateEvent`` calls in flight on these clients at once, where each session +# manager used to have a pool of its own. Past the pool size urllib3 still +# serves the request, but discards the extra connection afterwards (and warns). +SHARED_SESSION_MAX_POOL_CONNECTIONS = 50 + + +def _shared_client_key(service_name: str, region_name: Optional[str], config: Any) -> Tuple[str, Optional[str], str]: + # ``Config`` is not hashable; its user-provided options are what make two + # configs different (the SDKs differ from each other only in user agent). + options = getattr(config, "_user_provided_options", None) or {} + return service_name, region_name, repr(sorted(options.items())) + + +class ClientReusingSession(boto3.Session): + """A boto3 session whose ``client()`` returns one client per configuration. + + boto3 clients are thread-safe; building them is not, and building one is + the expensive part, so creation happens once, under a lock. Calls carrying + anything beyond region and config (explicit credentials, an endpoint + override) are not ours to share and pass straight through; a keyword + passed as ``None`` (Strands passes ``endpoint_url=None``) is not an + override and does not defeat the sharing. + """ + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + self._shared_clients: Dict[Tuple[str, Optional[str], str], Any] = {} + self._shared_clients_lock = threading.Lock() + + def client(self, service_name: str, region_name: Optional[str] = None, config: Any = None, **kwargs: Any) -> Any: # type: ignore[override] + overrides = {key: value for key, value in kwargs.items() if value is not None} + if overrides: + return super().client(service_name, region_name=region_name, config=config, **overrides) + key = _shared_client_key(service_name, region_name, config) + with self._shared_clients_lock: + client = self._shared_clients.get(key) + if client is None: + pooled = _BotocoreConfig(max_pool_connections=SHARED_SESSION_MAX_POOL_CONNECTIONS) + merged = pooled.merge(config) if config is not None else pooled + client = super().client(service_name, region_name=region_name, config=merged) + self._shared_clients[key] = client + return client + + +_shared_session: Optional[ClientReusingSession] = None + + +def shared_boto_session() -> ClientReusingSession: + """The process-wide session the agent build's SDK clients are built on. + + Built by ``apis/inference_api/warmup.py`` at container start, so a first + turn that reaches for it finds its service models already parsed; built + here on first use otherwise. + """ + global _shared_session + session = _shared_session + if session is not None: + return session + with _lock: + if _shared_session is None: + _shared_session = ClientReusingSession() + return _shared_session + + def reset_cached_clients() -> None: - """Drop every cached client and resource. + """Drop every cached client and resource, and the shared session. Exists for tests. `moto.mock_aws()` is entered per test, so a client built under one test's mock must never be reused by the next — see the module docstring. Production code has no reason to call this. """ + global _shared_session with _lock: _resources.clear() _clients.clear() + _shared_session = None diff --git a/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py b/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py index 62f8f2a5e..5e580e635 100644 --- a/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py +++ b/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py @@ -6,7 +6,8 @@ the event loop, plus a cold connection pool. These tests build REAL ``TurnBasedSessionManager`` instances through the factory, with only the SDK's network calls stubbed, so they fail if an SDK upgrade stops going through the -seams the factory relies on. +seams the factory relies on. The session itself lives in +``apis.shared.aws_clients`` (``shared_boto_session``), where warm-up builds it. """ import threading @@ -17,6 +18,7 @@ import pytest from agents.main_agent.session import session_factory as factory +from apis.shared import aws_clients from bedrock_agentcore.memory.integrations.strands import session_manager as sdk @@ -26,11 +28,13 @@ def memory_env(monkeypatch): monkeypatch.setenv("AGENTCORE_MEMORY_ID", "mem-test") monkeypatch.setenv("AWS_REGION", "us-west-2") monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients") - monkeypatch.setattr(factory, "_discover_strategy_ids", lambda memory_id, region: (None, None, None)) + monkeypatch.setattr(factory, "_discover_strategy_ids", lambda memory_id, region, **kwargs: (None, None, None)) monkeypatch.setattr(sdk.AgentCoreMemorySessionManager, "read_session", lambda self, session_id, **k: None) monkeypatch.setattr(sdk.AgentCoreMemorySessionManager, "create_session", lambda self, session, **k: session) # A fresh shared session per test, so client counts start from zero. - monkeypatch.setattr(factory, "_shared_memory_session", None) + aws_clients.reset_cached_clients() + yield + aws_clients.reset_cached_clients() @pytest.fixture @@ -79,11 +83,11 @@ def test_shared_clients_get_a_bigger_pool_and_keep_the_sdk_user_agent(self, memo manager = _build("session-a") config = manager.memory_client.gmdp_client.meta.config - assert config.max_pool_connections == factory._SHARED_MEMORY_MAX_POOL_CONNECTIONS + assert config.max_pool_connections == aws_clients.SHARED_SESSION_MAX_POOL_CONNECTIONS assert "strands-agents" in (config.user_agent_extra or "") def test_concurrent_first_builds_share_one_client(self, memory_env, client_builds): - """Builds run in worker threads; the first ones race to create the clients.""" + """Two builds racing for the first client must still end up with one.""" barrier = threading.Barrier(4) def build(i: int) -> Any: @@ -112,14 +116,14 @@ def test_the_memory_client_seam_is_inert_outside_a_shared_construction(self, mem """The rebinding in the SDK module must not change a MemoryClient built anywhere else (e.g. `_discover_strategy_ids`).""" client = sdk.MemoryClient(region_name="us-west-2") - assert client.gmdp_client is not factory.shared_memory_boto_session().client( + assert client.gmdp_client is not aws_clients.shared_boto_session().client( "bedrock-agentcore", region_name="us-west-2" ) class TestClientReusingSession: def test_explicit_credentials_are_never_shared(self): - session = factory._ClientReusingSession() + session = aws_clients.ClientReusingSession() kwargs = dict(region_name="us-west-2", aws_access_key_id="a", aws_secret_access_key="b") assert session.client("sts", **kwargs) is not session.client("sts", **kwargs) @@ -127,7 +131,7 @@ def test_explicit_credentials_are_never_shared(self): def test_different_configs_get_different_clients(self): from botocore.config import Config - session = factory._ClientReusingSession() + session = aws_clients.ClientReusingSession() a = session.client("sts", region_name="us-west-2", config=Config(user_agent_extra="a")) b = session.client("sts", region_name="us-west-2", config=Config(user_agent_extra="b")) From 5802a29d491e8c7b56506628ab6f213e4fc5fdb5 Mon Sep 17 00:00:00 2001 From: Phil Merrell Date: Wed, 30 Sep 2026 08:24:48 -0600 Subject: [PATCH 4/4] test(backend): cover the warm-up shared session and the arm's SDK seams No network anywhere: warm-up builds each client once on a patched shared session and discovers the strategy ids once with `shared_session=True`; the factory hands the SDK constructor and `MemoryClient` the shared session on the arm and nothing off it; `BedrockModel` receives `boto_session` and no `region_name` on the arm (and two models share one bedrock-runtime client), `region_name`-free and session-free off it; `reset_cached_clients` drops the shared session; a `None` keyword (Strands' `endpoint_url=None`) does not defeat the sharing. Co-Authored-By: Claude Fable 5.1 --- .../core/test_bedrock_model_shared_session.py | 82 +++++++++++++++++++ .../test_session_factory_shared_clients.py | 76 +++++++++++++++++ .../tests/apis/inference_api/test_warmup.py | 82 ++++++++++++++++++- backend/tests/shared/test_aws_clients.py | 48 +++++++++++ 4 files changed, 287 insertions(+), 1 deletion(-) create mode 100644 backend/tests/agents/main_agent/core/test_bedrock_model_shared_session.py diff --git a/backend/tests/agents/main_agent/core/test_bedrock_model_shared_session.py b/backend/tests/agents/main_agent/core/test_bedrock_model_shared_session.py new file mode 100644 index 000000000..8f00b566b --- /dev/null +++ b/backend/tests/agents/main_agent/core/test_bedrock_model_shared_session.py @@ -0,0 +1,82 @@ +"""On the `shared_clients` arm, Strands' BedrockModel is built on the process-wide +boto3 session instead of the fresh `boto3.Session()` it would construct itself. + +`BedrockModel.__init__` raises when `region_name` and `boto_session` are both +given (strands-agents 1.55.0), so the arm must never send both. Real objects, +no network: constructing a client opens no socket. +""" + +import pytest + +from agents.main_agent.core.agent_factory import AgentFactory +from agents.main_agent.core.model_config import ModelConfig +from apis.shared import aws_clients + +MODEL_ID = "us.anthropic.claude-haiku-4-5-20251001-v1:0" + + +@pytest.fixture(autouse=True) +def _region_and_fresh_session(monkeypatch): + monkeypatch.setenv("AWS_REGION", "us-west-2") + aws_clients.reset_cached_clients() + yield + aws_clients.reset_cached_clients() + + +class TestOnTheArm: + @pytest.fixture(autouse=True) + def _arm(self, monkeypatch): + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "shared_clients") + + def test_config_carries_the_shared_session_and_no_region(self): + config = ModelConfig(model_id=MODEL_ID).to_bedrock_config(session_id="s") + + assert config["boto_session"] is aws_clients.shared_boto_session() + assert "region_name" not in config + + def test_the_model_client_is_the_shared_sessions_client(self): + first = AgentFactory._create_bedrock_model(ModelConfig(model_id=MODEL_ID), session_id="s") + second = AgentFactory._create_bedrock_model(ModelConfig(model_id=MODEL_ID), session_id="s") + + assert first.client is second.client, "two models, one bedrock-runtime client" + assert first.client.meta.config.max_pool_connections == aws_clients.SHARED_SESSION_MAX_POOL_CONNECTIONS + + def test_the_arm_resolves_the_same_region_as_a_fresh_session(self, monkeypatch): + """Strands resolves the region from the session it is given, and a + fresh `boto3.Session()` reads the same environment the shared one + does, so the arm must land on the same region as control.""" + shared = AgentFactory._create_bedrock_model(ModelConfig(model_id=MODEL_ID), session_id="s") + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", "control") + control = AgentFactory._create_bedrock_model(ModelConfig(model_id=MODEL_ID), session_id="s") + + assert shared.client.meta.region_name == control.client.meta.region_name + + def test_a_deliberate_region_next_to_the_session_is_still_refused(self): + """Pins the Strands contract the arm is written around.""" + from strands.models import BedrockModel + + with pytest.raises(ValueError, match="both"): + BedrockModel(model_id=MODEL_ID, region_name="us-west-2", boto_session=aws_clients.shared_boto_session()) + + +class TestOffTheArm: + @pytest.mark.parametrize("value", [None, "", "control", "ab"]) + def test_config_carries_no_session(self, monkeypatch, value): + if value is None: + monkeypatch.delenv("AGENT_BUILD_EXPERIMENT", raising=False) + else: + monkeypatch.setenv("AGENT_BUILD_EXPERIMENT", value) + + config = ModelConfig(model_id=MODEL_ID).to_bedrock_config(session_id=None) + + assert "boto_session" not in config + assert "region_name" not in config + + def test_the_model_builds_its_own_client(self, monkeypatch): + monkeypatch.delenv("AGENT_BUILD_EXPERIMENT", raising=False) + + first = AgentFactory._create_bedrock_model(ModelConfig(model_id=MODEL_ID)) + second = AgentFactory._create_bedrock_model(ModelConfig(model_id=MODEL_ID)) + + assert first.client is not second.client + assert aws_clients._shared_session is None, "control never touches the shared session" diff --git a/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py b/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py index 5e580e635..3d5ac910f 100644 --- a/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py +++ b/backend/tests/agents/main_agent/session/test_session_factory_shared_clients.py @@ -13,6 +13,7 @@ import threading from concurrent.futures import ThreadPoolExecutor from typing import Any, List +from unittest.mock import MagicMock, patch import boto3 import pytest @@ -137,3 +138,78 @@ def test_different_configs_get_different_clients(self): assert a is not b assert session.client("sts", region_name="us-west-2", config=Config(user_agent_extra="a")) is a + + +class TestWhatTheFactoryHandsTheSdk: + """The seam the arm rides on: the SDK's constructor takes ``boto_session``, + and ``MemoryClient`` takes ``boto3_session``. On the arm both get the + shared session; off it, nothing — the SDKs build their own, as before.""" + + def test_the_sdk_constructor_gets_the_shared_session_on_the_arm(self, memory_env, monkeypatch): + captured = {} + original = sdk.AgentCoreMemorySessionManager.__init__ + + def spy(self, *args, **kwargs): + captured["boto_session"] = kwargs.get("boto_session") + original(self, *args, **kwargs) + + monkeypatch.setattr(sdk.AgentCoreMemorySessionManager, "__init__", spy) + + _build("session-a") + + assert captured["boto_session"] is aws_clients.shared_boto_session() + + def test_the_sdk_constructor_gets_nothing_off_the_arm(self, memory_env, monkeypatch): + monkeypatch.delenv("AGENT_BUILD_EXPERIMENT", raising=False) + captured = {} + original = sdk.AgentCoreMemorySessionManager.__init__ + + def spy(self, *args, **kwargs): + captured["boto_session"] = kwargs.get("boto_session") + original(self, *args, **kwargs) + + monkeypatch.setattr(sdk.AgentCoreMemorySessionManager, "__init__", spy) + + _build("session-a") + + assert captured["boto_session"] is None + + def test_strategy_discovery_builds_its_client_on_the_shared_session_on_the_arm(self): + fetch = factory._discover_strategy_ids.__wrapped__ + with patch.object(factory, "MemoryClient") as memory_client: + memory_client.return_value.get_memory_strategies.return_value = [] + fetch("mem-test", "us-west-2", shared_session=True) + + memory_client.assert_called_once_with( + region_name="us-west-2", boto3_session=aws_clients.shared_boto_session() + ) + + def test_strategy_discovery_builds_a_fresh_client_off_the_arm(self): + fetch = factory._discover_strategy_ids.__wrapped__ + with patch.object(factory, "MemoryClient") as memory_client: + memory_client.return_value.get_memory_strategies.return_value = [] + fetch("mem-test", "us-west-2", shared_session=False) + + memory_client.assert_called_once_with(region_name="us-west-2", boto3_session=None) + + def test_the_factory_asks_for_the_shared_entry_only_on_the_arm(self, memory_env, monkeypatch): + asked: List[bool] = [] + monkeypatch.setattr( + factory, + "_discover_strategy_ids", + lambda memory_id, region, **kwargs: (asked.append(kwargs["shared_session"]), (None, None, None))[1], + ) + + _build("session-a") + monkeypatch.delenv("AGENT_BUILD_EXPERIMENT", raising=False) + _build("session-b") + + assert asked == [True, False] + + def test_warm_strategy_ids_primes_the_shared_entry(self, memory_env, monkeypatch): + """Warm-up's call and the arm's first turn must hit the same cache key, + or the first turn pays the call warm-up already made.""" + with patch.object(factory, "_discover_strategy_ids", return_value=(None, None, None)) as discover: + factory.warm_strategy_ids() + + discover.assert_called_once_with("mem-test", "us-west-2", shared_session=True) diff --git a/backend/tests/apis/inference_api/test_warmup.py b/backend/tests/apis/inference_api/test_warmup.py index c7a330517..f6b8c3d87 100644 --- a/backend/tests/apis/inference_api/test_warmup.py +++ b/backend/tests/apis/inference_api/test_warmup.py @@ -8,7 +8,7 @@ import sys import threading import types -from unittest.mock import patch +from unittest.mock import MagicMock, patch import pytest @@ -89,6 +89,86 @@ def flaky(service, **kwargs): assert calls == ["bedrock-runtime", "s3"] +class TestWarmSharedSession: + """The agent build's SDK clients are built on one process-wide session + (`apis.shared.aws_clients.shared_boto_session`). Building them here is + what lets the `shared_clients` arm's first turn skip the service-model + parses; discovering the strategy ids here is what lets it skip the one + control-plane call.""" + + def test_builds_each_client_once_on_the_shared_session(self, monkeypatch): + monkeypatch.setenv("AWS_REGION", "us-west-2") + session = MagicMock(name="shared-session") + + with patch("apis.shared.aws_clients.shared_boto_session", return_value=session), patch( + "agents.main_agent.session.session_factory.warm_strategy_ids" + ): + warmup.warm_shared_session(["bedrock-agentcore", "bedrock-runtime"]) + + assert [c.args[0] for c in session.client.call_args_list] == ["bedrock-agentcore", "bedrock-runtime"] + assert all(c.kwargs["region_name"] == "us-west-2" for c in session.client.call_args_list) + + def test_the_default_list_covers_both_sdks(self): + # The Memory session manager (data + control plane) and Strands' + # BedrockModel (runtime) are the two SDKs handed the shared session. + assert set(warmup.WARM_SHARED_SESSION_SERVICES) == { + "bedrock-agentcore", + "bedrock-agentcore-control", + "bedrock-runtime", + } + + def test_discovers_the_strategy_ids_once_on_the_shared_session(self, monkeypatch): + monkeypatch.setenv("AWS_REGION", "us-west-2") + monkeypatch.setenv("AGENTCORE_MEMORY_ID", "mem-warm") + from agents.main_agent.session import session_factory + + with patch("apis.shared.aws_clients.shared_boto_session", return_value=MagicMock()), patch.object( + session_factory, "_discover_strategy_ids", return_value=(None, None, None) + ) as discover: + warmup.warm_shared_session([]) + + discover.assert_called_once_with("mem-warm", "us-west-2", shared_session=True) + + def test_no_memory_configured_skips_discovery_and_does_not_raise(self, monkeypatch): + monkeypatch.setenv("AWS_REGION", "us-west-2") + monkeypatch.delenv("AGENTCORE_MEMORY_ID", raising=False) + from agents.main_agent.session import session_factory + + with patch("apis.shared.aws_clients.shared_boto_session", return_value=MagicMock()), patch.object( + session_factory, "_discover_strategy_ids" + ) as discover: + warmup.warm_shared_session([]) + + discover.assert_not_called() + + def test_skips_without_a_region(self, monkeypatch): + monkeypatch.delenv("AWS_REGION", raising=False) + monkeypatch.delenv("AWS_DEFAULT_REGION", raising=False) + with patch("apis.shared.aws_clients.shared_boto_session") as session: + warmup.warm_shared_session() + session.assert_not_called() + + def test_a_failing_client_does_not_stop_the_rest(self, monkeypatch): + monkeypatch.setenv("AWS_REGION", "us-west-2") + session = MagicMock() + session.client.side_effect = [RuntimeError("no such service"), object()] + + with patch("apis.shared.aws_clients.shared_boto_session", return_value=session), patch( + "agents.main_agent.session.session_factory.warm_strategy_ids" + ) as ids: + warmup.warm_shared_session(["bedrock-agentcore-control", "bedrock-runtime"]) + + assert session.client.call_count == 2 + ids.assert_called_once() + + def test_run_warmup_includes_it(self): + with patch.object(warmup, "warm_modules"), patch.object(warmup, "warm_boto_clients"), patch.object( + warmup, "warm_shared_session" + ) as shared: + warmup.run_warmup() + shared.assert_called_once_with() + + class TestBackground: def test_runs_on_a_daemon_thread_and_returns_immediately(self, monkeypatch): monkeypatch.delenv(warmup.WARMUP_ENABLED_ENV, raising=False) diff --git a/backend/tests/shared/test_aws_clients.py b/backend/tests/shared/test_aws_clients.py index eb5fb3a64..91dd321e9 100644 --- a/backend/tests/shared/test_aws_clients.py +++ b/backend/tests/shared/test_aws_clients.py @@ -69,6 +69,54 @@ def test_the_aws_fixture_leaves_no_client_behind(self): assert aws_clients._clients == {} +class TestSharedSession: + """One process-wide `boto3.Session` for the SDKs that build their own + clients (docs/specs/turn-path-ttft.md §5 P2).""" + + def test_the_same_session_every_time(self): + assert aws_clients.shared_boto_session() is aws_clients.shared_boto_session() + + def test_reset_drops_the_shared_session_too(self): + """A session built under one moto backend is as stale as a client + built under it.""" + before = aws_clients.shared_boto_session() + + aws_clients.reset_cached_clients() + + assert aws_clients._shared_session is None + assert aws_clients.shared_boto_session() is not before + + def test_one_client_per_service_region_and_config(self): + session = aws_clients.ClientReusingSession() + + first = session.client("sts", region_name="us-west-2") + again = session.client(service_name="sts", region_name="us-west-2") + other_region = session.client("sts", region_name="us-east-1") + + assert first is again + assert first is not other_region + assert first.meta.config.max_pool_connections == aws_clients.SHARED_SESSION_MAX_POOL_CONNECTIONS + + def test_a_none_keyword_is_not_an_override(self): + """Strands' BedrockModel passes `endpoint_url=None`; that must still + land on the shared client, or the arm silently builds a fresh one.""" + session = aws_clients.ClientReusingSession() + + plain = session.client("sts", region_name="us-west-2") + with_none = session.client("sts", region_name="us-west-2", endpoint_url=None) + + assert with_none is plain + + def test_a_real_override_is_never_shared(self): + session = aws_clients.ClientReusingSession() + + plain = session.client("sts", region_name="us-west-2") + custom = session.client("sts", region_name="us-west-2", endpoint_url="https://sts.example") + + assert custom is not plain + assert session.client("sts", region_name="us-west-2") is plain + + class TestMetadataUsesIt: @pytest.mark.asyncio async def test_session_reads_go_through_one_cached_resource(