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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 18 additions & 15 deletions backend/src/agents/main_agent/session/compaction_summary.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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,
Expand All @@ -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}]}],
Expand Down Expand Up @@ -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}]}],
Expand Down
29 changes: 25 additions & 4 deletions backend/src/apis/shared/embeddings/bedrock_embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
import json
import logging
import os
import threading
from typing import Any, Dict, List

import boto3
Expand All @@ -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"""
Expand Down Expand Up @@ -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",
Expand Down
20 changes: 14 additions & 6 deletions backend/src/apis/shared/files/document_digest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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}],
Expand Down
23 changes: 14 additions & 9 deletions backend/src/apis/shared/tool_summaries/summarizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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``.

Expand All @@ -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}],
Expand Down
22 changes: 22 additions & 0 deletions backend/tests/agents/main_agent/session/test_compaction_summary.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,13 +16,22 @@
compress_with_model,
truncate_records_newest_first,
)
from apis.shared import aws_clients

from .conftest import make_conversation


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."""
Expand Down Expand Up @@ -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")
Expand Down
22 changes: 22 additions & 0 deletions backend/tests/apis/shared/tool_summaries/test_summarizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@

import pytest

from apis.shared import aws_clients
from apis.shared.tool_summaries.summarizer import (
_build_prompt,
_clean,
Expand Down Expand Up @@ -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."""
Expand All @@ -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 -------------------------------------------


Expand Down
36 changes: 36 additions & 0 deletions backend/tests/shared/test_bedrock_embeddings_client.py
Original file line number Diff line number Diff line change
@@ -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
23 changes: 23 additions & 0 deletions backend/tests/shared/test_document_digest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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
Expand Down
Loading