From 0342900ff17620a6a7600e3ab37bd6277525f3b5 Mon Sep 17 00:00:00 2001 From: Aakash Kattelu Date: Tue, 18 Aug 2026 10:37:36 -0400 Subject: [PATCH] fix(embedding): truncate on conclusion create paths and document BPE loop Storage callers in create_observations (API + agent tools) now pass on_oversize="truncate" so a single oversize item cannot drop the batch. Docstring on _truncate_to_token_limit notes why decode/re-encode is load-bearing. --- src/crud/document.py | 4 +- src/embedding_client.py | 5 ++- src/utils/agent_tools.py | 4 +- tests/crud/test_document.py | 40 ++++++++++++++++++++ tests/utils/test_agent_tools.py | 66 +++++++++++++++++++++++++++++++-- 5 files changed, 113 insertions(+), 6 deletions(-) diff --git a/src/crud/document.py b/src/crud/document.py index ba80712d..0cec85c9 100644 --- a/src/crud/document.py +++ b/src/crud/document.py @@ -984,7 +984,9 @@ async def create_observations( # Generate embeddings in batch contents = [obs.content for obs in observations] try: - embeddings = await embedding_client.simple_batch_embed(contents) + embeddings = await embedding_client.simple_batch_embed( + contents, on_oversize="truncate" + ) except ValueError as e: raise ValidationException(str(e)) from e diff --git a/src/embedding_client.py b/src/embedding_client.py index aeecae09..a380b40d 100644 --- a/src/embedding_client.py +++ b/src/embedding_client.py @@ -318,7 +318,10 @@ class _EmbeddingClient: ) def _truncate_to_token_limit(self, text: str) -> tuple[str, int]: - """Return a prefix of `text` whose re-encoded token count fits the cap.""" + """Return a prefix of `text` whose re-encoded token count fits the cap. + + Decode/re-encode after slicing: BPE boundaries can re-expand past the cap. + """ token_ids = self.encoding.encode(text) keep = self.max_embedding_tokens while len(token_ids) > self.max_embedding_tokens: diff --git a/src/utils/agent_tools.py b/src/utils/agent_tools.py index 9f214009..7fc79f62 100644 --- a/src/utils/agent_tools.py +++ b/src/utils/agent_tools.py @@ -930,7 +930,9 @@ async def create_observations( run_id=run_id, parent_category=parent_category, ): - embeddings = await embedding_client.simple_batch_embed(contents) + embeddings = await embedding_client.simple_batch_embed( + contents, on_oversize="truncate" + ) embeddings_by_index = dict( zip(range(len(normalized_observations)), embeddings, strict=True) ) diff --git a/tests/crud/test_document.py b/tests/crud/test_document.py index 5f03c80b..6686e688 100644 --- a/tests/crud/test_document.py +++ b/tests/crud/test_document.py @@ -1006,6 +1006,46 @@ class TestDocumentCRUD: assert documents[0].content in ["Observation 1", "Observation 2"] assert documents[1].content in ["Observation 1", "Observation 2"] + @pytest.mark.asyncio + async def test_create_observations_embeds_with_truncate_on_oversize( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """API conclusion creates must opt into truncation on oversize content.""" + test_workspace, test_peer = sample_data + test_peer2, test_session, _ = await self._setup_test_data( + db_session, test_workspace, test_peer + ) + + with patch( + "src.crud.document.embedding_client.simple_batch_embed", + new=AsyncMock(return_value=[[0.1] * 1536, [0.2] * 1536]), + ) as mock_embed: + created = await crud.create_observations( + db_session, + observations=[ + schemas.ConclusionCreate( + content="short conclusion", + observer_id=test_peer.name, + observed_id=test_peer2.name, + session_id=test_session.name, + ), + schemas.ConclusionCreate( + content="another conclusion", + observer_id=test_peer.name, + observed_id=test_peer2.name, + session_id=test_session.name, + ), + ], + workspace_name=test_workspace.name, + ) + + assert len(created) == 2 + mock_embed.assert_awaited_once_with( + ["short conclusion", "another conclusion"], on_oversize="truncate" + ) + class TestSessionPurityInvariant: """Regression tests for the explicit-document session-purity invariant. diff --git a/tests/utils/test_agent_tools.py b/tests/utils/test_agent_tools.py index c13fb1e7..5dd66a5b 100644 --- a/tests/utils/test_agent_tools.py +++ b/tests/utils/test_agent_tools.py @@ -334,7 +334,9 @@ class TestCreateObservations: """If batch embedding fails but individual embeds succeed, all observations are created.""" workspace, peer1, peer2, session, _, _ = tool_test_data - async def fail_batch_embed(_texts: list[str]) -> list[list[float]]: + async def fail_batch_embed( + _texts: list[str], **_kwargs: object + ) -> list[list[float]]: raise RuntimeError("embedding provider timeout") async def succeed_single_embed(_content: str) -> list[float]: @@ -393,7 +395,9 @@ class TestCreateObservations: """If batch embedding fails and some individual embeds also fail, only successful ones are created.""" workspace, peer1, peer2, session, _, _ = tool_test_data - async def fail_batch_embed(_texts: list[str]) -> list[list[float]]: + async def fail_batch_embed( + _texts: list[str], **_kwargs: object + ) -> list[list[float]]: raise RuntimeError("embedding provider timeout") async def embed_per_observation(content: str) -> list[float]: @@ -458,7 +462,9 @@ class TestCreateObservations: workspace, peer1, peer2, session, _, _ = tool_test_data created_documents: list[Any] = [] - async def fake_batch_embed(texts: list[str]) -> list[list[float]]: + async def fake_batch_embed( + texts: list[str], **_kwargs: object + ) -> list[list[float]]: assert texts == ["trimmed observation"] return [[0.4, 0.5, 0.6]] @@ -504,6 +510,60 @@ class TestCreateObservations: assert len(created_documents) == 1 assert created_documents[0].content == "trimmed observation" + async def test_create_observations_embeds_with_truncate_on_oversize( + self, + tool_test_data: Any, + monkeypatch: pytest.MonkeyPatch, + ): + """Storage path must opt into truncation so one long obs cannot drop the batch.""" + workspace, peer1, peer2, session, _, _ = tool_test_data + captured: dict[str, object] = {} + + async def fake_batch_embed( + texts: list[str], *, on_oversize: str = "raise", **_kwargs: object + ) -> list[list[float]]: + captured["texts"] = texts + captured["on_oversize"] = on_oversize + return [[0.1] for _ in texts] + + async def fake_create_documents( + _db: AsyncSession, + documents: list[Any], + workspace_name: str, + *, + observer: str, + observed: str, + deduplicate: bool = False, + ) -> crud.CreateDocumentsResult: + _ = (workspace_name, observer, observed, deduplicate) + return crud.CreateDocumentsResult(created_documents=documents) + + monkeypatch.setattr( + "src.utils.agent_tools.embedding_client.simple_batch_embed", + fake_batch_embed, + ) + monkeypatch.setattr( + "src.utils.agent_tools.crud.create_documents", fake_create_documents + ) + + result = await create_observations( + observations=[ + schemas.ObservationInput(content="short fact", level="explicit"), + schemas.ObservationInput(content="long fact", level="explicit"), + ], + observer=peer1.name, + observed=peer2.name, + session_name=session.name, + workspace_name=workspace.name, + message_ids=[], + message_created_at=str(datetime.now(timezone.utc)), + ) + + assert isinstance(result, ObservationsCreatedResult) + assert result.created_count == 2 + assert captured["on_oversize"] == "truncate" + assert captured["texts"] == ["short fact", "long fact"] + async def test_create_observations_skips_all_blank_content( self, tool_test_data: Any,