fix(deriver): truncate oversize observations so one cannot drop the batch

simple_batch_embed raised ValueError when any input exceeded the per-input
token cap, which failed the entire deriver save when a single observation
was over-length. Add on_oversize="truncate": oversize inputs are embedded
from a token-capped prefix (re-encoded until it fits, with a warning),
preserving one vector per input. Default stays "raise" so existing callers
are unchanged. RepresentationManager opts into truncate.

Also add a live embedding test that fails on main (raise / missing kwarg)
and passes once a mixed short+oversize batch survives.

Refs #569
This commit is contained in:
Aakash Kattelu 2026-08-12 11:49:10 -07:00
parent 7b8c2917f9
commit 199dae8821
7 changed files with 279 additions and 18 deletions

View File

@ -107,7 +107,7 @@ class RepresentationManager:
parent_category="representation",
):
embeddings = await embedding_client.simple_batch_embed(
observation_texts
observation_texts, on_oversize="truncate"
)
except ValueError as e:
raise exceptions.ValidationException(

View File

@ -309,39 +309,82 @@ class _EmbeddingClient:
fn=_call_openai,
)
async def simple_batch_embed(self, texts: list[str]) -> list[list[float]]:
def _truncate_to_token_limit(self, text: str) -> tuple[str, int]:
"""Return a prefix of `text` whose re-encoded token count fits the cap.
`decode(ids[:n])` can re-encode to more than `n` tokens at BPE /
pre-tokenizer boundaries, so we re-encode after each slice and trim
again until the verified count fits. If a slice does not shrink the
count we drop one extra token so the loop cannot spin.
"""
Batch-embed a list of text strings. Each input must already fit within
`max_embedding_tokens`; this method does not sub-chunk oversized inputs.
token_ids = self.encoding.encode(text)
while len(token_ids) > self.max_embedding_tokens:
keep = min(self.max_embedding_tokens, len(token_ids) - 1)
if keep < 1:
return "", 0
text = self.encoding.decode(token_ids[:keep])
token_ids = self.encoding.encode(text)
return text, len(token_ids)
async def simple_batch_embed(
self,
texts: list[str],
*,
on_oversize: Literal["raise", "truncate"] = "raise",
) -> list[list[float]]:
"""
Batch-embed a list of text strings. Does not sub-chunk oversized inputs.
Internally goes through the same token-aware batching pipeline as
`batch_embed()` so the per-request token cap is respected.
Args:
texts: List of text strings to embed
on_oversize: What to do when an input exceeds `max_embedding_tokens`.
``"raise"`` (default) preserves the historical contract.
``"truncate"`` embeds a token-capped prefix and keeps the
one-input one-vector mapping, so one long item cannot drop
the rest of the batch.
Returns:
List of embedding vectors, one per input text (in order)
Raises:
ValueError: If any text exceeds token limits
ValueError: If any text exceeds token limits and `on_oversize` is
``"raise"``
"""
if not texts:
return []
# Validate per-input token limit and collect token counts for batching
# Validate / cap per-input token limit and collect counts for batching
prepared_texts: list[str] = []
token_counts: list[int] = []
for idx, text in enumerate(texts):
tokens = len(self.encoding.encode(text))
if tokens > self.max_embedding_tokens:
raise ValueError(
f"Text at index {idx} exceeds maximum token limit of {self.max_embedding_tokens} tokens (got {tokens} tokens)"
)
token_ids = self.encoding.encode(text)
if len(token_ids) > self.max_embedding_tokens:
if on_oversize == "truncate":
original_count = len(token_ids)
text, tokens = self._truncate_to_token_limit(text)
logger.warning(
"truncated oversize embedding input at idx %d: %d->%d tokens",
idx,
original_count,
tokens,
)
else:
raise ValueError(
f"Text at index {idx} exceeds maximum token limit of "
+ f"{self.max_embedding_tokens} tokens (got {len(token_ids)} tokens)"
)
else:
tokens = len(token_ids)
prepared_texts.append(text)
token_counts.append(tokens)
# Use positional indices as text_ids so we can reassemble in input order.
text_chunks: dict[str, list[tuple[str, int]]] = {
str(i): [(text, token_counts[i])] for i, text in enumerate(texts)
str(i): [(prepared_texts[i], token_counts[i])]
for i in range(len(prepared_texts))
}
batches = self._create_batches(text_chunks)
@ -679,9 +722,16 @@ class EmbeddingClient:
"""Embed a single query string."""
return await self._get_client().embed(query)
async def simple_batch_embed(self, texts: list[str]) -> list[list[float]]:
async def simple_batch_embed(
self,
texts: list[str],
*,
on_oversize: Literal["raise", "truncate"] = "raise",
) -> list[list[float]]:
"""Batch embed a list of text strings (each must fit token limit)."""
return await self._get_client().simple_batch_embed(texts)
return await self._get_client().simple_batch_embed(
texts, on_oversize=on_oversize
)
def prepare_chunks(self, id_resource_dict: dict[str, str]) -> dict[str, list[str]]:
"""Chunk texts using the same rules as `batch_embed` (no network)."""

View File

@ -628,7 +628,9 @@ def mock_openai_embeddings(request: pytest.FixtureRequest):
mock_embed.side_effect = embed_side_effect
async def mock_simple_batch_embed_func(texts: list[str]) -> list[list[float]]:
async def mock_simple_batch_embed_func(
texts: list[str], **_kwargs: object
) -> list[list[float]]:
return [_content_to_embedding(text) for text in texts]
mock_simple_batch_embed.side_effect = mock_simple_batch_embed_func

View File

@ -520,7 +520,9 @@ class TestRepresentationManagerSave:
)
assert len(saved.created_documents) == 1
mock_embed.assert_awaited_once_with(["useful observation"])
mock_embed.assert_awaited_once_with(
["useful observation"], on_oversize="truncate"
)
saved_observations = _saved_observations(mock_save)
assert len(saved_observations) == 1
assert saved_observations[0].content == "useful observation"
@ -576,7 +578,9 @@ class TestRepresentationManagerSave:
)
assert len(saved.created_documents) == 1
mock_embed.assert_awaited_once_with(["inferred conclusion"])
mock_embed.assert_awaited_once_with(
["inferred conclusion"], on_oversize="truncate"
)
saved_observations = _saved_observations(mock_save)
assert len(saved_observations) == 1
assert isinstance(saved_observations[0], DeductiveObservation)
@ -630,6 +634,61 @@ class TestRepresentationManagerSave:
mock_embed.assert_not_awaited()
mock_save.assert_not_awaited()
@pytest.mark.asyncio
async def test_save_representation_embeds_with_truncate_on_oversize(self):
"""One oversize observation must not drop the rest of the batch (#569)."""
manager = RepresentationManager(
"workspace",
observer="observer",
observed="observed",
)
representation = Representation(
explicit=[
ExplicitObservation(
content="short fact",
created_at=datetime.now(timezone.utc),
message_ids=[1],
session_name="session",
)
],
deductive=[
DeductiveObservation(
conclusion="inferred fact",
premises=["premise"],
source_ids=["doc-a"],
created_at=datetime.now(timezone.utc),
message_ids=[1],
session_name="session",
)
],
)
with (
patch("src.crud.representation.tracked_db", _fake_tracked_db),
patch(
"src.crud.representation.embedding_client.simple_batch_embed",
new=AsyncMock(return_value=[[0.1], [0.2]]),
) as mock_embed,
patch.object(
manager,
"_save_representation_internal",
new=AsyncMock(
return_value=CreateDocumentsResult(created_documents=[MagicMock()])
),
),
):
await manager.save_representation(
representation,
message_ids=[1],
session_name="session",
message_created_at=datetime.now(timezone.utc),
message_level_configuration=_resolved_config(),
)
mock_embed.assert_awaited_once_with(
["inferred fact", "short fact"], on_oversize="truncate"
)
class TestVectorQueryTopKFloor:
"""Regression for HONCHO-19Q / HONCHO-4Q4.

View File

@ -68,5 +68,5 @@ Coverage by provider:
- OpenAI transport → OpenRouter non-reasoning models (e.g. `inception/mercury-2`): non-chat / diffusion architectures must stay on `max_tokens`, no `reasoning_effort`, tool-calling parameter-schema compatibility is the canary for exotic OR-served providers
- Gemini 2.5/3.0 classes: structured outputs, cached-content reuse, thought signatures, multi-turn tool replay
- Gemini 3.1 class: thinking and tool replay coverage by default; structured-output/caching coverage should only be added once Google documents support for that path
- Embeddings (`test_live_embeddings.py`): single embed, batched embed, batch-vs-single alignment, and chunk-to-id mapping for every configured embedding model. `gemini-embedding-2*` is the reason this exists — those models collapse a list of bare strings into one document (#745), and only a live call catches it
- Embeddings (`test_live_embeddings.py`): single embed, batched embed, batch-vs-single alignment, chunk-to-id mapping, and oversize-truncate survival (`on_oversize="truncate"`, #569) for every configured embedding model. `gemini-embedding-2*` is the reason this exists — those models collapse a list of bare strings into one document (#745), and only a live call catches it
- OpenAI-compatible embedding providers (e.g. OpenRouter's `google/gemini-embedding-001`): the #932 surface. Those providers reject a base64 embedding request outright (HTTP 400) or answer HTTP 200 with empty data, so the whole matrix fails without `encoding_format="float"`. Real OpenAI accepts base64 happily, so only a third-party provider catches it. Note that OpenRouter load-balances across upstreams, so the base64 failure is per-attempt rather than guaranteed: a retry can land on an endpoint that accepts it. `test_live_openai_float_encoding_matches_base64` covers the other side, that the float switch must not move vectors on real OpenAI

View File

@ -146,6 +146,30 @@ async def test_live_openai_float_encoding_matches_base64(
), f"{spec.id}: float encoding diverges from base64 (cosine={similarity:.8f})"
@pytest.mark.asyncio
@pytest.mark.parametrize("spec", ALL_SPECS, ids=lambda spec: spec.id)
async def test_live_batch_embed_truncates_oversize_instead_of_dropping_batch(
spec: LiveEmbeddingSpec,
) -> None:
"""Regression for #569. Fails on main: simple_batch_embed raises
ValueError (or TypeError, before `on_oversize` existed) when any input
exceeds the per-item cap, so the rest of the batch is never embedded.
After the truncate path, one oversize item cannot drop the others.
"""
# Tiny cap so the oversize input stays cheap to tokenize and send.
client = make_embedding_client(spec, max_input_tokens=32)
oversize = " ".join(f"oversize-token-{index}" for index in range(200))
assert len(client.encoding.encode(oversize)) > client.max_embedding_tokens
texts = [BATCH_TEXTS[0], oversize, BATCH_TEXTS[1]]
embeddings = await client.simple_batch_embed(texts, on_oversize="truncate")
assert len(embeddings) == len(texts)
assert all(len(embedding) == spec.dimensions for embedding in embeddings)
# A collapsed or dropped batch would reuse a vector or return fewer.
assert len({tuple(embedding) for embedding in embeddings}) == len(texts)
@pytest.mark.asyncio
@pytest.mark.parametrize("spec", GEMINI_SPECS, ids=lambda spec: spec.id)
async def test_live_gemini_batch_embed_survives_batch_split(

View File

@ -13,6 +13,7 @@ from src.config import (
)
from src.embedding_client import (
BatchItem,
EmbeddingClient,
_EmbeddingClient, # pyright: ignore[reportPrivateUsage]
)
@ -768,6 +769,131 @@ async def test_simple_batch_embed_rejects_oversized_input(
await client.simple_batch_embed([too_long])
@pytest.mark.asyncio
async def test_simple_batch_embed_truncates_oversize_when_requested(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""on_oversize='truncate' embeds a prefix instead of failing the batch."""
fake_embeddings = FakeOpenAIEmbeddingsAPI([0.1] * 4)
class FakeOpenAIClient:
def __init__(self, *, api_key: str | None, base_url: str | None) -> None:
self.embeddings: FakeOpenAIEmbeddingsAPI = fake_embeddings
monkeypatch.setattr("src.embedding_client.AsyncOpenAI", FakeOpenAIClient)
client = _EmbeddingClient(
EmbeddingModelConfig(
transport="openai",
model="text-embedding-3-small",
api_key="test-key",
base_url=None,
),
vector_dimensions=4,
max_input_tokens=10,
max_tokens_per_request=1000,
send_dimensions=False,
)
short = "hello"
too_long = ("word " * 50).strip()
assert len(client.encoding.encode(too_long)) > client.max_embedding_tokens
out = await client.simple_batch_embed([short, too_long], on_oversize="truncate")
assert len(out) == 2
assert fake_embeddings.calls, "expected a provider call after truncation"
received = fake_embeddings.calls[0]["input"]
assert received[0] == short
truncated = received[1]
assert isinstance(truncated, str)
assert truncated != too_long
assert len(client.encoding.encode(truncated)) <= client.max_embedding_tokens
@pytest.mark.asyncio
async def test_simple_batch_embed_truncate_reencodes_until_under_cap(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""decode(ids[:n]) can re-encode past n; truncate must re-verify the count."""
fake_embeddings = FakeOpenAIEmbeddingsAPI([0.1] * 4)
class FakeOpenAIClient:
def __init__(self, *, api_key: str | None, base_url: str | None) -> None:
self.embeddings: FakeOpenAIEmbeddingsAPI = fake_embeddings
monkeypatch.setattr("src.embedding_client.AsyncOpenAI", FakeOpenAIClient)
client = _EmbeddingClient(
EmbeddingModelConfig(
transport="openai",
model="text-embedding-3-small",
api_key="test-key",
base_url=None,
),
vector_dimensions=4,
max_input_tokens=10,
max_tokens_per_request=1000,
send_dimensions=False,
)
encode_calls = {"n": 0}
def encode(text: str) -> list[int]:
encode_calls["n"] += 1
if text.startswith("LONG"):
# 1: original oversize; 2: still over after first slice; 3+: fits.
if encode_calls["n"] == 1:
return list(range(20))
if encode_calls["n"] == 2:
return list(range(12))
return list(range(8))
return [1]
def decode(ids: list[int]) -> str:
return "LONG" + "x" * len(ids)
monkeypatch.setattr(client.encoding, "encode", encode)
monkeypatch.setattr(client.encoding, "decode", decode)
out = await client.simple_batch_embed(["LONG-input"], on_oversize="truncate")
assert len(out) == 1
received = fake_embeddings.calls[0]["input"][0]
assert isinstance(received, str)
# The provider must see the post-loop text, which encodes to 8 (<= cap).
assert encode(received) == list(range(8))
assert encode_calls["n"] >= 3
@pytest.mark.asyncio
async def test_public_embedding_client_forwards_on_oversize(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""The singleton wrapper must forward on_oversize to the inner client."""
captured: dict[str, object] = {}
class FakeInner:
async def simple_batch_embed(
self,
texts: list[str],
*,
on_oversize: str = "raise",
) -> list[list[float]]:
captured["texts"] = texts
captured["on_oversize"] = on_oversize
return [[0.1]]
wrapper = EmbeddingClient()
monkeypatch.setattr(wrapper, "_get_client", lambda: FakeInner())
out = await wrapper.simple_batch_embed(["hi"], on_oversize="truncate")
assert out == [[0.1]]
assert captured["texts"] == ["hi"]
assert captured["on_oversize"] == "truncate"
def test_prepare_chunks_returns_ordered_chunks(
monkeypatch: pytest.MonkeyPatch,
) -> None: