diff --git a/backend/src/agents/main_agent/session/compaction_summary.py b/backend/src/agents/main_agent/session/compaction_summary.py index 5d7c51384..12a8f6e01 100644 --- a/backend/src/agents/main_agent/session/compaction_summary.py +++ b/backend/src/agents/main_agent/session/compaction_summary.py @@ -58,7 +58,9 @@ import logging import os from dataclasses import dataclass -from typing import List, Optional, Sequence, Tuple +from typing import Any, Dict, List, Optional, Sequence, Tuple + +from apis.shared.aws_clients import get_client from .compaction_policy import CHARS_PER_TOKEN @@ -212,6 +214,17 @@ def _salvage(text: str, budget_tokens: int) -> Optional[str]: return _keep_head_lines(head, budget_tokens) if head else None +def _converse(region: Optional[str], **kwargs: Any) -> Dict[str, Any]: + """Run in a worker thread, so a first-use client build stays off the event loop. + + The client is process-cached: one per call paid botocore's service-model + load (~250ms the first time in a process) and, every time, a fresh + connection pool — a new TCP+TLS handshake per compaction call. + """ + region = region or os.environ.get("AWS_REGION", "us-west-2") + return get_client("bedrock-runtime", region).converse(**kwargs) + + async def _compress( records: Sequence[str], budget_tokens: int, @@ -224,17 +237,12 @@ async def _compress( if not text.strip(): return None, False try: - import boto3 - except ImportError: # pragma: no cover - dev without boto3 - return None, False - try: - region = region or os.environ.get("AWS_REGION", "us-west-2") - client = boto3.client("bedrock-runtime", region_name=region) # ~0.75 words/token; aim well under the budget so the chars/4 check # below passes with margin. word_budget = max(150, int(budget_tokens * 0.55)) response = await asyncio.to_thread( - client.converse, + _converse, + region, modelId=model_id, system=[{"text": _COMPRESSION_SYSTEM_PROMPT.replace("{word_budget}", f"{word_budget:,}")}], messages=[{"role": "user", "content": [{"text": "Summary notes, oldest first:\n\n" + text}]}], @@ -304,14 +312,9 @@ async def extract_with_model( if not text.strip(): return None try: - import boto3 - except ImportError: # pragma: no cover - dev without boto3 - return None - try: - region = region or os.environ.get("AWS_REGION", "us-west-2") - client = boto3.client("bedrock-runtime", region_name=region) response = await asyncio.to_thread( - client.converse, + _converse, + region, modelId=model_id, system=[{"text": _EXTRACTION_SYSTEM_PROMPT}], messages=[{"role": "user", "content": [{"text": "Summary notes, oldest first:\n\n" + text}]}], diff --git a/backend/src/apis/shared/embeddings/bedrock_embeddings.py b/backend/src/apis/shared/embeddings/bedrock_embeddings.py index 159b4e649..cb5495694 100644 --- a/backend/src/apis/shared/embeddings/bedrock_embeddings.py +++ b/backend/src/apis/shared/embeddings/bedrock_embeddings.py @@ -12,6 +12,7 @@ import json import logging import os +import threading from typing import Any, Dict, List import boto3 @@ -33,6 +34,27 @@ logger = logging.getLogger(__name__) +# One Bedrock client for every embedding, built on first use. A client per call +# cost ~250ms the first time in a process (botocore loads the service model) and, +# every time, a fresh connection pool — a new TCP+TLS handshake on each knowledge +# base search, which sits in front of the model's first token. boto3 clients are +# thread-safe, so the executor workers below can share it. Held here rather than +# in `apis.shared.aws_clients` because the kb-sync and rag-ingestion Lambda +# images copy this package without that module. +_bedrock_runtime_client: Any = None +# Guards the first build: boto3's default session is not thread-safe, and +# parallel embeddings can reach it from separate worker threads at once. +_bedrock_runtime_client_lock = threading.Lock() + + +def _get_bedrock_runtime_client() -> Any: + global _bedrock_runtime_client + if _bedrock_runtime_client is None: + with _bedrock_runtime_client_lock: + if _bedrock_runtime_client is None: + _bedrock_runtime_client = boto3.client("bedrock-runtime", region_name=AWS_REGION) + return _bedrock_runtime_client + def _get_vector_store_bucket() -> str: """Get vector store bucket name, validating if not set""" @@ -74,18 +96,17 @@ async def generate_embeddings(chunks: List[str]) -> List[List[float]]: Raises: Exception: If Bedrock API call fails """ - bedrock_runtime = boto3.client("bedrock-runtime", region_name=AWS_REGION) - logger.info(f"Generating embeddings for {len(chunks)} chunks in parallel...") async def get_single_embedding(chunk: str, index: int) -> List[float]: """Generate embedding for a single chunk""" loop = asyncio.get_event_loop() - # Run synchronous boto3 call in thread pool to avoid blocking + # Run synchronous boto3 call in thread pool to avoid blocking; the + # client is fetched there too, so its first build stays off the loop. response = await loop.run_in_executor( None, - lambda: bedrock_runtime.invoke_model( + lambda: _get_bedrock_runtime_client().invoke_model( modelId=BEDROCK_EMBEDDING_CONFIG["model_id"], contentType="application/json", accept="application/json", diff --git a/backend/src/apis/shared/files/document_digest.py b/backend/src/apis/shared/files/document_digest.py index fdb7a78a5..79e4cad36 100644 --- a/backend/src/apis/shared/files/document_digest.py +++ b/backend/src/apis/shared/files/document_digest.py @@ -45,6 +45,8 @@ from pydantic import BaseModel, Field +from apis.shared.aws_clients import get_client + from .document_read import _open_pdf, _pdf_page_text, docx_paragraphs, document_format_for logger = logging.getLogger(__name__) @@ -257,18 +259,24 @@ def _abstract_prompt(outline: DocumentDigest, sample: str) -> str: return "\n".join(lines) +def _converse(region: str, **kwargs: Any) -> Dict[str, Any]: + """Run in a worker thread, so a first-use client build stays off the event loop. + + The client is process-cached: one per upload paid botocore's service-model + load (~250ms the first time in a process) and, every time, a fresh + connection pool — a new TCP+TLS handshake per abstract. + """ + return get_client("bedrock-runtime", region).converse(**kwargs) + + async def generate_abstract(outline: DocumentDigest, sample: str, model_id: str = DOCUMENT_DIGEST_MODEL_ID) -> Optional[str]: """3–5 sentences from the cheap model, or ``None`` (never raises).""" if not sample.strip(): return None try: - import boto3 - except ImportError: # pragma: no cover - return None - try: - client = boto3.client("bedrock-runtime", region_name=os.environ.get("AWS_REGION", "us-west-2")) response = await asyncio.to_thread( - client.converse, + _converse, + os.environ.get("AWS_REGION", "us-west-2"), modelId=model_id, messages=[{"role": "user", "content": [{"text": _abstract_prompt(outline, sample)}]}], system=[{"text": _ABSTRACT_SYSTEM_PROMPT}], diff --git a/backend/src/apis/shared/tool_summaries/summarizer.py b/backend/src/apis/shared/tool_summaries/summarizer.py index 34a37d448..777c1bd6c 100644 --- a/backend/src/apis/shared/tool_summaries/summarizer.py +++ b/backend/src/apis/shared/tool_summaries/summarizer.py @@ -35,6 +35,8 @@ import os from typing import Any, Dict, List, Optional +from apis.shared.aws_clients import get_client + logger = logging.getLogger(__name__) # Nova Micro: the cheapest Bedrock model that reliably follows a one-line @@ -140,6 +142,16 @@ def _clean(text: str) -> str: return summary.rstrip(".").strip() +def _converse(region: str, **kwargs: Any) -> Dict[str, Any]: + """Run in a worker thread, so a first-use client build stays off the event loop. + + The client is process-cached: one per batch paid botocore's service-model + load (~250ms the first time in a process) and, every time, a fresh + connection pool — a new TCP+TLS handshake per summary. + """ + return get_client("bedrock-runtime", region).converse(**kwargs) + + async def summarize_tool_batch(calls: List[Dict[str, Any]]) -> Optional[str]: """Summarize one finished tool batch, or return ``None``. @@ -157,19 +169,12 @@ async def summarize_tool_batch(calls: List[Dict[str, Any]]) -> Optional[str]: return None try: - import boto3 - except ImportError: # pragma: no cover - dev without boto3 - return None - - try: - region = os.environ.get("AWS_REGION", "us-west-2") - client = boto3.client("bedrock-runtime", region_name=region) - # boto3's converse() is synchronous — awaited inline it would block # the event loop for the whole Nova round trip, stalling the agent # stream this task runs concurrently with. response = await asyncio.to_thread( - client.converse, + _converse, + os.environ.get("AWS_REGION", "us-west-2"), modelId=_MODEL_ID, messages=[{"role": "user", "content": [{"text": _build_prompt(calls)}]}], system=[{"text": _SUMMARY_SYSTEM_PROMPT}], diff --git a/backend/tests/agents/main_agent/session/test_compaction_summary.py b/backend/tests/agents/main_agent/session/test_compaction_summary.py index 344f6bc89..12c278f3f 100644 --- a/backend/tests/agents/main_agent/session/test_compaction_summary.py +++ b/backend/tests/agents/main_agent/session/test_compaction_summary.py @@ -16,6 +16,7 @@ compress_with_model, truncate_records_newest_first, ) +from apis.shared import aws_clients from .conftest import make_conversation @@ -23,6 +24,14 @@ BUDGET = 100 # tokens → 400 chars +@pytest.fixture(autouse=True) +def _fresh_bedrock_client(): + """The Bedrock client is cached per process; each test builds its own.""" + aws_clients.reset_cached_clients() + yield + aws_clients.reset_cached_clients() + + @pytest.fixture def bedrock(monkeypatch): """Patch boto3 so no test reaches Bedrock; returns the converse mock.""" @@ -105,6 +114,19 @@ async def test_a_refused_generation_falls_back_and_is_never_the_summary(self, be assert "Sorry" not in result.text and result.text.startswith("new") assert await compress_with_model(records, BUDGET, model_id="m") is None + @pytest.mark.asyncio + async def test_calls_reuse_one_bedrock_client_per_region(self, bedrock): + """A client per call paid botocore's model load and a fresh TLS handshake.""" + bedrock.return_value = _model_reply("summary") + factory = sys.modules["boto3"].client + + await compress_with_model(["r" * 900], BUDGET, model_id="m", region="us-west-2") + await compress_with_model(["r" * 900], BUDGET, model_id="m", region="us-west-2") + await compress_with_model(["r" * 900], BUDGET, model_id="m", region="us-east-1") + + assert [c.kwargs["region_name"] for c in factory.call_args_list] == ["us-west-2", "us-east-1"] + assert bedrock.call_count == 3 + @pytest.mark.asyncio async def test_model_overshoot_is_tail_trimmed(self, bedrock): bedrock.return_value = _model_reply("y" * 2000 + "END") diff --git a/backend/tests/apis/shared/tool_summaries/test_summarizer.py b/backend/tests/apis/shared/tool_summaries/test_summarizer.py index 40e9f13c8..889b00556 100644 --- a/backend/tests/apis/shared/tool_summaries/test_summarizer.py +++ b/backend/tests/apis/shared/tool_summaries/test_summarizer.py @@ -16,6 +16,7 @@ import pytest +from apis.shared import aws_clients from apis.shared.tool_summaries.summarizer import ( _build_prompt, _clean, @@ -50,6 +51,14 @@ def summaries_enabled(monkeypatch): monkeypatch.setenv("TOOL_SUMMARIES_ENABLED", "true") +@pytest.fixture(autouse=True) +def _fresh_bedrock_client(): + """The Bedrock client is cached per process; each test builds its own.""" + aws_clients.reset_cached_clients() + yield + aws_clients.reset_cached_clients() + + @pytest.fixture def bedrock(monkeypatch): """Patch boto3.client so no test ever reaches Bedrock.""" @@ -60,6 +69,19 @@ def bedrock(monkeypatch): return client +@pytest.mark.asyncio +async def test_batches_reuse_one_bedrock_client(bedrock): + """A client per batch paid botocore's model load and a fresh TLS handshake.""" + bedrock.converse.return_value = _response("Found the BIO 101 course") + factory = __import__("sys").modules["boto3"].client + + for _ in range(3): + assert await summarize_tool_batch(_calls()) == "Found the BIO 101 course" + + assert factory.call_count == 1 + assert bedrock.converse.call_count == 3 + + # -- the truncation regression ------------------------------------------- diff --git a/backend/tests/shared/test_bedrock_embeddings_client.py b/backend/tests/shared/test_bedrock_embeddings_client.py new file mode 100644 index 000000000..03afbc5fb --- /dev/null +++ b/backend/tests/shared/test_bedrock_embeddings_client.py @@ -0,0 +1,36 @@ +"""generate_embeddings shares one bedrock-runtime client across calls.""" + +import io +import json +from unittest.mock import MagicMock, patch + +import pytest + +from apis.shared.embeddings import bedrock_embeddings as be + + +@pytest.fixture(autouse=True) +def _fresh_bedrock_client(monkeypatch): + """The Bedrock client is cached per process; each test builds its own.""" + monkeypatch.setattr(be, "_bedrock_runtime_client", None) + + +def _client() -> MagicMock: + client = MagicMock() + client.invoke_model.side_effect = lambda **kwargs: { + "body": io.BytesIO(json.dumps({"embedding": [0.1, 0.2]}).encode()) + } + return client + + +@pytest.mark.asyncio +async def test_calls_reuse_one_bedrock_client(): + """A client per call paid botocore's model load and a fresh TLS handshake on + every knowledge base search, ahead of the model's first token.""" + client = _client() + with patch.object(be.boto3, "client", return_value=client) as factory: + for _ in range(3): + assert await be.generate_embeddings(["a", "b"]) == [[0.1, 0.2], [0.1, 0.2]] + + factory.assert_called_once_with("bedrock-runtime", region_name=be.AWS_REGION) + assert client.invoke_model.call_count == 6 diff --git a/backend/tests/shared/test_document_digest.py b/backend/tests/shared/test_document_digest.py index 49b426af2..251637071 100644 --- a/backend/tests/shared/test_document_digest.py +++ b/backend/tests/shared/test_document_digest.py @@ -13,6 +13,7 @@ import pytest +from apis.shared import aws_clients from apis.shared.files import document_digest as dd from tests.shared.test_document_read import build_docx, build_pdf @@ -123,9 +124,19 @@ def _bedrock(monkeypatch, text="A crisp abstract.", stop="end_turn", fail=False) boto3 = MagicMock() boto3.client.return_value = client monkeypatch.setitem(__import__("sys").modules, "boto3", boto3) + # The client is cached per process; drop the previous fake so this one is built. + aws_clients.reset_cached_clients() return client +@pytest.fixture(autouse=True) +def _fresh_bedrock_client(): + """The Bedrock client is cached per process; each test builds its own.""" + aws_clients.reset_cached_clients() + yield + aws_clients.reset_cached_clients() + + class TestAbstract: @pytest.mark.asyncio async def test_calls_the_cheap_model_with_outline_and_sample(self, monkeypatch): @@ -153,6 +164,18 @@ async def test_a_refused_generation_is_none(self, monkeypatch, stop): _bedrock(monkeypatch, text="Sorry, the model cannot answer this.", stop=stop) assert await dd.generate_abstract(dd.DocumentDigest(), "text") is None + @pytest.mark.asyncio + async def test_abstracts_reuse_one_bedrock_client(self, monkeypatch): + """A client per upload paid botocore's model load and a fresh TLS handshake.""" + client = _bedrock(monkeypatch) + factory = __import__("sys").modules["boto3"].client + + for _ in range(3): + assert await dd.generate_abstract(dd.DocumentDigest(), "text") == "A crisp abstract." + + assert factory.call_count == 1 + assert client.converse.call_count == 3 + class TestBuild: @pytest.mark.asyncio diff --git a/backend/tests/shared/test_side_channel_inference_config.py b/backend/tests/shared/test_side_channel_inference_config.py index 5d241fe7c..5a940537e 100644 --- a/backend/tests/shared/test_side_channel_inference_config.py +++ b/backend/tests/shared/test_side_channel_inference_config.py @@ -13,6 +13,7 @@ import pytest import apis.inference_api.chat.service as chat_service +from apis.shared import aws_clients from apis.shared.files import document_digest as dd from apis.shared.tool_summaries.summarizer import summarize_tool_batch @@ -33,6 +34,16 @@ def _patch_boto3_module(monkeypatch, client: MagicMock) -> None: module = MagicMock() module.client.return_value = client monkeypatch.setitem(__import__("sys").modules, "boto3", module) + # The side channels share a process-cached client; drop the previous fake. + aws_clients.reset_cached_clients() + + +@pytest.fixture(autouse=True) +def _fresh_bedrock_client(): + """The Bedrock client is cached per process; each test builds its own.""" + aws_clients.reset_cached_clients() + yield + aws_clients.reset_cached_clients() @pytest.mark.asyncio