From 66ca4bca0d84b563d4e1494836088790b98fe1a4 Mon Sep 17 00:00:00 2001 From: Dannong Xu Date: Sun, 26 Jul 2026 20:08:07 +0800 Subject: [PATCH 1/8] refactor(retain): replace orchestration with ingestion pipeline Route Retain through a database-neutral application service with bounded write windows, durable checkpoints, cancellation fences, and PostgreSQL/Oracle adapters. Preserve third-party attribution and add focused contract/live regression coverage. Refs #1 --- THIRD_PARTY_NOTICES.md | 29 + core/dataplane/LICENSE | 21 + core/dataplane/README.md | 3 +- core/dataplane/THIRD_PARTY_NOTICES.md | 29 + core/dataplane/hms_api/config.py | 20 + .../dataplane/hms_api/engine/db/ops_oracle.py | 11 +- .../hms_api/engine/db/ops_postgresql.py | 40 +- core/dataplane/hms_api/engine/db/oracle.py | 32 + .../hms_api/engine/embedding_fingerprint.py | 25 +- .../engine/entity_resolution_contracts.py | 105 + .../hms_api/engine/entity_resolver.py | 380 ++- .../hms_api/engine/ingestion/__init__.py | 61 + .../engine/ingestion/adapters/__init__.py | 10 + .../ingestion/adapters/embedding_model.py | 30 + .../ingestion/adapters/oracle_semantic.py | 77 + .../adapters/postgres_fresh_ownership.py | 109 + .../ingestion/adapters/storage_records.py | 179 ++ .../engine/ingestion/change_detection.py | 116 + .../hms_api/engine/ingestion/chunking.py | 152 ++ .../hms_api/engine/ingestion/contracts.py | 73 + .../engine/ingestion/document_planner.py | 210 ++ .../hms_api/engine/ingestion/domain.py | 158 ++ .../engine/ingestion/execution/__init__.py | 17 + .../engine/ingestion/execution/windowing.py | 200 ++ .../engine/ingestion/extraction/__init__.py | 54 + .../engine/ingestion/extraction/extractor.py | 495 ++++ .../engine/ingestion/extraction/layout.py | 324 +++ .../engine/ingestion/extraction/models.py | 145 ++ .../ingestion/extraction/passthrough.py | 151 ++ .../engine/ingestion/extraction/ports.py | 116 + .../hms_api/engine/ingestion/normalization.py | 276 ++ .../engine/ingestion/persistence/__init__.py | 14 + .../engine/ingestion/persistence/backend.py | 87 + .../engine/ingestion/persistence/models.py | 112 + .../ingestion/persistence/operation_fence.py | 92 + .../engine/ingestion/persistence/oracle.py | 480 ++++ .../engine/ingestion/persistence/ports.py | 53 + .../engine/ingestion/persistence/postgres.py | 566 +++++ .../ingestion/persistence/unit_of_work.py | 631 +++++ .../engine/ingestion/persistence/writer.py | 477 ++++ .../engine/ingestion/projection/__init__.py | 27 + .../engine/ingestion/projection/embeddings.py | 126 + .../engine/ingestion/projection/records.py | 178 ++ .../hms_api/engine/ingestion/redaction.py | 50 + .../hms_api/engine/ingestion/runtime.py | 419 ++++ .../hms_api/engine/ingestion/service.py | 1591 ++++++++++++ .../dataplane/hms_api/engine/memory_engine.py | 680 +++-- .../hms_api/engine/retain/chunk_storage.py | 13 +- .../hms_api/engine/retain/embedding_utils.py | 4 +- .../hms_api/engine/retain/entity_labels.py | 9 +- .../engine/retain/entity_processing.py | 54 + .../hms_api/engine/retain/fact_extraction.py | 265 +- .../hms_api/engine/retain/fact_storage.py | 27 +- .../hms_api/engine/retain/link_utils.py | 126 +- .../hms_api/engine/retain/orchestrator.py | 2209 ----------------- core/dataplane/hms_api/engine/retain/types.py | 10 + core/dataplane/hms_api/worker/poller.py | 44 +- core/dataplane/pyproject.toml | 2 + .../tests/e2e/test_multimodal_security_e2e.py | 10 +- .../tests/test_async_batch_retain.py | 91 +- .../dataplane/tests/test_config_validation.py | 35 + core/dataplane/tests/test_db_abstraction.py | 79 + core/dataplane/tests/test_delta_retain.py | 15 - .../tests/test_fact_extraction_retry.py | 31 + core/dataplane/tests/test_file_retain.py | 2 +- .../tests/test_ingestion_oracle_contracts.py | 602 +++++ .../tests/test_ingestion_oracle_live.py | 208 ++ .../test_ingestion_pipeline_contracts.py | 1473 +++++++++++ .../tests/test_ingestion_postgresql_live.py | 675 +++++ core/dataplane/tests/test_link_utils.py | 54 +- .../tests/test_multimodal_admission.py | 2 +- .../tests/test_multimodal_engine_bridge.py | 37 +- .../tests/test_multimodal_security.py | 30 +- .../tests/test_observation_invalidation.py | 4 +- core/dataplane/tests/test_op_cancellation.py | 261 +- .../tests/test_retain_orchestrator_mapping.py | 526 ---- deploy/containers/standalone/Dockerfile | 2 + docker-compose.yml | 1 + docs/multimodal_memory.md | 4 +- docs/system_architecture_and_multimodal.md | 2 +- 80 files changed, 12948 insertions(+), 3190 deletions(-) create mode 100644 THIRD_PARTY_NOTICES.md create mode 100644 core/dataplane/LICENSE create mode 100644 core/dataplane/THIRD_PARTY_NOTICES.md create mode 100644 core/dataplane/hms_api/engine/entity_resolution_contracts.py create mode 100644 core/dataplane/hms_api/engine/ingestion/__init__.py create mode 100644 core/dataplane/hms_api/engine/ingestion/adapters/__init__.py create mode 100644 core/dataplane/hms_api/engine/ingestion/adapters/embedding_model.py create mode 100644 core/dataplane/hms_api/engine/ingestion/adapters/oracle_semantic.py create mode 100644 core/dataplane/hms_api/engine/ingestion/adapters/postgres_fresh_ownership.py create mode 100644 core/dataplane/hms_api/engine/ingestion/adapters/storage_records.py create mode 100644 core/dataplane/hms_api/engine/ingestion/change_detection.py create mode 100644 core/dataplane/hms_api/engine/ingestion/chunking.py create mode 100644 core/dataplane/hms_api/engine/ingestion/contracts.py create mode 100644 core/dataplane/hms_api/engine/ingestion/document_planner.py create mode 100644 core/dataplane/hms_api/engine/ingestion/domain.py create mode 100644 core/dataplane/hms_api/engine/ingestion/execution/__init__.py create mode 100644 core/dataplane/hms_api/engine/ingestion/execution/windowing.py create mode 100644 core/dataplane/hms_api/engine/ingestion/extraction/__init__.py create mode 100644 core/dataplane/hms_api/engine/ingestion/extraction/extractor.py create mode 100644 core/dataplane/hms_api/engine/ingestion/extraction/layout.py create mode 100644 core/dataplane/hms_api/engine/ingestion/extraction/models.py create mode 100644 core/dataplane/hms_api/engine/ingestion/extraction/passthrough.py create mode 100644 core/dataplane/hms_api/engine/ingestion/extraction/ports.py create mode 100644 core/dataplane/hms_api/engine/ingestion/normalization.py create mode 100644 core/dataplane/hms_api/engine/ingestion/persistence/__init__.py create mode 100644 core/dataplane/hms_api/engine/ingestion/persistence/backend.py create mode 100644 core/dataplane/hms_api/engine/ingestion/persistence/models.py create mode 100644 core/dataplane/hms_api/engine/ingestion/persistence/operation_fence.py create mode 100644 core/dataplane/hms_api/engine/ingestion/persistence/oracle.py create mode 100644 core/dataplane/hms_api/engine/ingestion/persistence/ports.py create mode 100644 core/dataplane/hms_api/engine/ingestion/persistence/postgres.py create mode 100644 core/dataplane/hms_api/engine/ingestion/persistence/unit_of_work.py create mode 100644 core/dataplane/hms_api/engine/ingestion/persistence/writer.py create mode 100644 core/dataplane/hms_api/engine/ingestion/projection/__init__.py create mode 100644 core/dataplane/hms_api/engine/ingestion/projection/embeddings.py create mode 100644 core/dataplane/hms_api/engine/ingestion/projection/records.py create mode 100644 core/dataplane/hms_api/engine/ingestion/redaction.py create mode 100644 core/dataplane/hms_api/engine/ingestion/runtime.py create mode 100644 core/dataplane/hms_api/engine/ingestion/service.py delete mode 100644 core/dataplane/hms_api/engine/retain/orchestrator.py create mode 100644 core/dataplane/tests/test_ingestion_oracle_contracts.py create mode 100644 core/dataplane/tests/test_ingestion_oracle_live.py create mode 100644 core/dataplane/tests/test_ingestion_pipeline_contracts.py create mode 100644 core/dataplane/tests/test_ingestion_postgresql_live.py delete mode 100644 core/dataplane/tests/test_retain_orchestrator_mapping.py diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md new file mode 100644 index 0000000..f4e51a2 --- /dev/null +++ b/THIRD_PARTY_NOTICES.md @@ -0,0 +1,29 @@ +# Third-Party Notices + +## Hindsight + +Portions of this repository are derived from +[Hindsight](https://github.com/vectorize-io/hindsight), which is distributed +under the MIT License: + +> MIT License +> +> Copyright (c) 2025 Vectorize AI, Inc. +> +> Permission is hereby granted, free of charge, to any person obtaining a copy +> of this software and associated documentation files (the "Software"), to deal +> in the Software without restriction, including without limitation the rights +> to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +> copies of the Software, and to permit persons to whom the Software is +> furnished to do so, subject to the following conditions: +> +> The above copyright notice and this permission notice shall be included in +> all copies or substantial portions of the Software. +> +> THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +> IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +> FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +> AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +> LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +> OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +> SOFTWARE. diff --git a/core/dataplane/LICENSE b/core/dataplane/LICENSE new file mode 100644 index 0000000..ddca1cc --- /dev/null +++ b/core/dataplane/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2025 HMS AI, Inc. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/core/dataplane/README.md b/core/dataplane/README.md index f782585..b18d68c 100644 --- a/core/dataplane/README.md +++ b/core/dataplane/README.md @@ -193,4 +193,5 @@ Full documentation: [https://docs.hms.local](https://docs.hms.local) ## License -Apache 2.0 +MIT. See the package [LICENSE](LICENSE) and +[third-party notices](THIRD_PARTY_NOTICES.md). diff --git a/core/dataplane/THIRD_PARTY_NOTICES.md b/core/dataplane/THIRD_PARTY_NOTICES.md new file mode 100644 index 0000000..e47d02b --- /dev/null +++ b/core/dataplane/THIRD_PARTY_NOTICES.md @@ -0,0 +1,29 @@ +# Third-Party Notices + +## Hindsight + +Portions of this package are derived from +[Hindsight](https://github.com/vectorize-io/hindsight), which is distributed +under the MIT License: + +> MIT License +> +> Copyright (c) 2025 Vectorize AI, Inc. +> +> Permission is hereby granted, free of charge, to any person obtaining a copy +> of this software and associated documentation files (the "Software"), to deal +> in the Software without restriction, including without limitation the rights +> to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +> copies of the Software, and to permit persons to whom the Software is +> furnished to do so, subject to the following conditions: +> +> The above copyright notice and this permission notice shall be included in +> all copies or substantial portions of the Software. +> +> THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +> IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +> FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +> AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +> LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +> OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +> SOFTWARE. diff --git a/core/dataplane/hms_api/config.py b/core/dataplane/hms_api/config.py index 725adfa..3d571dc 100644 --- a/core/dataplane/hms_api/config.py +++ b/core/dataplane/hms_api/config.py @@ -358,6 +358,7 @@ def normalize_config_dict(config: dict[str, Any]) -> dict[str, Any]: ENV_RETAIN_BATCH_ENABLED = "HMS_API_RETAIN_BATCH_ENABLED" ENV_RETAIN_BATCH_POLL_INTERVAL_SECONDS = "HMS_API_RETAIN_BATCH_POLL_INTERVAL_SECONDS" ENV_RETAIN_CHUNK_BATCH_SIZE = "HMS_API_RETAIN_CHUNK_BATCH_SIZE" +ENV_RETAIN_EMBEDDING_FAILURE_POLICY = "HMS_API_RETAIN_EMBEDDING_FAILURE_POLICY" # File storage configuration ENV_FILE_STORAGE_TYPE = "HMS_API_FILE_STORAGE_TYPE" @@ -683,6 +684,8 @@ def normalize_config_dict(config: dict[str, Any]) -> dict[str, Any]: DEFAULT_RETAIN_ENTITY_LOOKUP = "trigram" # "full" or "trigram" DEFAULT_RETAIN_BATCH_ENABLED = False # Use LLM Batch API for fact extraction (only when async=True) DEFAULT_RETAIN_BATCH_POLL_INTERVAL_SECONDS = 60 # Batch API polling interval in seconds +DEFAULT_RETAIN_EMBEDDING_FAILURE_POLICY = "store_without_embedding" +RETAIN_EMBEDDING_FAILURE_POLICIES = ("store_without_embedding", "raise") # File storage defaults DEFAULT_FILE_STORAGE_TYPE = "native" # PostgreSQL BYTEA storage @@ -943,6 +946,16 @@ def _validate_extraction_mode(mode: str) -> str: return mode_lower +def _validate_retain_embedding_failure_policy(policy: str) -> str: + """Validate and normalize Retain's whole-batch embedding failure policy.""" + + normalized = policy.strip().lower() + if normalized not in RETAIN_EMBEDDING_FAILURE_POLICIES: + choices = ", ".join(RETAIN_EMBEDDING_FAILURE_POLICIES) + raise ValueError(f"{ENV_RETAIN_EMBEDDING_FAILURE_POLICY} must be one of: {choices}") + return normalized + + def _validate_recall_budget_function(function: str) -> str: """Validate and normalize recall budget function.""" function_lower = function.lower() @@ -1331,6 +1344,7 @@ class HMSConfig: # Defaulted fields (source-compatible additions — existing direct constructor callers keep working). # Keep at the end of the dataclass; Python forbids non-default fields after default fields. embeddings_openai_batch_size: int = DEFAULT_EMBEDDINGS_OPENAI_BATCH_SIZE + retain_embedding_failure_policy: str = DEFAULT_RETAIN_EMBEDDING_FAILURE_POLICY embedding_fingerprint_policy: Literal["strict", "warn", "off"] = DEFAULT_EMBEDDING_FINGERPRINT_POLICY embedding_fingerprint_legacy_attestation: str | None = None vector_index_provider: str = DEFAULT_VECTOR_INDEX_PROVIDER @@ -2140,6 +2154,12 @@ def from_env(cls) -> "HMSConfig": os.getenv(ENV_RETAIN_BATCH_POLL_INTERVAL_SECONDS, str(DEFAULT_RETAIN_BATCH_POLL_INTERVAL_SECONDS)) ), retain_chunk_batch_size=int(os.getenv(ENV_RETAIN_CHUNK_BATCH_SIZE, str(DEFAULT_RETAIN_CHUNK_BATCH_SIZE))), + retain_embedding_failure_policy=_validate_retain_embedding_failure_policy( + os.getenv( + ENV_RETAIN_EMBEDDING_FAILURE_POLICY, + DEFAULT_RETAIN_EMBEDDING_FAILURE_POLICY, + ) + ), # File storage file_storage_type=os.getenv(ENV_FILE_STORAGE_TYPE, DEFAULT_FILE_STORAGE_TYPE), file_storage_s3_bucket=os.getenv(ENV_FILE_STORAGE_S3_BUCKET) or None, diff --git a/core/dataplane/hms_api/engine/db/ops_oracle.py b/core/dataplane/hms_api/engine/db/ops_oracle.py index 656b49a..8b6f8bf 100644 --- a/core/dataplane/hms_api/engine/db/ops_oracle.py +++ b/core/dataplane/hms_api/engine/db/ops_oracle.py @@ -197,8 +197,15 @@ async def fetch_missing_entity_ids( orig_name, ) if row: - # Wrap in a dict-like to include input_name for downstream compat - results.append(row) + results.append( + ResultRow( + { + "id": row["id"], + "name_lower": row["name_lower"], + "input_name": orig_name, + } + ) + ) return results async def bulk_insert_unit_entities( diff --git a/core/dataplane/hms_api/engine/db/ops_postgresql.py b/core/dataplane/hms_api/engine/db/ops_postgresql.py index 56f0cd3..695ec6f 100644 --- a/core/dataplane/hms_api/engine/db/ops_postgresql.py +++ b/core/dataplane/hms_api/engine/db/ops_postgresql.py @@ -7,7 +7,7 @@ import json from datetime import UTC, datetime from typing import Any -from uuid import UUID +from uuid import UUID, uuid4 from .base import DatabaseConnection from .ops import DataAccessOps, TagListingParts @@ -74,23 +74,30 @@ async def insert_facts_batch( config = get_config() table = self._get_mu_table() + # Generate IDs in caller/input order and insert those exact values. + # PostgreSQL does not guarantee row order for INSERT ... RETURNING; + # treating its result order as the fact order can attach entity/causal + # bindings to the wrong unit. Client-owned IDs make the mapping an + # explicit contract and match the existing Oracle implementation. + unit_uuids = [uuid4() for _ in fact_texts] + unit_ids = [str(unit_id) for unit_id in unit_uuids] if config.text_search_extension == "vchord": query = f""" WITH input_data AS ( SELECT * FROM unnest( - $2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[], - $8::text[], $9::text[], $10::jsonb[], $11::text[], $12::text[], $13::jsonb[], $14::jsonb[], $15::text[], - $16::jsonb[] - ) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at, + $2::uuid[], $3::text[], $4::vector[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[], + $8::timestamptz[], $9::text[], $10::text[], $11::jsonb[], $12::text[], $13::text[], $14::jsonb[], + $15::jsonb[], $16::text[], $17::jsonb[] + ) AS t(id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at, context, fact_type, metadata, chunk_id, document_id, tags_json, observation_scopes_json, text_signals, projection) ) - INSERT INTO {table} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at, + INSERT INTO {table} (id, bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at, context, fact_type, metadata, chunk_id, document_id, tags, observation_scopes, text_signals, projection, search_vector) SELECT - $1, + id, $1, text, embedding, event_date, occurred_start, occurred_end, mentioned_at, context, fact_type, metadata, chunk_id, document_id, COALESCE( @@ -105,24 +112,23 @@ async def insert_facts_batch( 'llmlingua2' )::bm25_catalog.bm25vector FROM input_data - RETURNING id """ else: query = f""" WITH input_data AS ( SELECT * FROM unnest( - $2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[], - $8::text[], $9::text[], $10::jsonb[], $11::text[], $12::text[], $13::jsonb[], $14::jsonb[], $15::text[], - $16::jsonb[] - ) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at, + $2::uuid[], $3::text[], $4::vector[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[], + $8::timestamptz[], $9::text[], $10::text[], $11::jsonb[], $12::text[], $13::text[], $14::jsonb[], + $15::jsonb[], $16::text[], $17::jsonb[] + ) AS t(id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at, context, fact_type, metadata, chunk_id, document_id, tags_json, observation_scopes_json, text_signals, projection) ) - INSERT INTO {table} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at, + INSERT INTO {table} (id, bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at, context, fact_type, metadata, chunk_id, document_id, tags, observation_scopes, text_signals, projection) SELECT - $1, + id, $1, text, embedding, event_date, occurred_start, occurred_end, mentioned_at, context, fact_type, metadata, chunk_id, document_id, COALESCE( @@ -133,12 +139,12 @@ async def insert_facts_batch( text_signals, COALESCE(projection, '{{}}'::jsonb) FROM input_data - RETURNING id """ - results = await conn.fetch( + await conn.execute( query, bank_id, + unit_uuids, fact_texts, embeddings, event_dates, @@ -155,7 +161,7 @@ async def insert_facts_batch( text_signals_list, projection_jsons, ) - return [str(row["id"]) for row in results] + return unit_ids async def bulk_insert_links( self, diff --git a/core/dataplane/hms_api/engine/db/oracle.py b/core/dataplane/hms_api/engine/db/oracle.py index 0d3e70d..80f9957 100644 --- a/core/dataplane/hms_api/engine/db/oracle.py +++ b/core/dataplane/hms_api/engine/db/oracle.py @@ -90,6 +90,14 @@ class RewriteResult(NamedTuple): _NOT_ALL_RE = re.compile(r"!=\s*ALL\s*\(\s*:(\d+)\s*\)", re.IGNORECASE) _JSON_ARROW_TEXT_RE = re.compile(r'("?\w+"?)\s*->>\s*\'(\w+)\'') # handles both col and "col" +_JSON_UUID_TEXT_EQ_RE = re.compile( + r"""\(\s*("?\w+"?)\s*->>\s*'(\w+)'\s*\)\s*::uuid\s*=\s*:(\d+)""", + re.IGNORECASE, +) +_JSON_NUMBER_TEXT_CAST_RE = re.compile( + r"""\(\s*("?\w+"?)\s*->>\s*'(\w+)'\s*\)\s*::(?:int|integer|bigint|numeric|float)\b""", + re.IGNORECASE, +) _JSON_HAS_KEY_RE = re.compile(r"(\w+)\s*\?\s*'(\w+)'") _JSONB_CONTAINS_RE = re.compile(r"(\w+)\s*@>\s*:(\d+)") @@ -356,6 +364,30 @@ def _rewrite_json_bool(m: re.Match) -> str: flags=re.IGNORECASE, ) + # PostgreSQL can cast a JSON text scalar to UUID before comparing it with + # a UUID bind. Preserve that typed comparison on Oracle by converting the + # canonical UUID text stored in JSON to RAW(16). The normal UUID argument + # conversion can then remain authoritative for RAW-backed identifiers. + # + # This must run before generic cast stripping and JSON arrow rewriting: + # (result_metadata->>'parent_operation_id')::uuid = :2 + # -> HEXTORAW(REPLACE(JSON_VALUE(...), '-', '')) = :2 + def _rewrite_json_uuid_text_equality(match: re.Match) -> str: + column, key, bind_index = match.groups() + json_text = f"JSON_VALUE({column}, '$.{key}')" + return f"HEXTORAW(REPLACE({json_text}, '-', '')) = :{bind_index}" + + query = _JSON_UUID_TEXT_EQ_RE.sub(_rewrite_json_uuid_text_equality, query) + + # Preserve numeric ordering and comparisons for JSON text casts. Stripping + # the PostgreSQL cast would leave JSON_VALUE as VARCHAR2, so values such as + # sub-batch 10 would sort before sub-batch 2 on Oracle. + def _rewrite_json_number_text_cast(match: re.Match) -> str: + column, key = match.groups() + return f"TO_NUMBER(JSON_VALUE({column}, '$.{key}'))" + + query = _JSON_NUMBER_TEXT_CAST_RE.sub(_rewrite_json_number_text_cast, query) + # Strip ::type casts (including bare ::jsonb on literals in generic contexts) query = _PG_CAST_RE.sub("", query) diff --git a/core/dataplane/hms_api/engine/embedding_fingerprint.py b/core/dataplane/hms_api/engine/embedding_fingerprint.py index 097918c..7ca96e1 100644 --- a/core/dataplane/hms_api/engine/embedding_fingerprint.py +++ b/core/dataplane/hms_api/engine/embedding_fingerprint.py @@ -17,7 +17,7 @@ import logging import re from collections.abc import Mapping -from typing import Any +from typing import Any, Protocol from urllib.parse import urlsplit, urlunsplit from .schema import fq_table @@ -32,6 +32,12 @@ FINGERPRINT_POLICIES = frozenset({"strict", "warn", "off"}) +class _IdentifierLogSanitizer(Protocol): + """Minimal request-local identifier renderer used by security-sensitive callers.""" + + def identifier(self, value: Any) -> str: ... + + class EmbeddingFingerprintError(RuntimeError): """Base class for embedding compatibility failures.""" @@ -521,6 +527,7 @@ async def validate_bank_embedding_fingerprint( policy: str = "strict", for_write: bool = False, legacy_attestation: Any = None, + log_sanitizer: _IdentifierLogSanitizer | None = None, ) -> dict[str, Any]: """Validate a bank fingerprint, optionally initialising an empty bank. @@ -532,6 +539,7 @@ async def validate_bank_embedding_fingerprint( """ policy_value = _policy(policy) + log_bank_id = log_sanitizer.identifier(bank_id) if log_sanitizer is not None else bank_id current_fp = canonical_fingerprint(current) if current_fp is None or current_fp.get("legacy"): raise EmbeddingFingerprintSchemaError("Current embedding fingerprint is not a modern JSON object") @@ -544,7 +552,7 @@ async def validate_bank_embedding_fingerprint( # does not exist; retain's get_or_create path normally makes this # impossible, but a clear error is safer inside a write transaction. if for_write: - raise EmbeddingFingerprintError(f"Cannot fingerprint unknown bank {bank_id!r}: bank row does not exist") + raise EmbeddingFingerprintError(f"Cannot fingerprint unknown bank {log_bank_id!r}: bank row does not exist") return current_fp stored_fp = canonical_fingerprint(stored_raw) if stored_raw is not None else None @@ -554,12 +562,12 @@ async def validate_bank_embedding_fingerprint( current_fp, legacy_version=str(stored_fp.get("version") or ""), ): - logger.warning("Using explicitly attested legacy embedding fingerprint for bank %s", bank_id) + logger.warning("Using explicitly attested legacy embedding fingerprint for bank %s", log_bank_id) if for_write: await _persist_fingerprint(conn, bank_id, current_fp) return current_fp message = ( - f"Bank {bank_id!r} has a legacy embedding fingerprint ({stored_fp.get('version')!r}); " + f"Bank {log_bank_id!r} has a legacy embedding fingerprint ({stored_fp.get('version')!r}); " "projection metadata cannot establish vector compatibility. " "Re-index the bank or provide an explicit legacy attestation." ) @@ -576,12 +584,12 @@ async def validate_bank_embedding_fingerprint( # A read against an empty bank need not mutate it. return current_fp if _attestation_matches(legacy_attestation, current_fp): - logger.warning("Using explicit embedding attestation for legacy bank %s", bank_id) + logger.warning("Using explicit embedding attestation for legacy bank %s", log_bank_id) if for_write: await _persist_fingerprint(conn, bank_id, current_fp) return current_fp message = ( - f"Bank {bank_id!r} contains memory units but has no embedding fingerprint. " + f"Bank {log_bank_id!r} contains memory units but has no embedding fingerprint. " "Refusing to compare vectors without an explicit legacy attestation." ) if policy_value == "warn": @@ -600,7 +608,8 @@ async def validate_bank_embedding_fingerprint( return current_fp message = ( - f"Embedding fingerprint mismatch for bank {bank_id!r}: stored={stored_fp.get('hash', 'unknown')[:16]} " + f"Embedding fingerprint mismatch for bank {log_bank_id!r}: " + f"stored={stored_fp.get('hash', 'unknown')[:16]} " f"current={current_fp.get('hash', 'unknown')[:16]}. " "Use the encoder that created the bank or explicitly re-index/attest it; " "semantic recall and writes must not silently mix vector spaces." @@ -619,6 +628,7 @@ async def ensure_bank_embedding_fingerprint( policy: str = "strict", for_write: bool = False, legacy_attestation: Any = None, + log_sanitizer: _IdentifierLogSanitizer | None = None, ) -> dict[str, Any]: """Convenience wrapper that fingerprints an encoder object and validates it.""" @@ -630,6 +640,7 @@ async def ensure_bank_embedding_fingerprint( policy=policy, for_write=for_write, legacy_attestation=legacy_attestation, + log_sanitizer=log_sanitizer, ) diff --git a/core/dataplane/hms_api/engine/entity_resolution_contracts.py b/core/dataplane/hms_api/engine/entity_resolution_contracts.py new file mode 100644 index 0000000..0eb2c58 --- /dev/null +++ b/core/dataplane/hms_api/engine/entity_resolution_contracts.py @@ -0,0 +1,105 @@ +"""Immutable contracts for read-only entity planning and transactional finalization. + +The candidate lookup and scoring phase can be expensive, so Retain performs +it before opening the core write transaction. These values carry only the +decision produced by that read phase; unresolved canonical rows are created +later, on the Retain unit-of-work connection. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime +from typing import Any + + +@dataclass(frozen=True, slots=True) +class EntityOccurrenceBinding: + """Stable source identity for one entity mention in one projected fact.""" + + occurrence_key: str + unit_key: str + local_index: int + event_date: datetime | None + + def __post_init__(self) -> None: + if not self.occurrence_key: + raise ValueError("occurrence_key must be non-empty") + if not self.unit_key: + raise ValueError("unit_key must be non-empty") + if isinstance(self.local_index, bool) or not isinstance(self.local_index, int): + raise TypeError("local_index must be an integer") + if self.local_index < 0: + raise ValueError("local_index must be non-negative") + + +@dataclass(frozen=True, slots=True) +class ExistingEntityBinding: + """Read-phase decision binding an occurrence to an existing canonical row.""" + + occurrence_key: str + entity_id: Any + + def __post_init__(self) -> None: + if not self.occurrence_key: + raise ValueError("occurrence_key must be non-empty") + if self.entity_id is None or str(self.entity_id) == "": + raise ValueError("entity_id must be non-empty") + + +@dataclass(frozen=True, slots=True) +class UnresolvedEntityDescriptor: + """Canonical entity data that may need to be inserted during finalization.""" + + occurrence_key: str + canonical_name: str + entity_type: str + event_date: datetime | None + nearby_occurrence_keys: tuple[str, ...] = () + validated_labels: tuple[str, ...] = () + + def __post_init__(self) -> None: + if not self.occurrence_key: + raise ValueError("occurrence_key must be non-empty") + if not self.canonical_name: + raise ValueError("canonical_name must be non-empty") + if not self.entity_type: + raise ValueError("entity_type must be non-empty") + if len(self.nearby_occurrence_keys) != len(set(self.nearby_occurrence_keys)): + raise ValueError("nearby_occurrence_keys must be unique") + + +@dataclass(frozen=True, slots=True) +class EntityResolutionReadPlan: + """Complete, provider-neutral result of read-only entity resolution.""" + + bank_id: str + occurrences: tuple[EntityOccurrenceBinding, ...] + existing_bindings: tuple[ExistingEntityBinding, ...] = () + unresolved_descriptors: tuple[UnresolvedEntityDescriptor, ...] = () + + def __post_init__(self) -> None: + if not self.bank_id: + raise ValueError("bank_id must be non-empty") + occurrence_keys = tuple(item.occurrence_key for item in self.occurrences) + if len(occurrence_keys) != len(set(occurrence_keys)): + raise ValueError("occurrences contain duplicate occurrence keys") + existing_keys = tuple(item.occurrence_key for item in self.existing_bindings) + unresolved_keys = tuple(item.occurrence_key for item in self.unresolved_descriptors) + if len(existing_keys) != len(set(existing_keys)): + raise ValueError("existing_bindings contain duplicate occurrence keys") + if len(unresolved_keys) != len(set(unresolved_keys)): + raise ValueError("unresolved_descriptors contain duplicate occurrence keys") + if set(existing_keys) & set(unresolved_keys): + raise ValueError("an occurrence cannot be both existing and unresolved") + if set(occurrence_keys) != set(existing_keys) | set(unresolved_keys): + raise ValueError("entity read plan must resolve every occurrence exactly once") + + +@dataclass(frozen=True, slots=True) +class FinalizedEntityResolution: + """Resolved graph inputs after missing canonical rows are finalized.""" + + resolved_entity_ids: tuple[Any, ...] + entity_to_unit: tuple[tuple[str, int, datetime | None], ...] + unit_to_entity_ids: tuple[tuple[str, tuple[Any, ...]], ...] diff --git a/core/dataplane/hms_api/engine/entity_resolver.py b/core/dataplane/hms_api/engine/entity_resolver.py index fbf2e88..bc5d772 100644 --- a/core/dataplane/hms_api/engine/entity_resolver.py +++ b/core/dataplane/hms_api/engine/entity_resolver.py @@ -12,12 +12,22 @@ from dataclasses import dataclass, field from datetime import UTC, datetime from difflib import SequenceMatcher -from typing import Any, Final +from typing import TYPE_CHECKING, Any, Final from .db_utils import acquire_with_retry +from .entity_resolution_contracts import ( + EntityOccurrenceBinding, + EntityResolutionReadPlan, + ExistingEntityBinding, + FinalizedEntityResolution, + UnresolvedEntityDescriptor, +) from .memory_engine import fq_table from .retain.entity_labels import build_labels_lookup as _build_labels_lookup_from_config +if TYPE_CHECKING: + from .db.ops import DataAccessOps + logger = logging.getLogger(__name__) @@ -47,17 +57,18 @@ class _EntityStatAgg: # Sentinel distinguishing "key not in dict" from "key present with value None". -# Needed when merging event_date across duplicate unit rows: legacy two-tuple +# Needed when merging event_date across duplicate unit rows: two-tuple # callers surface `None`, which must not clobber a real datetime from another # caller for the same unit. _SENTINEL_MISSING: Final = object() +_ORACLE_IN_CHUNK_SIZE: Final = 900 def _later_date(a: datetime | None, b: datetime | None) -> datetime | None: """Return whichever of ``a`` / ``b`` is later (None loses to any datetime). - Used to fold duplicate co-occurrence pairs across a retain batch: legacy - two-tuple callers surface ``None``, which must not clobber a real + Used to fold duplicate co-occurrence pairs across a retain batch: two-tuple + callers surface ``None``, which must not clobber a real datetime that arrived for the same pair from an aware caller. """ if a is None: @@ -103,7 +114,7 @@ def __init__(self, pool: Any, entity_lookup: str = "full"): self.entity_lookup = entity_lookup self._pg_trgm_checked = False # Backend-specific operations — accessed via pool.ops (Django pattern). - self._ops = pool.ops if pool is not None else None + self._ops: DataAccessOps | None = pool.ops if pool is not None else None # Keyed by asyncio task id so concurrent retain batches never mix their # pending updates. flush_pending_stats() pops only the calling task's items. self._pending_stats: dict[int, list[_EntityStat]] = {} @@ -114,6 +125,13 @@ def _task_key(self) -> int: task = asyncio.current_task() return id(task) if task is not None else 0 + def _require_ops(self) -> "DataAccessOps": + """Return backend operations for methods that require a database pool.""" + + if self._ops is None: + raise RuntimeError("entity resolution database operations require a configured pool") + return self._ops + def discard_pending_stats(self) -> None: """ Discard accumulated entity stats and co-occurrence counts for the current task. @@ -238,6 +256,85 @@ async def resolve_entities_batch( conn, bank_id, entities_data, context, unit_event_date, taxonomy_lookup ) + async def plan_entities_batch( + self, + bank_id: str, + entities_data: list[dict], + context: str, + unit_event_date, + conn=None, + entity_labels: list | None = None, + ) -> EntityResolutionReadPlan: + """Score entity candidates without creating rows or queuing statistics. + + Retain calls this method on its Phase-1 connection. Missing canonical + rows remain immutable descriptors until ``finalize_entity_read_plan`` is + invoked on the core unit-of-work connection. + """ + + if not entities_data: + return EntityResolutionReadPlan(bank_id=bank_id, occurrences=()) + taxonomy_lookup = self._build_labels_lookup(entity_labels) + if conn is None: + async with acquire_with_retry(self.pool) as conn: + return await self._plan_entities_batch_impl( + conn, bank_id, entities_data, context, unit_event_date, taxonomy_lookup + ) + return await self._plan_entities_batch_impl( + conn, bank_id, entities_data, context, unit_event_date, taxonomy_lookup + ) + + async def _plan_entities_batch_impl( + self, + conn, + bank_id: str, + entities_data: list[dict], + context: str, + unit_event_date, + taxonomy_lookup: set[str] | None = None, + ) -> EntityResolutionReadPlan: + del context, taxonomy_lookup # Reserved for future provider-neutral strategies. + if self.entity_lookup == "trigram": + backend_strategy = self._require_ops().get_entity_resolution_strategy() + if backend_strategy == "oracle_fuzzy": + return await self._resolve_entities_batch_oracle_fuzzy( + conn, + bank_id, + entities_data, + unit_event_date, + read_plan=True, + ) + if not self._pg_trgm_checked: + self._pg_trgm_checked = True + has_trgm = await conn.fetchval("SELECT EXISTS(SELECT 1 FROM pg_extension WHERE extname = 'pg_trgm')") + if not has_trgm: + logger.warning( + "pg_trgm extension is not available — falling back to 'full' " + "entity lookup strategy. Install pg_trgm for faster entity resolution." + ) + self.entity_lookup = "full" + return await self._resolve_entities_batch_full( + conn, + bank_id, + entities_data, + unit_event_date, + read_plan=True, + ) + return await self._resolve_entities_batch_trigram( + conn, + bank_id, + entities_data, + unit_event_date, + read_plan=True, + ) + return await self._resolve_entities_batch_full( + conn, + bank_id, + entities_data, + unit_event_date, + read_plan=True, + ) + async def _resolve_entities_batch_impl( self, conn, @@ -250,7 +347,7 @@ async def _resolve_entities_batch_impl( if self.entity_lookup == "trigram": # Route to backend-specific fuzzy strategy. # Non-PG backends (Oracle) use UTL_MATCH instead of pg_trgm. - backend_strategy = self._ops.get_entity_resolution_strategy() + backend_strategy = self._require_ops().get_entity_resolution_strategy() if backend_strategy == "oracle_fuzzy": return await self._resolve_entities_batch_oracle_fuzzy(conn, bank_id, entities_data, unit_event_date) # Auto-detect pg_trgm availability on first call and fall back to @@ -271,8 +368,14 @@ async def _resolve_entities_batch_impl( return await self._resolve_entities_batch_full(conn, bank_id, entities_data, unit_event_date) async def _resolve_entities_batch_full( - self, conn, bank_id: str, entities_data: list[dict], unit_event_date - ) -> list[str]: + self, + conn, + bank_id: str, + entities_data: list[dict], + unit_event_date, + *, + read_plan: bool = False, + ) -> list[str] | EntityResolutionReadPlan: """Original strategy: load all bank entities then match in Python.""" # Query ALL candidates for this bank all_entities = await conn.fetch( @@ -337,13 +440,23 @@ async def _resolve_entities_batch_full( matching.append((ent_id, canonical_name, metadata, last_seen, mention_count)) all_candidates[entity_text] = matching + if read_plan: + return self._build_entity_read_plan( + bank_id, entities_data, unit_event_date, all_candidates, cooccurrence_map + ) return await self._resolve_from_candidates( conn, bank_id, entities_data, unit_event_date, all_candidates, cooccurrence_map ) async def _resolve_entities_batch_trigram( - self, conn, bank_id: str, entities_data: list[dict], unit_event_date - ) -> list[str]: + self, + conn, + bank_id: str, + entities_data: list[dict], + unit_event_date, + *, + read_plan: bool = False, + ) -> list[str] | EntityResolutionReadPlan: """ Trigram strategy: fetch only similar candidates per entity name using pg_trgm. @@ -417,13 +530,23 @@ async def _resolve_entities_batch_trigram( if eid1 in id_to_name: cooccurrence_map[eid2].add(id_to_name[eid1]) + if read_plan: + return self._build_entity_read_plan( + bank_id, entities_data, unit_event_date, all_candidates, cooccurrence_map + ) return await self._resolve_from_candidates( conn, bank_id, entities_data, unit_event_date, all_candidates, cooccurrence_map ) async def _resolve_entities_batch_oracle_fuzzy( - self, conn: Any, bank_id: str, entities_data: list[dict], unit_event_date: datetime | None - ) -> list[str]: + self, + conn: Any, + bank_id: str, + entities_data: list[dict], + unit_event_date: datetime | None, + *, + read_plan: bool = False, + ) -> list[str] | EntityResolutionReadPlan: """ Oracle strategy: fetch similar candidates using UTL_MATCH.JARO_WINKLER_SIMILARITY. @@ -463,7 +586,13 @@ async def _resolve_entities_batch_oracle_fuzzy( e, ) self.entity_lookup = "full" - return await self._resolve_entities_batch_full(conn, bank_id, entities_data, unit_event_date) + return await self._resolve_entities_batch_full( + conn, + bank_id, + entities_data, + unit_event_date, + read_plan=read_plan, + ) # Group candidates by query_text (same structure as trigram strategy) all_candidates: dict[str, list] = {t: [] for t in entity_texts} @@ -505,10 +634,199 @@ async def _resolve_entities_batch_oracle_fuzzy( if eid1 in id_to_name: cooccurrence_map[eid2].add(id_to_name[eid1]) + if read_plan: + return self._build_entity_read_plan( + bank_id, entities_data, unit_event_date, all_candidates, cooccurrence_map + ) return await self._resolve_from_candidates( conn, bank_id, entities_data, unit_event_date, all_candidates, cooccurrence_map ) + @staticmethod + def _build_entity_read_plan( + bank_id: str, + entities_data: list[dict], + unit_event_date, + all_candidates: dict[str, list], + cooccurrence_map: dict[str, set[str]], + ) -> EntityResolutionReadPlan: + """Apply established candidate scoring while deferring every database write.""" + + occurrences: list[EntityOccurrenceBinding] = [] + existing: list[ExistingEntityBinding] = [] + unresolved: list[UnresolvedEntityDescriptor] = [] + + for idx, entity_data in enumerate(entities_data): + entity_text = entity_data["text"] + occurrence_key = entity_data.get("occurrence_key") + unit_key = entity_data.get("unit_key") + local_index = entity_data.get("local_index") + if not isinstance(occurrence_key, str) or not occurrence_key: + raise ValueError(f"entities_data[{idx}] requires a stable occurrence_key") + if not isinstance(unit_key, str) or not unit_key: + raise ValueError(f"entities_data[{idx}] requires a stable unit_key") + if isinstance(local_index, bool) or not isinstance(local_index, int) or local_index < 0: + raise ValueError(f"entities_data[{idx}] requires a non-negative local_index") + + event_date = entity_data.get("event_date", unit_event_date) + occurrences.append( + EntityOccurrenceBinding( + occurrence_key=occurrence_key, + unit_key=unit_key, + local_index=local_index, + event_date=event_date, + ) + ) + nearby_entities = entity_data.get("nearby_entities", []) + nearby_entity_set = { + item["text"].lower() for item in nearby_entities if item.get("text") and item["text"] != entity_text + } + best_candidate = None + best_score = 0.0 + for candidate_id, canonical_name, _metadata, last_seen, _mention_count in all_candidates.get( + entity_text, [] + ): + score = SequenceMatcher(None, entity_text.lower(), canonical_name.lower()).ratio() * 0.5 + if nearby_entity_set: + overlap = len(nearby_entity_set & cooccurrence_map.get(candidate_id, set())) + score += (overlap / len(nearby_entity_set)) * 0.3 + if last_seen and event_date: + event_date_utc = event_date if event_date.tzinfo else event_date.replace(tzinfo=UTC) + last_seen_utc = last_seen if last_seen.tzinfo else last_seen.replace(tzinfo=UTC) + days_diff = abs((event_date_utc - last_seen_utc).total_seconds() / 86400) + if days_diff < 7: + score += max(0, 1.0 - (days_diff / 7)) * 0.2 + if score > best_score: + best_score = score + best_candidate = candidate_id + + if best_score > 0.6: + existing.append(ExistingEntityBinding(occurrence_key, best_candidate)) + continue + + unresolved.append( + UnresolvedEntityDescriptor( + occurrence_key=occurrence_key, + canonical_name=entity_text, + entity_type=entity_data.get("type") or "CONCEPT", + event_date=event_date, + nearby_occurrence_keys=tuple(entity_data.get("nearby_occurrence_keys") or ()), + validated_labels=tuple(entity_data.get("validated_labels") or ()), + ) + ) + + return EntityResolutionReadPlan( + bank_id=bank_id, + occurrences=tuple(occurrences), + existing_bindings=tuple(existing), + unresolved_descriptors=tuple(unresolved), + ) + + async def finalize_entity_read_plan( + self, + conn: Any, + bank_id: str, + plan: EntityResolutionReadPlan, + *, + entities_table: str | None = None, + ) -> FinalizedEntityResolution: + """Create unresolved entities on the core transaction connection. + + Only indexed exact-name reads occur here. Candidate scans and scoring + have already completed in Phase 1. + """ + + if plan.bank_id != bank_id: + raise ValueError("entity read plan bank does not match the core write bank") + table = entities_table or fq_table("entities") + resolved_by_occurrence = {binding.occurrence_key: binding.entity_id for binding in plan.existing_bindings} + + existing_ids = tuple(dict.fromkeys(binding.entity_id for binding in plan.existing_bindings)) + if existing_ids: + backend_type = getattr(conn, "backend_type", "postgresql") + batch_size = ( + _ORACLE_IN_CHUNK_SIZE + if isinstance(backend_type, str) and backend_type.lower() == "oracle" + else len(existing_ids) + ) + found_ids: set[str] = set() + for start in range(0, len(existing_ids), batch_size): + rows = await conn.fetch( + f"SELECT id FROM {table} WHERE bank_id = $1 AND id = ANY($2::uuid[])", + bank_id, + list(existing_ids[start : start + batch_size]), + ) + found_ids.update(str(row["id"]) for row in rows) + missing_ids = [entity_id for entity_id in existing_ids if str(entity_id) not in found_ids] + if missing_ids: + raise ValueError(f"entity read plan contains missing or cross-bank IDs: {missing_ids!r}") + + descriptors = plan.unresolved_descriptors + pending_date_by_occurrence = { + occurrence.occurrence_key: occurrence.event_date for occurrence in plan.occurrences + } + if descriptors: + ops = self._require_ops() + # Preserve established grouping semantics: Python lower() defines + # in-batch groups, the first occurrence supplies spelling/first_seen, + # and groups are inserted in lower-key order. PostgreSQL LOWER remains + # authoritative for conflicts between Python-distinct Unicode spellings. + groups: dict[str, list[UnresolvedEntityDescriptor]] = {} + for descriptor in descriptors: + groups.setdefault(descriptor.canonical_name.lower(), []).append(descriptor) + sorted_groups = sorted(groups.items()) + names = [group[0].canonical_name for _name_lower, group in sorted_groups] + dates = [group[0].event_date for _name_lower, group in sorted_groups] + await ops.bulk_insert_entities(conn, table, bank_id, names, dates) + rows = await ops.fetch_missing_entity_ids(conn, table, bank_id, names) + id_by_input_name: dict[str, Any] = {} + for row in rows: + input_name = row.get("input_name") if hasattr(row, "get") else row["input_name"] + if input_name is not None: + id_by_input_name[input_name] = row["id"] + + missing_names = [name for name in names if name not in id_by_input_name] + if missing_names: + raise ValueError( + "transactional entity finalization could not resolve canonical names " + f"{missing_names!r} in bank {bank_id!r}" + ) + for _name_lower, group in sorted_groups: + representative = group[0] + entity_id = id_by_input_name[representative.canonical_name] + for descriptor in group: + resolved_by_occurrence[descriptor.occurrence_key] = entity_id + pending_date_by_occurrence[descriptor.occurrence_key] = representative.event_date + + occurrence_keys = tuple(item.occurrence_key for item in plan.occurrences) + if set(resolved_by_occurrence) != set(occurrence_keys): + raise ValueError("transactional entity finalization returned an incomplete occurrence mapping") + resolved_ids = tuple(resolved_by_occurrence[key] for key in occurrence_keys) + if any(entity_id is None or str(entity_id) == "" for entity_id in resolved_ids): + raise ValueError("transactional entity finalization returned an empty entity ID") + + unit_to_ids: dict[str, list[Any]] = {} + entity_to_unit: list[tuple[str, int, datetime | None]] = [] + pending: list[_EntityStat] = [] + for occurrence, entity_id in zip(plan.occurrences, resolved_ids, strict=True): + entity_to_unit.append((occurrence.unit_key, occurrence.local_index, occurrence.event_date)) + unit_to_ids.setdefault(occurrence.unit_key, []).append(entity_id) + pending.append( + _EntityStat( + entity_id=entity_id, + event_date=pending_date_by_occurrence[occurrence.occurrence_key], + ) + ) + + # Queue stats only after every occurrence has a verified in-bank ID. The + # service discards this task-local buffer on core rollback/ownership loss. + self._pending_stats.setdefault(self._task_key(), []).extend(pending) + return FinalizedEntityResolution( + resolved_entity_ids=resolved_ids, + entity_to_unit=tuple(entity_to_unit), + unit_to_entity_ids=tuple((unit_key, tuple(ids)) for unit_key, ids in unit_to_ids.items()), + ) + async def _resolve_from_candidates( self, conn, @@ -587,13 +905,15 @@ async def _resolve_from_candidates( # Existing entities: IDs already known from the candidate SELECT above. # No in-transaction UPDATE — mention_count/last_seen are stats deferred to - # flush_pending_stats() which the orchestrator calls after the transaction. + # flush_pending_stats(), which the Retain service calls after the transaction. pending: list[_EntityStat] = list(entities_to_update) # New entities: INSERT with DO NOTHING to avoid row locks on concurrent races. # ON CONFLICT DO NOTHING returns nothing for rows that conflicted; we handle # that rare case with a fallback SELECT. if entities_to_create: + ops = self._require_ops() + # Group by lowercase name — deduplicate within the batch. @dataclass class _NameGroup: @@ -618,7 +938,7 @@ class _NameGroup: # truth for mention counting (one stat per original mention in the batch). entities_table = fq_table("entities") - id_by_name = await self._ops.bulk_insert_entities( + id_by_name = await ops.bulk_insert_entities( conn, entities_table, bank_id, @@ -638,7 +958,7 @@ class _NameGroup: # a NOT NULL constraint violation on unit_entities.entity_id. missing_original = [g.name for name_lower, g in sorted_groups if name_lower not in id_by_name] if missing_original: - existing_rows = await self._ops.fetch_missing_entity_ids( + existing_rows = await ops.fetch_missing_entity_ids( conn, entities_table, bank_id, @@ -662,7 +982,7 @@ class _NameGroup: entity_ids[original_idx] = entity_id pending.append(_EntityStat(entity_id=entity_id, event_date=g.event_date)) - # Accumulate into the resolver's pending list; the orchestrator flushes + # Accumulate into the resolver's pending list; the Retain service flushes # these with await entity_resolver.flush_pending_stats() after the txn. key = self._task_key() self._pending_stats.setdefault(key, []).extend(pending) @@ -846,6 +1166,7 @@ async def link_unit_to_entity(self, unit_id: str, entity_id: str): entity_id: Entity ID """ async with acquire_with_retry(self.pool) as conn: + ops = self._require_ops() # Insert unit-entity link await conn.execute( f""" @@ -856,7 +1177,7 @@ async def link_unit_to_entity(self, unit_id: str, entity_id: str): unit_id, entity_id, ) - await self._ops.refresh_entity_fact_counts( + await ops.refresh_entity_fact_counts( conn, fq_table("entities"), fq_table("unit_entities"), @@ -925,17 +1246,21 @@ async def link_units_to_entities_batch( observed in that unit advances to the event time instead of ``now()``, which matters for backfilled corpora where ingest time is a single spike unrelated to the underlying timeline. - Legacy two-tuples remain accepted. + Backward-compatible two-tuples remain accepted. conn: Optional connection to use (if None, acquires from pool) """ if not unit_entity_pairs: return # Normalize to 3-tuples internally so downstream code doesn't branch. - normalized: list[tuple[str, str, datetime | None]] = [ - (t[0], t[1], t[2] if len(t) >= 3 else None) # type: ignore[misc] - for t in unit_entity_pairs - ] + normalized: list[tuple[str, str, datetime | None]] = [] + for pair in unit_entity_pairs: + if len(pair) == 2: + unit_id, entity_id = pair + event_date = None + else: + unit_id, entity_id, event_date = pair + normalized.append((unit_id, entity_id, event_date)) if conn is None: async with acquire_with_retry(self.pool) as conn: @@ -944,19 +1269,20 @@ async def link_units_to_entities_batch( return await self._link_units_to_entities_batch_impl(conn, normalized) async def _link_units_to_entities_batch_impl(self, conn, unit_entity_pairs: list[tuple[str, str, datetime | None]]): + ops = self._require_ops() # Sorted bulk insert to prevent deadlocks from inconsistent lock ordering # across concurrent transactions on the unit_entities unique index. sorted_pairs = sorted(unit_entity_pairs, key=lambda t: (t[0], t[1])) unit_ids = [p[0] for p in sorted_pairs] entity_ids = [p[1] for p in sorted_pairs] - await self._ops.bulk_insert_unit_entities( + await ops.bulk_insert_unit_entities( conn, fq_table("unit_entities"), unit_ids, entity_ids, ) - await self._ops.refresh_entity_fact_counts( + await ops.refresh_entity_fact_counts( conn, fq_table("entities"), fq_table("unit_entities"), @@ -966,8 +1292,8 @@ async def _link_units_to_entities_batch_impl(self, conn, unit_entity_pairs: list # Build maps keyed by unit_id: # unit_to_entities: entity set per unit (for the co-occurrence cross-product) # unit_event_date: event time per unit (propagated onto every pair from that unit) - # When a unit shows up more than once with conflicting event_dates (legacy - # callers passing None interleaved with aware callers), prefer the first + # When a unit shows up more than once with conflicting event_dates + # (two-tuple callers interleaved with aware callers), prefer the first # non-None value so we don't accidentally erase an explicit timestamp. unit_to_entities: dict[str, set[str]] = {} unit_event_date: dict[str, datetime | None] = {} diff --git a/core/dataplane/hms_api/engine/ingestion/__init__.py b/core/dataplane/hms_api/engine/ingestion/__init__.py new file mode 100644 index 0000000..6cdf364 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/__init__.py @@ -0,0 +1,61 @@ +"""Application boundary for the Retain ingestion pipeline.""" + +from .contracts import ( + RetainExecutionContext, + RetainInvocation, + RetainOperationInactiveError, + RetainOutcome, + RetainPipeline, +) +from .domain import ( + ChunkPlan, + ChunkPolicy, + ContentItem, + ContentOrigin, + DocumentChangeKind, + DocumentChangePlan, + DocumentIntent, + EventDateState, + EventDateValue, + ExistingChunkFingerprint, + UpdateMode, +) +from .service import ( + RetainCheckpointRecoveryError, + RetainDatabaseUnsupportedError, + RetainError, + RetainExtractionModeUnsupportedError, + RetainOwnershipLostError, + RetainPipelineService, + RetainPublicationAborted, + RetainResultMappingError, + RetainUnsupportedError, +) + +__all__ = [ + "ChunkPlan", + "ChunkPolicy", + "ContentItem", + "ContentOrigin", + "DocumentChangeKind", + "DocumentChangePlan", + "DocumentIntent", + "EventDateState", + "EventDateValue", + "ExistingChunkFingerprint", + "RetainCheckpointRecoveryError", + "RetainDatabaseUnsupportedError", + "RetainError", + "RetainExecutionContext", + "RetainExtractionModeUnsupportedError", + "RetainInvocation", + "RetainOperationInactiveError", + "RetainOutcome", + "RetainOwnershipLostError", + "RetainPipeline", + "RetainPipelineService", + "RetainPublicationAborted", + "RetainResultMappingError", + "RetainUnsupportedError", + "UpdateMode", +] diff --git a/core/dataplane/hms_api/engine/ingestion/adapters/__init__.py b/core/dataplane/hms_api/engine/ingestion/adapters/__init__.py new file mode 100644 index 0000000..596df2f --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/adapters/__init__.py @@ -0,0 +1,10 @@ +"""Runtime adapters for the Retain ingestion pipeline.""" + +from .embedding_model import EmbeddingModelAdapter +from .postgres_fresh_ownership import FreshDocumentOwnershipConflict, FreshPostgresDocumentOwnership + +__all__ = [ + "EmbeddingModelAdapter", + "FreshDocumentOwnershipConflict", + "FreshPostgresDocumentOwnership", +] diff --git a/core/dataplane/hms_api/engine/ingestion/adapters/embedding_model.py b/core/dataplane/hms_api/engine/ingestion/adapters/embedding_model.py new file mode 100644 index 0000000..ce1ebaa --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/adapters/embedding_model.py @@ -0,0 +1,30 @@ +"""Adapter from the async embedding port to the configured model API.""" + +from __future__ import annotations + +from collections.abc import Sequence +from typing import Any + +from ...retain import embedding_processing +from ..projection.embeddings import EmbeddingVector + + +class EmbeddingModelAdapter: + """Expose the existing synchronous ``encode`` model as an async port. + + The delegated helper already moves CPU-bound model work off the event loop + and validates one-output-per-input cardinality. The ingestion pipeline + performs its own + cardinality check as a second boundary guard before positional projection. + """ + + def __init__(self, embeddings_model: Any) -> None: + if embeddings_model is None: + raise ValueError("Retain requires an initialized embeddings model") + self._embeddings_model = embeddings_model + + async def embed_batch(self, texts: tuple[str, ...]) -> Sequence[EmbeddingVector]: + return await embedding_processing.generate_embeddings_batch( + self._embeddings_model, + list(texts), + ) diff --git a/core/dataplane/hms_api/engine/ingestion/adapters/oracle_semantic.py b/core/dataplane/hms_api/engine/ingestion/adapters/oracle_semantic.py new file mode 100644 index 0000000..7537a68 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/adapters/oracle_semantic.py @@ -0,0 +1,77 @@ +"""Oracle semantic-neighbor reads used while retaining memories.""" + +from __future__ import annotations + +import json +from array import array +from collections.abc import Sequence +from typing import Any + +from ...memory_engine import fq_table + + +def _vector_bind(embedding: Any) -> array: + if isinstance(embedding, str): + try: + embedding = json.loads(embedding) + except json.JSONDecodeError as exc: + raise ValueError("Oracle semantic planning received an invalid vector string") from exc + if not isinstance(embedding, Sequence): + raise TypeError("Oracle semantic planning requires a vector sequence") + if not embedding: + raise ValueError("Oracle semantic planning requires a non-empty vector") + try: + return array("f", (float(value) for value in embedding)) + except (TypeError, ValueError) as exc: + raise ValueError("Oracle semantic planning requires numeric vector values") from exc + + +async def compute_oracle_semantic_links_ann( + connection: Any, + bank_id: str, + unit_ids: Sequence[str], + embeddings: Sequence[Any], + *, + fact_types: Sequence[str] | None = None, + top_k: int = 50, + threshold: float = 0.7, +) -> list[tuple[Any, ...]]: + """Find Oracle VECTOR neighbors without PostgreSQL temp tables or arrays.""" + + if isinstance(top_k, bool) or not isinstance(top_k, int) or top_k <= 0: + raise ValueError("top_k must be a positive integer") + if not unit_ids or not embeddings: + return [] + if len(unit_ids) != len(embeddings): + raise ValueError("Oracle semantic planning requires one embedding per unit") + if fact_types is None: + fact_types = ("world",) * len(unit_ids) + if len(fact_types) != len(unit_ids): + raise ValueError("Oracle semantic planning requires one fact type per unit") + + links: list[tuple[Any, ...]] = [] + for unit_id, embedding, fact_type in zip(unit_ids, embeddings, fact_types, strict=True): + vector = _vector_bind(embedding) + rows = await connection.fetch( + f""" + SELECT id AS to_id, + 1 - VECTOR_DISTANCE(embedding, $3, COSINE) AS similarity + FROM {fq_table("memory_units")} + WHERE bank_id = $1 + AND fact_type = $2 + AND embedding IS NOT NULL + ORDER BY VECTOR_DISTANCE(embedding, $3, COSINE) + FETCH FIRST {top_k} ROWS ONLY + """, + bank_id, + fact_type, + vector, + ) + for row in rows: + similarity = float(min(1.0, max(0.0, row["similarity"]))) + if similarity >= threshold: + links.append((unit_id, str(row["to_id"]), "semantic", similarity, None)) + return links + + +__all__ = ["compute_oracle_semantic_links_ann"] diff --git a/core/dataplane/hms_api/engine/ingestion/adapters/postgres_fresh_ownership.py b/core/dataplane/hms_api/engine/ingestion/adapters/postgres_fresh_ownership.py new file mode 100644 index 0000000..92a447c --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/adapters/postgres_fresh_ownership.py @@ -0,0 +1,109 @@ +"""Strict PostgreSQL ownership gate for fresh-document writes.""" + +from __future__ import annotations + +from typing import Any + +from ...schema import fq_table_explicit + + +class FreshDocumentOwnershipConflict(RuntimeError): + """A document appeared after Retain's read-side fresh check.""" + + +class FreshPostgresDocumentOwnership: + """Claim a previously absent document without replacing concurrent data. + + The general full-write ownership adapter intentionally permits replacement. + This stricter adapter uses ``INSERT ... RETURNING`` to turn the + preflight/write race into a typed failure instead of allowing another + writer's just-committed document to be replaced. + """ + + def __init__(self, *, schema: str | None = None) -> None: + self._schema = schema + + async def prepare_first_window(self, connection: Any, *, bank_id: str, document_id: str) -> None: + if not isinstance(bank_id, str) or not bank_id: + raise ValueError("bank_id must be a non-empty string") + if not isinstance(document_id, str) or not document_id: + raise ValueError("document_id must be a non-empty string") + + documents = fq_table_explicit("documents", self._schema) + claimed = await connection.fetchval( + f""" + INSERT INTO {documents} (id, bank_id, original_text, content_hash) + VALUES ($1, $2, '', '__pending__') + ON CONFLICT (id, bank_id) DO NOTHING + RETURNING id + """, + document_id, + bank_id, + ) + if claimed is None: + raise FreshDocumentOwnershipConflict(f"Document {document_id!r} in bank {bank_id!r} is no longer fresh") + + locked = await connection.fetchval( + f""" + SELECT id + FROM {documents} + WHERE id = $1 AND bank_id = $2 + FOR UPDATE + """, + document_id, + bank_id, + ) + if locked is None: # pragma: no cover - same transaction inserted it + raise RuntimeError("Fresh document ownership row disappeared inside its transaction") + + async def validate_later_window( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + expected_content_hash: str, + ) -> bool: + del connection, bank_id, document_id, expected_content_hash + raise RuntimeError("Fresh-document ownership does not support later full-write windows") + + async def validate_unhashed_window( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + ) -> bool: + del connection, bank_id, document_id + raise RuntimeError("Fresh-document ownership does not support existing unhashed rows") + + async def transition_content_hash( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + expected_content_hash: str, + new_content_hash: str, + ) -> bool: + for field_name, value in ( + ("bank_id", bank_id), + ("document_id", document_id), + ("expected_content_hash", expected_content_hash), + ("new_content_hash", new_content_hash), + ): + if not isinstance(value, str) or not value: + raise ValueError(f"{field_name} must be a non-empty string") + updated = await connection.fetchval( + f""" + UPDATE {fq_table_explicit("documents", self._schema)} + SET content_hash = $1, updated_at = now() + WHERE id = $2 AND bank_id = $3 AND content_hash = $4 + RETURNING id + """, + new_content_hash, + document_id, + bank_id, + expected_content_hash, + ) + return updated is not None diff --git a/core/dataplane/hms_api/engine/ingestion/adapters/storage_records.py b/core/dataplane/hms_api/engine/ingestion/adapters/storage_records.py new file mode 100644 index 0000000..b99c66d --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/adapters/storage_records.py @@ -0,0 +1,179 @@ +"""Conversions from ingestion domain objects to durable storage records.""" + +from __future__ import annotations + +import hashlib +from collections import Counter +from collections.abc import Mapping, Sequence +from typing import Any + +from ...retain.types import ChunkMetadata, ExtractedFact, RetainContent +from ..domain import ChunkPlan, ContentItem, EventDateState, thaw_json +from ..extraction import build_content_position_map +from ..projection import MemoryRecord, to_processed_fact + + +def content_to_storage(item: ContentItem) -> RetainContent: + """Snapshot one immutable content item as the current mutable DTO.""" + + metadata = thaw_json(item.metadata) + if not isinstance(metadata, dict): + raise TypeError("ContentItem.metadata must thaw to an object") + + entities: list[dict[str, Any]] = [] + for index, frozen_entity in enumerate(item.entities): + entity = thaw_json(frozen_entity) + if not isinstance(entity, dict): + raise TypeError(f"ContentItem.entities[{index}] must thaw to an object") + entities.append(entity) + + observation_scopes: str | list[list[str]] | None + if isinstance(item.observation_scopes, tuple): + observation_scopes = [list(scope) for scope in item.observation_scopes] + else: + observation_scopes = item.observation_scopes + + return RetainContent( + content=item.content, + context=item.context, + event_date=item.event_date.value, + metadata=metadata, + entities=entities, + tags=list(item.tags), + observation_scopes=observation_scopes, + ) + + +def compute_document_hash(combined_content: str) -> str: + """Hash normalized document text for durable change tracking.""" + + if not isinstance(combined_content, str): + raise TypeError("combined_content must be a string") + from ...retain.fact_extraction import _sanitize_text + + sanitized = _sanitize_text(combined_content) or "" + return hashlib.sha256(sanitized.encode()).hexdigest() + + +def retain_document_metadata(items: Sequence[ContentItem]) -> tuple[dict[str, Any], tuple[str, ...]]: + """Build the document metadata snapshot from normalized content. + + Defaulted timestamps are deliberately omitted: ``event_date`` is recorded + in ``retain_params`` only when the caller supplied a truthy value, even + though extraction itself receives a default timestamp. + """ + + if not items: + return {}, () + + first = items[0] + retain_params: dict[str, Any] = {} + if first.context: + retain_params["context"] = first.context + if first.event_date.state is EventDateState.EXPLICIT and first.event_date.value is not None: + retain_params["event_date"] = first.event_date.value.isoformat() + metadata = thaw_json(first.metadata) + if not isinstance(metadata, dict): + raise TypeError("ContentItem.metadata must thaw to an object") + if metadata: + retain_params["metadata"] = metadata + + seen: set[str] = set() + tags: list[str] = [] + for item in items: + for tag in item.tags: + if tag not in seen: + seen.add(tag) + tags.append(tag) + return retain_params, tuple(tags) + + +def chunks_to_storage( + chunks: Sequence[ChunkPlan], + items: Sequence[ContentItem], + records: Sequence[MemoryRecord], +) -> tuple[ChunkMetadata, ...]: + """Create chunk DTOs with fact counts derived by stable chunk key.""" + + positions = build_content_position_map(items) + fact_counts = Counter(record.chunk_key for record in records) + known_keys = {chunk.chunk_key for chunk in chunks} + unknown_keys = sorted(set(fact_counts) - known_keys) + if unknown_keys: + raise ValueError(f"Projected records reference unknown chunk keys: {unknown_keys!r}") + + result: list[ChunkMetadata] = [] + for chunk in chunks: + try: + content_index = positions[chunk.source_index] + except KeyError as exc: # pragma: no cover - extraction validates this first + raise ValueError(f"Missing content position for source_index={chunk.source_index!r}") from exc + result.append( + ChunkMetadata( + chunk_text=chunk.text, + fact_count=fact_counts[chunk.chunk_key], + content_index=content_index, + chunk_index=chunk.global_index, + ) + ) + return tuple(result) + + +def record_to_extracted_fact(record: MemoryRecord, *, content_index: int) -> ExtractedFact: + """Convert a projected record into the raw DTO required by Phase 2.""" + + if isinstance(content_index, bool) or not isinstance(content_index, int) or content_index < 0: + raise ValueError("content_index must be a non-negative integer") + metadata = thaw_json(record.metadata) + if not isinstance(metadata, dict): + raise TypeError("MemoryRecord.metadata must thaw to an object") + if isinstance(record.observation_scopes, tuple): + observation_scopes = [list(scope) for scope in record.observation_scopes] + else: + observation_scopes = record.observation_scopes + + return ExtractedFact( + fact_text=record.text, + fact_type=record.fact_type, + entities=list(record.entity_mentions), + occurred_start=record.occurred_start, + occurred_end=record.occurred_end, + where=getattr(record, "where", None), + causal_relations=[], + content_index=content_index, + chunk_index=record.global_index, + context=record.context, + mentioned_at=record.mentioned_at, + metadata=metadata, + tags=list(record.tags), + observation_scopes=observation_scopes, + ) + + +def record_to_processed_fact( + record: MemoryRecord, + *, + document_id: str, + content_index: int, + fact_positions: Mapping[str, int] | None = None, +): + """Create a processed DTO before its durable chunk ID is known. + + The stable chunk key is used only as a non-empty boundary placeholder. + ``PersistenceWriter`` replaces it with the exact chunk ID returned by the + chunk upsert inside the same transaction. + """ + + return to_processed_fact( + record, + document_id=document_id, + chunk_id=record.chunk_key, + content_index=content_index, + fact_positions=fact_positions, + ) + + +def content_positions(items: Sequence[ContentItem]) -> Mapping[int | None, int]: + """Expose the validated stable-source-to-current-position mapping.""" + + return build_content_position_map(items) diff --git a/core/dataplane/hms_api/engine/ingestion/change_detection.py b/core/dataplane/hms_api/engine/ingestion/change_detection.py new file mode 100644 index 0000000..5e30dc3 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/change_detection.py @@ -0,0 +1,116 @@ +"""Pure document/chunk change classification for Retain.""" + +from __future__ import annotations + +from collections.abc import Sequence +from datetime import datetime, timezone + +from .domain import ( + ChunkPlan, + DocumentChangeKind, + DocumentChangePlan, + ExistingChunkFingerprint, +) + + +def detect_document_change( + new_chunks: Sequence[ChunkPlan], + existing_chunks: Sequence[ExistingChunkFingerprint], + *, + document_exists: bool, + existing_document_content_hash: str | None, + new_document_content_hash: str | None, + updated_at: datetime | None, + request_started_at: datetime, + policy_compatible: bool, +) -> DocumentChangePlan: + """Classify a document update without performing I/O. + + Chunk identity is the pair ``(global/chunk index, UTF-8 SHA-256)``. Unsafe + or ambiguous stored state falls back to ``FULL`` with a stable reason. A + partial delta is used only when at least one existing chunk remains + unchanged, matching the conservative fallback behavior. + """ + + if not document_exists: + return DocumentChangePlan(kind=DocumentChangeKind.FULL, reason="document_not_found") + + if updated_at is not None and _as_utc(updated_at) > _as_utc(request_started_at): + return DocumentChangePlan( + kind=DocumentChangeKind.STALE_SKIP, + reason="document_updated_after_request_started", + ) + + if not policy_compatible: + return DocumentChangePlan(kind=DocumentChangeKind.FULL, reason="chunk_policy_incompatible") + + if not existing_document_content_hash or not new_document_content_hash: + return DocumentChangePlan(kind=DocumentChangeKind.FULL, reason="missing_document_content_hash") + + if not existing_chunks: + return DocumentChangePlan(kind=DocumentChangeKind.FULL, reason="no_existing_chunks") + + existing_by_index: dict[int, ExistingChunkFingerprint] = {} + for chunk in existing_chunks: + if chunk.chunk_index in existing_by_index: + return DocumentChangePlan(kind=DocumentChangeKind.FULL, reason="duplicate_existing_chunk_index") + if not chunk.content_hash: + return DocumentChangePlan(kind=DocumentChangeKind.FULL, reason="missing_existing_chunk_hash") + existing_by_index[chunk.chunk_index] = chunk + + new_by_index: dict[int, ChunkPlan] = {} + for chunk in new_chunks: + if chunk.global_index in new_by_index: + return DocumentChangePlan(kind=DocumentChangeKind.FULL, reason="duplicate_new_chunk_index") + if not chunk.content_hash: + return DocumentChangePlan(kind=DocumentChangeKind.FULL, reason="missing_new_chunk_hash") + new_by_index[chunk.global_index] = chunk + + unchanged: list[int] = [] + changed: list[int] = [] + added: list[int] = [] + removed: list[int] = [] + + for index, chunk in sorted(new_by_index.items()): + existing = existing_by_index.get(index) + if existing is None: + added.append(index) + elif existing.content_hash == chunk.content_hash: + unchanged.append(index) + else: + changed.append(index) + + for index in sorted(existing_by_index): + if index not in new_by_index: + removed.append(index) + + if not changed and not added and not removed: + if existing_document_content_hash != new_document_content_hash: + return DocumentChangePlan( + kind=DocumentChangeKind.FULL, + reason="document_hash_mismatch_without_chunk_changes", + ) + return DocumentChangePlan( + kind=DocumentChangeKind.METADATA_ONLY, + unchanged=tuple(unchanged), + ) + + # Retain deliberately abandons delta when every stored chunk changed. + # Keeping that rule avoids a delta transaction that is effectively a full + # replacement but has different observation/outbox/recovery semantics. + if not unchanged: + return DocumentChangePlan(kind=DocumentChangeKind.FULL, reason="no_unchanged_chunks") + + return DocumentChangePlan( + kind=DocumentChangeKind.DELTA, + unchanged=tuple(unchanged), + changed=tuple(changed), + added=tuple(added), + removed=tuple(removed), + ) + + +def _as_utc(value: datetime) -> datetime: + if value.tzinfo is None: + return value.replace(tzinfo=timezone.utc) + return value.astimezone(timezone.utc) diff --git a/core/dataplane/hms_api/engine/ingestion/chunking.py b/core/dataplane/hms_api/engine/ingestion/chunking.py new file mode 100644 index 0000000..bf28947 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/chunking.py @@ -0,0 +1,152 @@ +"""Deterministic, side-effect-free chunk planning for Retain.""" + +from __future__ import annotations + +import hashlib +import json +from collections.abc import Sequence +from typing import Any + +from .domain import ChunkPlan, ChunkPolicy, ContentItem + +PLAIN_TEXT_SEPARATORS: tuple[str, ...] = ( + "\n\n", + "\n", + ". ", + "! ", + "? ", + "; ", + ", ", + " ", + "", +) + + +def compute_content_hash(text: str) -> str: + """Return the lowercase SHA-256 digest of the text's UTF-8 bytes.""" + + if not isinstance(text, str): + raise TypeError(f"text must be str, got {type(text).__name__}") + return hashlib.sha256(text.encode("utf-8")).hexdigest() + + +def split_text(text: str, policy: ChunkPolicy) -> tuple[str, ...]: + """Split one content item according to the versioned chunk policy. + + The default behavior follows the Retain chunking contract: + + * text at or below ``max_chars`` is returned byte-for-byte; + * a JSON array containing only objects is treated as a conversation and + split only between complete turns; + * all other text uses ``RecursiveCharacterTextSplitter`` with the + configured ordered separator list. + + A single conversation turn is never split, even if it is larger than the + configured limit. Conversation overlap is deliberately unsupported because + repeating turns would change Retain's fact semantics; callers can disable + conversation mode when character overlap is required. + """ + + if not isinstance(text, str): + raise TypeError(f"text must be str, got {type(text).__name__}") + + # Empty text is represented by one empty chunk. + if len(text) <= policy.max_chars: + return (text,) + + if policy.conversation_mode: + turns = _parse_conversation(text) + if turns is not None: + if policy.overlap: + raise ValueError("conversation chunking does not support overlap") + return _split_conversation(turns, policy.max_chars) + + from langchain_text_splitters import RecursiveCharacterTextSplitter + + splitter = RecursiveCharacterTextSplitter( + chunk_size=policy.max_chars, + chunk_overlap=policy.overlap, + length_function=len, + is_separator_regex=False, + separators=list(PLAIN_TEXT_SEPARATORS), + ) + return tuple(splitter.split_text(text)) + + +def build_chunk_plans( + document_id: str, + items: Sequence[ContentItem], + policy: ChunkPolicy, +) -> tuple[ChunkPlan, ...]: + """Build stable chunk plans for all items in one document. + + ``local_index`` restarts for every content item while ``global_index`` is + assigned synchronously across the complete document. Consequently neither + index depends on task completion order. ``chunk_key`` is an unambiguous, + deterministic composition of document ID, global index, and content hash; + it never relies on Python's process-randomized :func:`hash`. + """ + + if not isinstance(document_id, str) or not document_id: + raise ValueError("document_id must be a non-empty string") + + plans: list[ChunkPlan] = [] + global_index = 0 + for item in items: + for local_index, chunk in enumerate(split_text(item.content, policy)): + content_hash = compute_content_hash(chunk) + chunk_key = _build_chunk_key(document_id, global_index, content_hash) + plans.append( + ChunkPlan( + chunk_key=chunk_key, + source_index=item.source_index, + global_index=global_index, + local_index=local_index, + text=chunk, + content_hash=content_hash, + ) + ) + global_index += 1 + return tuple(plans) + + +def _build_chunk_key(document_id: str, global_index: int, content_hash: str) -> str: + # Prefixing the document length makes the representation unambiguous even + # when a caller-supplied document ID contains colons. + return f"chunk:{len(document_id)}:{document_id}:{global_index}:{content_hash}" + + +def _parse_conversation(text: str) -> list[dict[str, Any]] | None: + try: + parsed = json.loads(text) + except (json.JSONDecodeError, ValueError): + return None + + if isinstance(parsed, list) and all(isinstance(turn, dict) for turn in parsed): + return parsed + return None + + +def _split_conversation(turns: list[dict[str, Any]], max_chars: int) -> tuple[str, ...]: + chunks: list[str] = [] + current_chunk: list[dict[str, Any]] = [] + current_size = 2 # Serialized [] + + for turn in turns: + turn_json = json.dumps(turn, ensure_ascii=False) + turn_size = len(turn_json) + 1 # Comma between adjacent turns. + + if current_size + turn_size > max_chars and current_chunk: + chunks.append(json.dumps(current_chunk, ensure_ascii=False)) + current_chunk = [] + current_size = 2 + + current_chunk.append(turn) + current_size += turn_size + + if current_chunk: + chunks.append(json.dumps(current_chunk, ensure_ascii=False)) + + # The branch is normally reached only for text larger than max_chars, but + # retaining the fallback keeps the helper total for an empty JSON array. + return tuple(chunks) if chunks else (json.dumps(turns, ensure_ascii=False),) diff --git a/core/dataplane/hms_api/engine/ingestion/contracts.py b/core/dataplane/hms_api/engine/ingestion/contracts.py new file mode 100644 index 0000000..9a20df4 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/contracts.py @@ -0,0 +1,73 @@ +"""Contracts for the Retain ingestion pipeline.""" + +from __future__ import annotations + +import asyncio +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from typing import Any, Protocol + +from ..response_models import TokenUsage +from ..retain.types import RetainContentDict + +OutboxCallback = Callable[[Any], Awaitable[None]] +CoreCommitCallback = Callable[[Any, tuple[tuple[str, ...], ...]], Awaitable[None]] + + +class RetainOperationInactiveError(RuntimeError): + """A tracked Retain operation became terminal before its core write.""" + + +@dataclass(frozen=True, slots=True) +class RetainInvocation: + """Raw invocation received below ``MemoryEngine`` batch handling.""" + + bank_id: str + raw_contents: tuple[RetainContentDict, ...] + request_context: Any + batch_document_id: str | None = None + is_first_batch: bool = True + fact_type_override: str | None = None + document_tags: tuple[str, ...] | None = None + operation_id: str | None = None + outbox_callback: OutboxCallback | None = None + strategy: str | None = None + sanitize_log_identifiers: bool = False + + +@dataclass(frozen=True, slots=True) +class RetainExecutionContext: + """Resolved dependencies for one Retain submission shard.""" + + pool: Any + embeddings_model: Any + llm_config: Any + entity_resolver: Any + format_date_fn: Callable[..., str] + resolved_config: Any + schema: str | None = None + db_semaphore: asyncio.Semaphore | None = None + + +@dataclass(frozen=True, slots=True) +class RetainOutcome: + """Result returned by the Retain pipeline.""" + + unit_ids_by_input: list[list[str]] + usage: TokenUsage + processed_content_tokens: int | None + + def as_tuple(self) -> tuple[list[list[str]], TokenUsage, int | None]: + """Return the tuple consumed by ``MemoryEngine``.""" + + return self.unit_ids_by_input, self.usage, self.processed_content_tokens + + +class RetainPipeline(Protocol): + """Interface implemented by the Retain service.""" + + async def retain( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + ) -> RetainOutcome: ... diff --git a/core/dataplane/hms_api/engine/ingestion/document_planner.py b/core/dataplane/hms_api/engine/ingestion/document_planner.py new file mode 100644 index 0000000..a0061d0 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/document_planner.py @@ -0,0 +1,210 @@ +"""Pure document grouping and append planning for Retain.""" + +from __future__ import annotations + +import uuid +from collections.abc import Callable, Iterable, Sequence +from dataclasses import replace +from datetime import UTC, datetime +from typing import TypeGuard + +from .domain import ContentItem, ContentOrigin, DocumentIntent, UpdateMode, freeze_json +from .normalization import Clock, parse_event_date + +DocumentIdFactory = Callable[[], str] + + +def _uuid_document_id() -> str: + return str(uuid.uuid4()) + + +def _utcnow() -> datetime: + return datetime.now(UTC) + + +def _is_valid_document_id(document_id: object) -> TypeGuard[str]: + return isinstance(document_id, str) and bool(document_id.strip()) + + +def _generated_document_id(id_factory: DocumentIdFactory, *, reserved: set[str]) -> str: + document_id = id_factory() + if not _is_valid_document_id(document_id): + raise ValueError("id_factory must return a non-empty document ID") + if document_id in reserved: + raise ValueError(f"id_factory returned a duplicate document ID: {document_id!r}") + return document_id + + +def _resolve_shared_document_id( + items: Sequence[ContentItem], + *, + batch_document_id: str | None, + recovered_document_id: str | None, + id_factory: DocumentIdFactory, +) -> str: + if _is_valid_document_id(batch_document_id): + return batch_document_id + if _is_valid_document_id(recovered_document_id): + return recovered_document_id + if any(item.update_mode is UpdateMode.APPEND for item in items): + raise ValueError("update_mode='append' requires a valid document ID") + return _generated_document_id(id_factory, reserved=set()) + + +def _validate_update_mode(document_id: str, items: Sequence[ContentItem]) -> UpdateMode: + if not items: + raise ValueError("A document intent requires at least one content item") + + update_mode = items[0].update_mode + if any(item.update_mode is not update_mode for item in items[1:]): + modes = ", ".join(sorted({item.update_mode.value for item in items})) + raise ValueError(f"Conflicting update_mode values for document {document_id!r}: {modes}") + return update_mode + + +def _build_intent(document_id: str, items: Sequence[ContentItem]) -> DocumentIntent: + update_mode = _validate_update_mode(document_id, items) + if update_mode is UpdateMode.APPEND and not _is_valid_document_id(document_id): + raise ValueError("update_mode='append' requires a valid document ID") + + source_indices = tuple(item.source_index for item in items if item.source_index is not None) + return DocumentIntent( + document_id=document_id, + items=tuple(items), + source_indices=source_indices, + expected_input_slots=source_indices, + update_mode=update_mode, + ) + + +def plan_documents( + items: Iterable[ContentItem], + *, + batch_document_id: str | None = None, + recovered_document_id: str | None = None, + id_factory: DocumentIdFactory = _uuid_document_id, +) -> tuple[DocumentIntent, ...]: + """Group normalized Retain items into deterministic document intents. + + Explicit IDs determine grouping. One explicit ID absorbs items without an + ID, while multiple explicit IDs leave every missing-ID item as its own + generated document. If every item is missing an ID, the batch ID wins over + a recovered operation ID, which in turn wins over the injected factory. + """ + + normalized_items = tuple(items) + if not normalized_items: + return () + + # A non-empty batch-level document ID is authoritative before per-item IDs + # are inspected. ``MemoryEngine`` only fills missing item IDs, so a caller + # can still reach this boundary with both values present. + if _is_valid_document_id(batch_document_id): + _validate_update_mode(batch_document_id, normalized_items) + return (_build_intent(batch_document_id, normalized_items),) + + explicit_groups: dict[str, list[ContentItem]] = {} + for item in normalized_items: + if _is_valid_document_id(item.document_id): + explicit_groups.setdefault(item.document_id, []).append(item) + + explicit_ids = tuple(explicit_groups) + + if not explicit_ids: + effective_id_hint = ( + batch_document_id + if _is_valid_document_id(batch_document_id) + else recovered_document_id + if _is_valid_document_id(recovered_document_id) + else "" + ) + _validate_update_mode(effective_id_hint, normalized_items) + document_id = _resolve_shared_document_id( + normalized_items, + batch_document_id=batch_document_id, + recovered_document_id=recovered_document_id, + id_factory=id_factory, + ) + return (_build_intent(document_id, normalized_items),) + + if len(explicit_ids) == 1: + return (_build_intent(explicit_ids[0], normalized_items),) + + # Validate every explicit group before invoking the ID factory. Factory + # calls are observable injected side effects, so an invalid batch must fail + # before generating IDs for otherwise unrelated missing-ID items. + for document_id, group_items in explicit_groups.items(): + _validate_update_mode(document_id, group_items) + + groups: dict[str, list[ContentItem]] = {} + reserved = set(explicit_ids) + for item in normalized_items: + if _is_valid_document_id(item.document_id): + document_id = item.document_id + else: + if item.update_mode is UpdateMode.APPEND: + raise ValueError("update_mode='append' requires a valid document ID") + document_id = _generated_document_id(id_factory, reserved=reserved) + reserved.add(document_id) + groups.setdefault(document_id, []).append(item) + + return tuple(_build_intent(document_id, group_items) for document_id, group_items in groups.items()) + + +def make_append_synthetic_item( + existing_content: str, + *, + document_id: str, + template: ContentItem, + clock: Clock = _utcnow, +) -> ContentItem: + """Create the non-result-bearing item that represents existing text. + + Append copies only context and tags from the first submitted item. + Its missing event date is defaulted at append execution time, while + metadata, declared entities, and observation scopes remain empty. Keeping + that behavior here avoids a structural refactor silently changing what is + projected from unchanged source text. + """ + + if template.update_mode is not UpdateMode.APPEND: + raise ValueError("An append synthetic item requires an append-mode template") + if not _is_valid_document_id(document_id): + raise ValueError("update_mode='append' requires a valid document ID") + if not isinstance(existing_content, str): + raise TypeError("existing_content must be a string") + + return replace( + template, + content=existing_content, + event_date=parse_event_date(clock=clock), + metadata=freeze_json({}), + entities=(), + observation_scopes=None, + document_id=document_id, + update_mode=UpdateMode.REPLACE, + source_index=None, + origin=ContentOrigin.EXISTING_DOCUMENT, + ) + + +def prepend_existing_document( + intent: DocumentIntent, + existing_content: str, + *, + clock: Clock = _utcnow, +) -> DocumentIntent: + """Prepend existing text to an append intent without adding a result slot.""" + + if intent.update_mode is not UpdateMode.APPEND: + raise ValueError("Existing document text can only be prepended to an append intent") + if not intent.items: + raise ValueError("An append intent requires at least one content item") + + synthetic = make_append_synthetic_item( + existing_content, + document_id=intent.document_id, + template=intent.items[0], + clock=clock, + ) + return replace(intent, items=(synthetic, *intent.items)) diff --git a/core/dataplane/hms_api/engine/ingestion/domain.py b/core/dataplane/hms_api/engine/ingestion/domain.py new file mode 100644 index 0000000..b4601b8 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/domain.py @@ -0,0 +1,158 @@ +"""Immutable domain values shared by Retain planning stages.""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime +from enum import StrEnum +from typing import Any, TypeAlias + +FrozenJsonScalar: TypeAlias = str | int | float | bool | None + + +@dataclass(frozen=True, slots=True) +class FrozenObject: + items: tuple[tuple[str, "FrozenJson"], ...] + + +@dataclass(frozen=True, slots=True) +class FrozenArray: + items: tuple["FrozenJson", ...] + + +FrozenJson: TypeAlias = FrozenJsonScalar | FrozenObject | FrozenArray + + +def freeze_json(value: Any) -> FrozenJson: + """Create a recursively immutable, order-preserving JSON-like value.""" + + if value is None or isinstance(value, (str, int, float, bool)): + return value + if isinstance(value, dict): + return FrozenObject(tuple((str(key), freeze_json(item)) for key, item in value.items())) + if isinstance(value, (list, tuple)): + return FrozenArray(tuple(freeze_json(item) for item in value)) + raise TypeError(f"Expected a JSON-compatible value, got {type(value).__name__}") + + +def thaw_json(value: FrozenJson) -> Any: + """Convert a value produced by :func:`freeze_json` back to containers.""" + + if isinstance(value, FrozenObject): + return {key: thaw_json(item) for key, item in value.items} + if isinstance(value, FrozenArray): + return [thaw_json(item) for item in value.items] + return value + + +class EventDateState(StrEnum): + """Meaning of an item's event-date input after compatibility parsing.""" + + DEFAULTED = "defaulted" + TIMELESS = "timeless" + EXPLICIT = "explicit" + + +@dataclass(frozen=True, slots=True) +class EventDateValue: + state: EventDateState + value: datetime | None + + def __post_init__(self) -> None: + if self.state is EventDateState.TIMELESS and self.value is not None: + raise ValueError("A timeless event date cannot carry a datetime") + if self.state is not EventDateState.TIMELESS and self.value is None: + raise ValueError(f"{self.state.value} event date requires a datetime") + + +class ContentOrigin(StrEnum): + SUBMITTED = "submitted" + EXISTING_DOCUMENT = "existing_document" + + +class UpdateMode(StrEnum): + REPLACE = "replace" + APPEND = "append" + + +ObservationScopes: TypeAlias = str | tuple[tuple[str, ...], ...] | None + + +@dataclass(frozen=True, slots=True) +class ContentItem: + """Normalized, immutable Retain input item.""" + + content: str + context: str + event_date: EventDateValue + metadata: FrozenJson + entities: tuple[FrozenJson, ...] + tags: tuple[str, ...] + observation_scopes: ObservationScopes + document_id: str | None + update_mode: UpdateMode + source_index: int | None + origin: ContentOrigin = ContentOrigin.SUBMITTED + + +@dataclass(frozen=True, slots=True) +class DocumentIntent: + """All normalized content that will form one tracked document.""" + + document_id: str + items: tuple[ContentItem, ...] + source_indices: tuple[int, ...] + expected_input_slots: tuple[int, ...] + update_mode: UpdateMode + + +@dataclass(frozen=True, slots=True) +class ChunkPolicy: + version: str + max_chars: int + conversation_mode: bool = True + overlap: int = 0 + + def __post_init__(self) -> None: + if self.max_chars <= 0: + raise ValueError("max_chars must be greater than zero") + if self.overlap < 0 or self.overlap >= self.max_chars: + raise ValueError("overlap must be non-negative and smaller than max_chars") + + +@dataclass(frozen=True, slots=True) +class ChunkPlan: + chunk_key: str + source_index: int | None + global_index: int + local_index: int + text: str + content_hash: str + + +@dataclass(frozen=True, slots=True) +class ExistingChunkFingerprint: + chunk_id: str + chunk_index: int + content_hash: str | None + + +class DocumentChangeKind(StrEnum): + FULL = "full" + DELTA = "delta" + METADATA_ONLY = "metadata_only" + STALE_SKIP = "stale_skip" + + +@dataclass(frozen=True, slots=True) +class DocumentChangePlan: + kind: DocumentChangeKind + unchanged: tuple[int, ...] = () + changed: tuple[int, ...] = () + added: tuple[int, ...] = () + removed: tuple[int, ...] = () + reason: str | None = None + + @property + def chunks_to_process(self) -> tuple[int, ...]: + return tuple(sorted((*self.changed, *self.added))) diff --git a/core/dataplane/hms_api/engine/ingestion/execution/__init__.py b/core/dataplane/hms_api/engine/ingestion/execution/__init__.py new file mode 100644 index 0000000..069f25c --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/execution/__init__.py @@ -0,0 +1,17 @@ +"""Pure execution planning primitives for Retain.""" + +from .windowing import ( + FactRecordIdentity, + FullWriteWindowPlan, + WindowUnitResult, + merge_window_unit_ids, + plan_full_write_windows, +) + +__all__ = [ + "FactRecordIdentity", + "FullWriteWindowPlan", + "WindowUnitResult", + "merge_window_unit_ids", + "plan_full_write_windows", +] diff --git a/core/dataplane/hms_api/engine/ingestion/execution/windowing.py b/core/dataplane/hms_api/engine/ingestion/execution/windowing.py new file mode 100644 index 0000000..bb3214d --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/execution/windowing.py @@ -0,0 +1,200 @@ +"""Pure planning and result-mapping helpers for FULL Retain write windows.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass +from typing import Protocol, TypeAlias + +from ..domain import ChunkPlan + + +class FactRecordIdentity(Protocol): + """Minimum stable identity required to map one committed fact result.""" + + fact_key: str + source_index: int | None + + +WindowUnitResult: TypeAlias = tuple[ + Sequence[FactRecordIdentity], + Sequence[tuple[str, str]], +] + + +@dataclass(frozen=True, slots=True) +class FullWriteWindowPlan: + """One immutable, globally ordered FULL-document write window.""" + + window_index: int + chunks: tuple[ChunkPlan, ...] + global_indices: tuple[int, ...] + is_first: bool + is_last: bool + + def __post_init__(self) -> None: + if isinstance(self.window_index, bool) or not isinstance(self.window_index, int): + raise TypeError("window_index must be an integer") + if self.window_index < 0: + raise ValueError("window_index must be non-negative") + if not isinstance(self.chunks, tuple) or any(not isinstance(chunk, ChunkPlan) for chunk in self.chunks): + raise TypeError("chunks must be a tuple of ChunkPlan values") + if not isinstance(self.global_indices, tuple) or any( + isinstance(index, bool) or not isinstance(index, int) for index in self.global_indices + ): + raise TypeError("global_indices must be a tuple of integers") + expected_indices = tuple(chunk.global_index for chunk in self.chunks) + if self.global_indices != expected_indices: + raise ValueError("global_indices must exactly match chunks in their stored order") + if any(left >= right for left, right in zip(self.global_indices, self.global_indices[1:])): + raise ValueError("global_indices must be strictly increasing") + if not isinstance(self.is_first, bool) or not isinstance(self.is_last, bool): + raise TypeError("is_first and is_last must be booleans") + if self.is_first != (self.window_index == 0): + raise ValueError("is_first must be true exactly for window_index=0") + if not self.chunks and not (self.window_index == 0 and self.is_first and self.is_last): + raise ValueError("an empty FULL window must be the sole first/final window") + + +def plan_full_write_windows( + chunks: Sequence[ChunkPlan], + batch_size: int, +) -> tuple[FullWriteWindowPlan, ...]: + """Partition a complete, ordered chunk plan into deterministic windows. + + The configured ``retain_chunk_batch_size=0`` convention disables batching, + so all chunks are placed in one window. Empty documents still receive one + first/final zero-fact window so document tracking and a final transactional + callback have a durable execution point. + """ + + if isinstance(batch_size, bool) or not isinstance(batch_size, int): + raise TypeError("batch_size must be an integer") + if batch_size < 0: + raise ValueError("batch_size must be non-negative; 0 disables batching") + + chunk_batch = tuple(chunks) + if any(not isinstance(chunk, ChunkPlan) for chunk in chunk_batch): + raise TypeError("chunks must contain only ChunkPlan values") + observed_indices = tuple(chunk.global_index for chunk in chunk_batch) + expected_indices = tuple(range(len(chunk_batch))) + if observed_indices != expected_indices: + raise ValueError("chunks must be ordered by unique, continuous global_index values 0..N-1") + + if not chunk_batch: + return ( + FullWriteWindowPlan( + window_index=0, + chunks=(), + global_indices=(), + is_first=True, + is_last=True, + ), + ) + + effective_batch_size = len(chunk_batch) if batch_size == 0 else batch_size + windows: list[FullWriteWindowPlan] = [] + for window_index, start in enumerate(range(0, len(chunk_batch), effective_batch_size)): + window_chunks = chunk_batch[start : start + effective_batch_size] + windows.append( + FullWriteWindowPlan( + window_index=window_index, + chunks=window_chunks, + global_indices=tuple(chunk.global_index for chunk in window_chunks), + is_first=window_index == 0, + is_last=start + effective_batch_size >= len(chunk_batch), + ) + ) + return tuple(windows) + + +def merge_window_unit_ids( + document_source_indices: Sequence[int | None], + window_results: Sequence[WindowUnitResult], +) -> tuple[tuple[str, ...], ...]: + """Map committed window results back to ordered submitted document items. + + ``document_source_indices`` is the document-item order, not a dense numeric + range. A single ``None`` denotes append's synthetic existing-content item; + its committed unit IDs are validated but deliberately omitted from the + public buckets. Mapping order never determines output order: records do. + """ + + sources = tuple(document_source_indices) + seen_sources: set[int | None] = set() + for source in sources: + _validate_source_index(source, field_name="document_source_indices") + if source in seen_sources: + raise ValueError(f"document_source_indices contains duplicate source_index={source!r}") + seen_sources.add(source) + + public_sources = tuple(source for source in sources if source is not None) + buckets: dict[int, list[str]] = {source: [] for source in public_sources} + seen_record_keys: set[str] = set() + seen_mapping_keys: set[str] = set() + seen_unit_ids: set[str] = set() + + for window_index, window_result in enumerate(window_results): + if not isinstance(window_result, tuple) or len(window_result) != 2: + raise TypeError(f"window_results[{window_index}] must be a (records, unit_ids_by_fact_key) tuple") + records = tuple(window_result[0]) + bindings = tuple(window_result[1]) + + local_record_keys: list[str] = [] + for record_index, record in enumerate(records): + fact_key = getattr(record, "fact_key", None) + source_index = getattr(record, "source_index", object()) + if not isinstance(fact_key, str) or not fact_key: + raise ValueError(f"window_results[{window_index}].records[{record_index}] has an invalid fact_key") + _validate_source_index( + source_index, + field_name=f"window_results[{window_index}].records[{record_index}].source_index", + ) + if source_index not in seen_sources: + raise ValueError(f"fact_key={fact_key!r} references unknown source_index={source_index!r}") + if fact_key in seen_record_keys: + raise ValueError(f"duplicate record fact_key across windows: {fact_key!r}") + seen_record_keys.add(fact_key) + local_record_keys.append(fact_key) + + units_by_key: dict[str, str] = {} + for binding_index, binding in enumerate(bindings): + if not isinstance(binding, tuple) or len(binding) != 2: + raise TypeError( + f"window_results[{window_index}].unit_ids_by_fact_key[{binding_index}] " + "must be a (fact_key, unit_id) tuple" + ) + fact_key, unit_id = binding + if not isinstance(fact_key, str) or not fact_key: + raise ValueError("unit_ids_by_fact_key contains an invalid fact_key") + if not isinstance(unit_id, str) or not unit_id: + raise ValueError(f"unit_ids_by_fact_key[{fact_key!r}] has an invalid unit_id") + if fact_key in seen_mapping_keys: + raise ValueError(f"duplicate mapped fact_key across windows: {fact_key!r}") + if unit_id in seen_unit_ids: + raise ValueError(f"duplicate unit_id across windows: {unit_id!r}") + seen_mapping_keys.add(fact_key) + seen_unit_ids.add(unit_id) + units_by_key[fact_key] = unit_id + + record_key_set = set(local_record_keys) + mapping_key_set = set(units_by_key) + if record_key_set != mapping_key_set: + missing = sorted(record_key_set - mapping_key_set) + unexpected = sorted(mapping_key_set - record_key_set) + raise ValueError(f"window fact-key mapping is incomplete (missing={missing}, unexpected={unexpected})") + + for record in records: + if record.source_index is not None: + buckets[record.source_index].append(units_by_key[record.fact_key]) + + return tuple(tuple(buckets[source]) for source in public_sources) + + +def _validate_source_index(value: object, *, field_name: str) -> None: + if value is None: + return + if isinstance(value, bool) or not isinstance(value, int): + raise TypeError(f"{field_name} must contain only integers or None") + if value < 0: + raise ValueError(f"{field_name} must contain only non-negative integers or None") diff --git a/core/dataplane/hms_api/engine/ingestion/extraction/__init__.py b/core/dataplane/hms_api/engine/ingestion/extraction/__init__.py new file mode 100644 index 0000000..bec5378 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/extraction/__init__.py @@ -0,0 +1,54 @@ +"""Provider-neutral extraction contracts and strategies.""" + +from .extractor import FactExtractorAdapter +from .layout import ( + PrechunkedExtractionLayout, + PrechunkedLayoutError, + build_prechunked_extraction_layout, +) +from .models import FACT_KEY_VERSION, CausalFactRelation, FactCandidate, compute_fact_key +from .passthrough import ( + SECONDS_PER_FACT, + build_content_position_map, + extract_passthrough, + to_chunk_metadata, + to_chunk_metadata_batch, +) +from .ports import ( + BatchExtractionUnsupportedError, + ChunkFactCount, + ExtractionAdapterError, + ExtractionContractError, + ExtractionMode, + ExtractionModeMismatchError, + ExtractionPolicy, + ExtractionRequest, + ExtractionResult, + FactExtractor, +) + +__all__ = [ + "FACT_KEY_VERSION", + "SECONDS_PER_FACT", + "BatchExtractionUnsupportedError", + "CausalFactRelation", + "ChunkFactCount", + "ExtractionAdapterError", + "ExtractionContractError", + "ExtractionMode", + "ExtractionModeMismatchError", + "ExtractionPolicy", + "ExtractionRequest", + "ExtractionResult", + "FactCandidate", + "FactExtractor", + "FactExtractorAdapter", + "PrechunkedExtractionLayout", + "PrechunkedLayoutError", + "build_content_position_map", + "build_prechunked_extraction_layout", + "compute_fact_key", + "extract_passthrough", + "to_chunk_metadata", + "to_chunk_metadata_batch", +] diff --git a/core/dataplane/hms_api/engine/ingestion/extraction/extractor.py b/core/dataplane/hms_api/engine/ingestion/extraction/extractor.py new file mode 100644 index 0000000..f091d7f --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/extraction/extractor.py @@ -0,0 +1,495 @@ +"""Fact extractor adapter for the ingestion pipeline.""" + +from __future__ import annotations + +from collections import Counter +from collections.abc import Awaitable, Callable, Mapping, Sequence +from copy import copy, deepcopy +from dataclasses import replace +from datetime import UTC, datetime +from typing import Any + +from ...response_models import TokenUsage +from ..domain import ChunkPlan, ContentItem, FrozenJson, freeze_json, thaw_json +from .models import CausalFactRelation, FactCandidate, compute_fact_key +from .passthrough import build_content_position_map, extract_passthrough +from .ports import ( + BatchExtractionUnsupportedError, + ChunkFactCount, + ExtractionContractError, + ExtractionMode, + ExtractionModeMismatchError, + ExtractionRequest, + ExtractionResult, +) + +ExtractionPrimitive = Callable[..., Awaitable[tuple[Any, Any, Any]]] +BatchCheckpointClearer = Callable[[], Awaitable[None]] + + +def _frozen_metadata(value: Any, *, fact_index: int) -> FrozenJson: + if not isinstance(value, Mapping): + raise ExtractionContractError(f"fact[{fact_index}].metadata must be a mapping") + try: + return freeze_json(deepcopy(dict(value))) + except TypeError as exc: + raise ExtractionContractError(f"fact[{fact_index}].metadata is not JSON-compatible") from exc + + +def _string_tuple(value: Any, *, field_name: str) -> tuple[str, ...]: + if value is None: + return () + if isinstance(value, (str, bytes)) or not isinstance(value, Sequence): + raise ExtractionContractError(f"{field_name} must be a sequence of strings") + result = tuple(value) + if any(not isinstance(item, str) for item in result): + raise ExtractionContractError(f"{field_name} must contain only strings") + return result + + +def _optional_string(value: Any, *, field_name: str) -> str | None: + if value is not None and not isinstance(value, str): + raise ExtractionContractError(f"{field_name} must be a string or None") + return value + + +def _aware_datetime(value: Any, *, field_name: str) -> datetime | None: + """Interpret naive datetimes as UTC while preserving aware time zones.""" + + if value is None: + return None + if not isinstance(value, datetime): + raise ExtractionContractError(f"{field_name} must be a datetime or None") + if value.tzinfo is None or value.utcoffset() is None: + return value.replace(tzinfo=UTC) + return value + + +def _observation_scopes(value: Any, *, field_name: str): + if value is None: + return None + if isinstance(value, str): + if value not in {"per_tag", "combined", "all_combinations"}: + raise ExtractionContractError(f"{field_name} contains an unsupported named scope") + return value + if isinstance(value, bytes) or not isinstance(value, Sequence): + raise ExtractionContractError(f"{field_name} must be a named scope or nested string sequences") + + scopes: list[tuple[str, ...]] = [] + for scope_index, scope in enumerate(value): + if isinstance(scope, (str, bytes)) or not isinstance(scope, Sequence): + raise ExtractionContractError(f"{field_name}[{scope_index}] must be a sequence of strings") + frozen_scope = tuple(scope) + if any(not isinstance(tag, str) for tag in frozen_scope): + raise ExtractionContractError(f"{field_name}[{scope_index}] must contain only strings") + scopes.append(frozen_scope) + return tuple(scopes) + + +def _storage_content(item: ContentItem): + metadata = thaw_json(item.metadata) + if not isinstance(metadata, dict): + raise ExtractionContractError("ContentItem.metadata must thaw to an object") + + entities = [] + for index, frozen_entity in enumerate(item.entities): + entity = thaw_json(frozen_entity) + if not isinstance(entity, dict): + raise ExtractionContractError(f"ContentItem.entities[{index}] must thaw to an object") + entities.append(entity) + + if isinstance(item.observation_scopes, tuple): + observation_scopes = [list(scope) for scope in item.observation_scopes] + else: + observation_scopes = item.observation_scopes + + from ...retain.types import RetainContent + + return RetainContent( + content=item.content, + context=item.context, + event_date=item.event_date.value, + metadata=metadata, + entities=entities, + tags=list(item.tags), + observation_scopes=observation_scopes, + ) + + +def _validate_planned_chunks( + chunks: tuple[ChunkPlan, ...], + items: tuple[ContentItem, ...], +) -> dict[int | None, int]: + try: + positions = build_content_position_map(items) + except (TypeError, ValueError) as exc: + raise ExtractionContractError(str(exc)) from exc + + if bool(items) != bool(chunks): + raise ExtractionContractError("items and planned chunks must either both be empty or both be non-empty") + + seen_keys: set[str] = set() + seen_indices: set[int] = set() + planned_sources: set[int | None] = set() + for position, chunk in enumerate(chunks): + if chunk.chunk_key in seen_keys: + raise ExtractionContractError(f"planned chunk[{position}] duplicates chunk_key={chunk.chunk_key!r}") + if chunk.global_index in seen_indices: + raise ExtractionContractError(f"planned chunk[{position}] duplicates global_index={chunk.global_index}") + if chunk.source_index not in positions: + raise ExtractionContractError( + f"planned chunk[{position}] references unknown source_index={chunk.source_index!r}" + ) + seen_keys.add(chunk.chunk_key) + seen_indices.add(chunk.global_index) + planned_sources.add(chunk.source_index) + + missing_sources = set(positions) - planned_sources + if missing_sources: + raise ExtractionContractError(f"No planned chunk for source indices: {sorted(missing_sources, key=str)!r}") + return positions + + +def _validate_extracted_chunks( + extracted_chunks: Sequence[Any], + planned_chunks: tuple[ChunkPlan, ...], + content_positions: Mapping[int | None, int], +) -> tuple[int, ...]: + if len(extracted_chunks) != len(planned_chunks): + raise ExtractionContractError( + f"Extractor returned {len(extracted_chunks)} chunks for {len(planned_chunks)} planned chunks" + ) + + declared_counts: list[int] = [] + for position, (extracted_chunk, planned_chunk) in enumerate(zip(extracted_chunks, planned_chunks, strict=True)): + chunk_index = getattr(extracted_chunk, "chunk_index", None) + content_index = getattr(extracted_chunk, "content_index", None) + fact_count = getattr(extracted_chunk, "fact_count", None) + chunk_text = getattr(extracted_chunk, "chunk_text", None) + + if isinstance(chunk_index, bool) or not isinstance(chunk_index, int) or chunk_index != position: + raise ExtractionContractError( + f"chunk[{position}].chunk_index must equal its zero-based position; got {chunk_index!r}" + ) + expected_content_index = content_positions[planned_chunk.source_index] + if ( + isinstance(content_index, bool) + or not isinstance(content_index, int) + or content_index != expected_content_index + ): + raise ExtractionContractError( + f"chunk[{position}].content_index={content_index!r} does not match " + f"source_index={planned_chunk.source_index!r} at content position {expected_content_index}" + ) + if chunk_text != planned_chunk.text: + raise ExtractionContractError(f"chunk[{position}] text does not match its planned chunk") + if isinstance(fact_count, bool) or not isinstance(fact_count, int) or fact_count < 0: + raise ExtractionContractError(f"chunk[{position}].fact_count must be a non-negative integer") + declared_counts.append(fact_count) + return tuple(declared_counts) + + +class FactExtractorAdapter: + """Expose the configured extractor through the ingestion extraction port.""" + + def __init__( + self, + *, + llm_config: Any, + config: Any, + agent_name: str, + pool: Any = None, + operation_id: str | None = None, + schema: str | None = None, + batch_checkpoint_clearer: BatchCheckpointClearer | None = None, + sync_primitive: ExtractionPrimitive | None = None, + batch_primitive: ExtractionPrimitive | None = None, + ) -> None: + if sync_primitive is None or batch_primitive is None: + from ...retain import fact_extraction + + sync_primitive = sync_primitive or fact_extraction.extract_facts_from_contents + batch_primitive = batch_primitive or fact_extraction.extract_facts_from_contents_batch_api + self._llm_config = llm_config + self._config = config + self._agent_name = agent_name + self._pool = pool + self._operation_id = operation_id + self._schema = schema + self._batch_checkpoint_clearer = batch_checkpoint_clearer + self._sync_primitive = sync_primitive + self._batch_primitive = batch_primitive + + def _sync_fallback_config(self) -> Any: + """Disable Batch API on a shallow config copy for one safe fallback. + + The batch entry point falls back by calling the sync entry point + with ``retain_batch_enabled`` still true, which immediately routes back + to Batch API and recurses forever. This adapter makes the one-shot + boundary explicit and never mutates the resolved bank configuration + shared by the request. + """ + + fallback = copy(self._config) + setattr(fallback, "retain_batch_enabled", False) + return fallback + + async def _batch_primitive_if_supported( + self, + mode: ExtractionMode, + ) -> tuple[ExtractionPrimitive, Any]: + if mode is ExtractionMode.VERBATIM: + # The Batch result parser requires ``what`` while verbatim + # deliberately omits it. Running the established sync primitive + # is lossless and avoids silently producing zero facts. + return self._sync_primitive, self._sync_fallback_config() + + provider = getattr(self._llm_config, "_provider_impl", None) + supports_batch_api = getattr(provider, "supports_batch_api", None) + if not callable(supports_batch_api): + return self._sync_primitive, self._sync_fallback_config() + try: + supported = await supports_batch_api() + except Exception as exc: + raise BatchExtractionUnsupportedError("Failed to determine provider Batch API capability") from exc + if supported is not True: + return self._sync_primitive, self._sync_fallback_config() + return self._batch_primitive, self._config + + async def extract(self, request: ExtractionRequest) -> ExtractionResult: + configured_mode = getattr(self._config, "retain_extraction_mode", None) + if configured_mode != request.policy.mode.value: + raise ExtractionModeMismatchError( + f"Extraction policy mode={request.policy.mode.value!r} does not match " + f"resolved config mode={configured_mode!r}" + ) + + content_positions = _validate_planned_chunks(request.chunks, request.items) + if request.policy.mode is ExtractionMode.CHUNKS: + candidates = extract_passthrough( + request.chunks, + request.items, + fact_type_override=request.policy.fact_type_override, + ) + return ExtractionResult( + candidates=candidates, + chunk_fact_counts=tuple(ChunkFactCount(chunk.chunk_key, 1) for chunk in request.chunks), + usage=TokenUsage(), + ) + + primitive = self._sync_primitive + primitive_config = self._config + if getattr(self._config, "retain_batch_enabled", False): + primitive, primitive_config = await self._batch_primitive_if_supported(request.policy.mode) + + storage_contents = [_storage_content(item) for item in request.items] + primitive_result = await primitive( + contents=storage_contents, + llm_config=self._llm_config, + agent_name=self._agent_name, + config=primitive_config, + pool=self._pool, + operation_id=self._operation_id, + schema=self._schema, + ) + if not isinstance(primitive_result, tuple) or len(primitive_result) != 3: + raise ExtractionContractError("Extractor primitive must return a three-tuple of facts, chunks, and usage") + + extracted_facts, extracted_chunks, usage = primitive_result + if isinstance(extracted_facts, (str, bytes)) or not isinstance(extracted_facts, Sequence): + raise ExtractionContractError("Extractor facts must be a sequence") + if isinstance(extracted_chunks, (str, bytes)) or not isinstance(extracted_chunks, Sequence): + raise ExtractionContractError("Extractor chunks must be a sequence") + if not isinstance(usage, TokenUsage): + raise ExtractionContractError("Extractor usage must be a TokenUsage") + + declared_counts = _validate_extracted_chunks(extracted_chunks, request.chunks, content_positions) + if request.policy.mode is ExtractionMode.VERBATIM: + # Sync verbatim collapses multiple metadata candidates to one fact + # per non-empty chunk but leaves the pre-collapse ChunkMetadata + # count unchanged. Normalize that known compatibility artifact at + # this boundary; all other modes retain strict count equality. + declared_counts = tuple(1 if count else 0 for count in declared_counts) + candidates, actual_counts = self._convert_facts( + tuple(extracted_facts), + tuple(extracted_chunks), + request, + content_positions, + ) + if actual_counts != declared_counts: + raise ExtractionContractError( + f"Chunk fact-count mismatch: metadata={declared_counts!r}, actual={actual_counts!r}" + ) + + causal_relations = self._convert_causal_relations(tuple(extracted_facts), candidates) + relations_by_source: dict[str, list[CausalFactRelation]] = {} + for relation in causal_relations: + relations_by_source.setdefault(relation.source_fact_key, []).append(relation) + candidates = tuple( + replace(candidate, causal_relations=tuple(relations_by_source.get(candidate.fact_key, ()))) + for candidate in candidates + ) + result = ExtractionResult( + candidates=candidates, + chunk_fact_counts=tuple( + ChunkFactCount(chunk.chunk_key, actual_counts[position]) + for position, chunk in enumerate(request.chunks) + ), + usage=usage, + causal_relations=causal_relations, + ) + if primitive is self._batch_primitive and self._batch_checkpoint_clearer is not None: + # A single async operation may contain multiple documents/windows. + # Retire this completed provider job before the next extraction so + # it cannot accidentally resume results belonging to another + # chunk set. Crashes while polling still retain the checkpoint. + await self._batch_checkpoint_clearer() + return result + + def _convert_facts( + self, + extracted_facts: tuple[Any, ...], + extracted_chunks: tuple[Any, ...], + request: ExtractionRequest, + content_positions: Mapping[int | None, int], + ) -> tuple[tuple[FactCandidate, ...], tuple[int, ...]]: + actual_counts = [0] * len(request.chunks) + per_chunk_ordinals: Counter[int] = Counter() + candidates: list[FactCandidate] = [] + seen_fact_keys: set[str] = set() + + for fact_index, fact in enumerate(extracted_facts): + chunk_index = getattr(fact, "chunk_index", None) + content_index = getattr(fact, "content_index", None) + if ( + isinstance(chunk_index, bool) + or not isinstance(chunk_index, int) + or not 0 <= chunk_index < len(request.chunks) + ): + raise ExtractionContractError( + f"fact[{fact_index}].chunk_index={chunk_index!r} is outside the returned chunk range" + ) + + planned_chunk = request.chunks[chunk_index] + expected_content_index = content_positions[planned_chunk.source_index] + if ( + isinstance(content_index, bool) + or not isinstance(content_index, int) + or content_index != expected_content_index + ): + raise ExtractionContractError( + f"fact[{fact_index}].content_index={content_index!r} does not match its chunk content position " + f"{expected_content_index}" + ) + if getattr(extracted_chunks[chunk_index], "content_index", None) != content_index: + raise ExtractionContractError(f"fact[{fact_index}] and chunk metadata disagree on content_index") + + text = getattr(fact, "fact_text", None) + raw_fact_type = getattr(fact, "fact_type", None) + context = getattr(fact, "context", None) + if not isinstance(text, str): + raise ExtractionContractError(f"fact[{fact_index}].fact_text must be a string") + if not isinstance(raw_fact_type, str) or not raw_fact_type: + raise ExtractionContractError(f"fact[{fact_index}].fact_type must be a non-empty string") + if not isinstance(context, str): + raise ExtractionContractError(f"fact[{fact_index}].context must be a string") + + fact_type = request.policy.fact_type_override or raw_fact_type + extractor_local_index = per_chunk_ordinals[chunk_index] + per_chunk_ordinals[chunk_index] += 1 + fact_key = compute_fact_key( + chunk_key=planned_chunk.chunk_key, + source_index=planned_chunk.source_index, + global_index=planned_chunk.global_index, + extractor_local_index=extractor_local_index, + text=text, + fact_type=fact_type, + ) + if fact_key in seen_fact_keys: + raise ExtractionContractError(f"fact[{fact_index}] produced a duplicate stable fact key") + seen_fact_keys.add(fact_key) + + item = request.items[content_index] + try: + candidate = FactCandidate( + fact_key=fact_key, + chunk_key=planned_chunk.chunk_key, + source_index=planned_chunk.source_index, + global_index=planned_chunk.global_index, + extractor_local_index=extractor_local_index, + text=text, + fact_type=fact_type, + context=context, + where=_optional_string( + getattr(fact, "where", None), + field_name=f"fact[{fact_index}].where", + ), + occurred_start=_aware_datetime( + getattr(fact, "occurred_start", None), + field_name=f"fact[{fact_index}].occurred_start", + ), + occurred_end=_aware_datetime( + getattr(fact, "occurred_end", None), + field_name=f"fact[{fact_index}].occurred_end", + ), + mentioned_at=_aware_datetime( + getattr(fact, "mentioned_at", None), + field_name=f"fact[{fact_index}].mentioned_at", + ), + metadata=_frozen_metadata(getattr(fact, "metadata", None), fact_index=fact_index), + declared_entities=item.entities, + tags=_string_tuple(getattr(fact, "tags", None), field_name=f"fact[{fact_index}].tags"), + observation_scopes=_observation_scopes( + getattr(fact, "observation_scopes", None), + field_name=f"fact[{fact_index}].observation_scopes", + ), + entity_mentions=_string_tuple( + getattr(fact, "entities", None), + field_name=f"fact[{fact_index}].entities", + ), + causal_relations=(), + ) + except ExtractionContractError: + raise + except (TypeError, ValueError) as exc: + raise ExtractionContractError(f"fact[{fact_index}] violates FactCandidate invariants: {exc}") from exc + + candidates.append(candidate) + actual_counts[chunk_index] += 1 + + return tuple(candidates), tuple(actual_counts) + + @staticmethod + def _convert_causal_relations( + extracted_facts: tuple[Any, ...], + candidates: tuple[FactCandidate, ...], + ) -> tuple[CausalFactRelation, ...]: + relations: list[CausalFactRelation] = [] + for source_index, fact in enumerate(extracted_facts): + raw_relations = getattr(fact, "causal_relations", None) or () + if isinstance(raw_relations, (str, bytes)) or not isinstance(raw_relations, Sequence): + raise ExtractionContractError(f"fact[{source_index}].causal_relations must be a sequence") + for relation_index, relation in enumerate(raw_relations): + target_index = getattr(relation, "target_fact_index", None) + relation_type = getattr(relation, "relation_type", None) + if ( + isinstance(target_index, bool) + or not isinstance(target_index, int) + or not 0 <= target_index < source_index + ): + raise ExtractionContractError( + f"fact[{source_index}].causal_relations[{relation_index}] target index " + f"must reference an earlier fact; got {target_index!r}" + ) + if relation_type != "caused_by": + raise ExtractionContractError( + f"fact[{source_index}].causal_relations[{relation_index}] has unsupported relation_type" + ) + relations.append( + CausalFactRelation( + source_fact_key=candidates[source_index].fact_key, + target_fact_key=candidates[target_index].fact_key, + relation_type=relation_type, + ) + ) + return tuple(relations) diff --git a/core/dataplane/hms_api/engine/ingestion/extraction/layout.py b/core/dataplane/hms_api/engine/ingestion/extraction/layout.py new file mode 100644 index 0000000..7e4e713 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/extraction/layout.py @@ -0,0 +1,324 @@ +"""Lossless pre-chunked layout for windowed Retain extraction. + +The extraction primitive identifies input content by its zero-based position, +while a document can have several independently active chunks that all belong +to the same stable ``ContentItem.source_index``. This module gives every active +chunk a temporary one-to-one content position and then restores document-level +identities on the extraction result. +""" + +from __future__ import annotations + +from collections import Counter +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, replace +from types import MappingProxyType + +from ..domain import ChunkPlan, ContentItem +from .models import CausalFactRelation, FactCandidate, compute_fact_key +from .passthrough import build_content_position_map +from .ports import ( + ChunkFactCount, + ExtractionContractError, + ExtractionPolicy, + ExtractionRequest, + ExtractionResult, +) + + +class PrechunkedLayoutError(ExtractionContractError): + """A pre-chunked request or its extraction result is inconsistent.""" + + +@dataclass(frozen=True, slots=True) +class PrechunkedExtractionLayout: + """Temporary one-content-per-chunk extraction layout. + + ``temporary_items`` and ``temporary_chunks`` are suitable for an + :class:`ExtractionRequest`. Temporary source indices are exactly their + local content positions. ``temporary_to_original_source`` is an immutable + map back to the stable source indices in ``document_items``. + """ + + document_items: tuple[ContentItem, ...] + active_chunks: tuple[ChunkPlan, ...] + temporary_items: tuple[ContentItem, ...] + temporary_chunks: tuple[ChunkPlan, ...] + temporary_to_original_source: Mapping[int, int | None] + + def __post_init__(self) -> None: + for field_name, values, value_type in ( + ("document_items", self.document_items, ContentItem), + ("active_chunks", self.active_chunks, ChunkPlan), + ("temporary_items", self.temporary_items, ContentItem), + ("temporary_chunks", self.temporary_chunks, ChunkPlan), + ): + if not isinstance(values, tuple) or any(not isinstance(value, value_type) for value in values): + raise TypeError(f"{field_name} must be a tuple of {value_type.__name__} values") + + try: + original_positions = build_content_position_map(self.document_items) + except (TypeError, ValueError) as exc: + raise PrechunkedLayoutError(str(exc)) from exc + + cardinality = len(self.active_chunks) + if len(self.temporary_items) != cardinality or len(self.temporary_chunks) != cardinality: + raise PrechunkedLayoutError( + "active chunks, temporary items, and temporary chunks must have equal cardinality" + ) + + if not isinstance(self.temporary_to_original_source, Mapping): + raise TypeError("temporary_to_original_source must be a mapping") + source_map = dict(self.temporary_to_original_source) + expected_temporary_sources = set(range(cardinality)) + if set(source_map) != expected_temporary_sources: + raise PrechunkedLayoutError("temporary source mapping keys must exactly equal local content positions") + for temporary_source, original_source in source_map.items(): + if isinstance(temporary_source, bool) or not isinstance(temporary_source, int): + raise PrechunkedLayoutError("temporary source mapping keys must be integers") + if original_source is not None and ( + isinstance(original_source, bool) or not isinstance(original_source, int) or original_source < 0 + ): + raise PrechunkedLayoutError("original source mapping values must be non-negative integers or None") + object.__setattr__(self, "temporary_to_original_source", MappingProxyType(source_map)) + + seen_chunk_keys: set[str] = set() + seen_global_indices: set[int] = set() + for position, (original_chunk, temporary_item, temporary_chunk) in enumerate( + zip(self.active_chunks, self.temporary_items, self.temporary_chunks, strict=True) + ): + if original_chunk.chunk_key in seen_chunk_keys: + raise PrechunkedLayoutError(f"duplicate active chunk_key={original_chunk.chunk_key!r}") + if original_chunk.global_index in seen_global_indices: + raise PrechunkedLayoutError(f"duplicate active chunk global_index={original_chunk.global_index}") + seen_chunk_keys.add(original_chunk.chunk_key) + seen_global_indices.add(original_chunk.global_index) + + if original_chunk.source_index not in original_positions: + raise PrechunkedLayoutError( + f"active chunk[{position}] references unknown source_index={original_chunk.source_index!r}" + ) + if source_map[position] != original_chunk.source_index: + raise PrechunkedLayoutError(f"temporary source mapping disagrees at content position {position}") + + original_item = self.document_items[original_positions[original_chunk.source_index]] + expected_item = replace(original_item, content=original_chunk.text, source_index=position) + expected_chunk = replace(original_chunk, source_index=position) + if temporary_item != expected_item: + raise PrechunkedLayoutError( + f"temporary item[{position}] does not losslessly represent its active chunk" + ) + if temporary_chunk != expected_chunk: + raise PrechunkedLayoutError(f"temporary chunk[{position}] changed fields other than source_index") + + def extraction_request(self, policy: ExtractionPolicy) -> ExtractionRequest: + """Build the extractor request for this temporary layout.""" + + return ExtractionRequest(items=self.temporary_items, chunks=self.temporary_chunks, policy=policy) + + def remap_result(self, result: ExtractionResult) -> ExtractionResult: + """Restore original source indices and every derived fact identity.""" + + if not isinstance(result, ExtractionResult): + raise TypeError("result must be an ExtractionResult") + + self._validate_chunk_counts(result.chunk_fact_counts, result.candidates) + + old_to_new_fact_key: dict[str, str] = {} + remapped_without_relations: list[FactCandidate] = [] + seen_old_keys: set[str] = set() + seen_new_keys: set[str] = set() + seen_local_indices: set[tuple[int, int]] = set() + + for candidate_position, candidate in enumerate(result.candidates): + temporary_source = candidate.source_index + if isinstance(temporary_source, bool) or not isinstance(temporary_source, int): + raise PrechunkedLayoutError( + f"candidate[{candidate_position}] has invalid temporary source_index={temporary_source!r}" + ) + if temporary_source not in self.temporary_to_original_source: + raise PrechunkedLayoutError( + f"candidate[{candidate_position}] references unknown temporary source_index={temporary_source}" + ) + + expected_chunk = self.temporary_chunks[temporary_source] + if candidate.chunk_key != expected_chunk.chunk_key: + raise PrechunkedLayoutError( + f"candidate[{candidate_position}] chunk_key does not match its temporary source" + ) + if candidate.global_index != expected_chunk.global_index: + raise PrechunkedLayoutError( + f"candidate[{candidate_position}] global_index does not match its temporary source" + ) + + local_identity = (temporary_source, candidate.extractor_local_index) + if local_identity in seen_local_indices: + raise PrechunkedLayoutError( + f"candidate[{candidate_position}] duplicates extractor-local index " + f"{candidate.extractor_local_index} for temporary source {temporary_source}" + ) + seen_local_indices.add(local_identity) + + expected_old_key = compute_fact_key( + chunk_key=candidate.chunk_key, + source_index=temporary_source, + global_index=candidate.global_index, + extractor_local_index=candidate.extractor_local_index, + text=candidate.text, + fact_type=candidate.fact_type, + ) + if candidate.fact_key != expected_old_key: + raise PrechunkedLayoutError( + f"candidate[{candidate_position}] fact_key is inconsistent with its temporary identity" + ) + if candidate.fact_key in seen_old_keys: + raise PrechunkedLayoutError(f"candidate[{candidate_position}] duplicates fact_key") + seen_old_keys.add(candidate.fact_key) + + original_source = self.temporary_to_original_source[temporary_source] + new_fact_key = compute_fact_key( + chunk_key=candidate.chunk_key, + source_index=original_source, + global_index=candidate.global_index, + extractor_local_index=candidate.extractor_local_index, + text=candidate.text, + fact_type=candidate.fact_type, + ) + if new_fact_key in seen_new_keys: + raise PrechunkedLayoutError( + f"candidate[{candidate_position}] collides after restoring its original source" + ) + seen_new_keys.add(new_fact_key) + old_to_new_fact_key[candidate.fact_key] = new_fact_key + remapped_without_relations.append( + replace( + candidate, + fact_key=new_fact_key, + source_index=original_source, + causal_relations=(), + ) + ) + + remapped_relations = tuple( + self._remap_relation(relation, old_to_new_fact_key, relation_position) + for relation_position, relation in enumerate(result.causal_relations) + ) + relations_by_source: dict[str, list[CausalFactRelation]] = {} + for relation in remapped_relations: + relations_by_source.setdefault(relation.source_fact_key, []).append(relation) + + candidates = tuple( + replace( + candidate, + causal_relations=tuple(relations_by_source.get(candidate.fact_key, ())), + ) + for candidate in remapped_without_relations + ) + if tuple(relation for candidate in candidates for relation in candidate.causal_relations) != remapped_relations: + raise PrechunkedLayoutError("causal relation ordering cannot be represented by the remapped candidates") + + # Counts and usage contain no source-derived identity, so retain their + # exact values (and object identity) across the layout boundary. + return ExtractionResult( + candidates=candidates, + chunk_fact_counts=result.chunk_fact_counts, + usage=result.usage, + causal_relations=remapped_relations, + ) + + def _validate_chunk_counts( + self, + counts: tuple[ChunkFactCount, ...], + candidates: tuple[FactCandidate, ...], + ) -> None: + if len(counts) != len(self.temporary_chunks): + raise PrechunkedLayoutError("extraction result must contain exactly one fact count for every active chunk") + + actual_counts = Counter(candidate.chunk_key for candidate in candidates) + for position, (count, chunk) in enumerate(zip(counts, self.temporary_chunks, strict=True)): + if count.chunk_key != chunk.chunk_key: + raise PrechunkedLayoutError(f"chunk_fact_counts[{position}] does not match the active chunk order") + if actual_counts[count.chunk_key] != count.fact_count: + raise PrechunkedLayoutError( + f"chunk_fact_counts[{position}]={count.fact_count} does not match " + f"{actual_counts[count.chunk_key]} candidates" + ) + + known_chunk_keys = {chunk.chunk_key for chunk in self.temporary_chunks} + unknown_chunk_keys = set(actual_counts) - known_chunk_keys + if unknown_chunk_keys: + raise PrechunkedLayoutError(f"candidates reference unknown chunk keys: {sorted(unknown_chunk_keys)!r}") + + @staticmethod + def _remap_relation( + relation: CausalFactRelation, + fact_key_map: Mapping[str, str], + relation_position: int, + ) -> CausalFactRelation: + try: + source_fact_key = fact_key_map[relation.source_fact_key] + except KeyError as exc: + raise PrechunkedLayoutError( + f"causal relation[{relation_position}] source has no candidate mapping" + ) from exc + try: + target_fact_key = fact_key_map[relation.target_fact_key] + except KeyError as exc: + raise PrechunkedLayoutError( + f"causal relation[{relation_position}] target has no candidate mapping" + ) from exc + try: + return CausalFactRelation( + source_fact_key=source_fact_key, + target_fact_key=target_fact_key, + relation_type=relation.relation_type, + ) + except (TypeError, ValueError) as exc: + raise PrechunkedLayoutError( + f"causal relation[{relation_position}] violates remapped relation invariants: {exc}" + ) from exc + + +def build_prechunked_extraction_layout( + document_items: Sequence[ContentItem], + active_chunks: Sequence[ChunkPlan], +) -> PrechunkedExtractionLayout: + """Build an immutable one-to-one extraction layout for active chunks.""" + + if isinstance(document_items, (str, bytes)) or not isinstance(document_items, Sequence): + raise TypeError("document_items must be a sequence of ContentItem values") + if isinstance(active_chunks, (str, bytes)) or not isinstance(active_chunks, Sequence): + raise TypeError("active_chunks must be a sequence of ChunkPlan values") + items = tuple(document_items) + chunks = tuple(active_chunks) + if any(not isinstance(item, ContentItem) for item in items): + raise TypeError("document_items must contain only ContentItem values") + if any(not isinstance(chunk, ChunkPlan) for chunk in chunks): + raise TypeError("active_chunks must contain only ChunkPlan values") + + try: + original_positions = build_content_position_map(items) + except (TypeError, ValueError) as exc: + raise PrechunkedLayoutError(str(exc)) from exc + + temporary_items: list[ContentItem] = [] + temporary_chunks: list[ChunkPlan] = [] + source_map: dict[int, int | None] = {} + for temporary_source, chunk in enumerate(chunks): + try: + original_item = items[original_positions[chunk.source_index]] + except KeyError as exc: + raise PrechunkedLayoutError( + f"active chunk[{temporary_source}] references unknown source_index={chunk.source_index!r}" + ) from exc + temporary_items.append(replace(original_item, content=chunk.text, source_index=temporary_source)) + temporary_chunks.append(replace(chunk, source_index=temporary_source)) + source_map[temporary_source] = chunk.source_index + + return PrechunkedExtractionLayout( + document_items=items, + active_chunks=chunks, + temporary_items=tuple(temporary_items), + temporary_chunks=tuple(temporary_chunks), + temporary_to_original_source=source_map, + ) diff --git a/core/dataplane/hms_api/engine/ingestion/extraction/models.py b/core/dataplane/hms_api/engine/ingestion/extraction/models.py new file mode 100644 index 0000000..90297af --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/extraction/models.py @@ -0,0 +1,145 @@ +"""Provider-neutral fact candidates produced by Retain extraction strategies.""" + +from __future__ import annotations + +import hashlib +import json +from dataclasses import dataclass +from datetime import datetime + +from ..domain import FrozenJson, ObservationScopes + +FACT_KEY_VERSION = "retain-fact-v1" + + +def compute_fact_key( + *, + chunk_key: str, + source_index: int | None, + global_index: int, + extractor_local_index: int, + text: str, + fact_type: str, +) -> str: + """Build an unambiguous deterministic SHA-256 fact identity. + + Canonical JSON avoids delimiter collisions, while the explicit version + leaves room for a future identity policy without silently changing keys. + """ + + if isinstance(extractor_local_index, bool) or not isinstance(extractor_local_index, int): + raise TypeError("extractor_local_index must be an integer") + if extractor_local_index < 0: + raise ValueError("extractor_local_index must be non-negative") + + payload = json.dumps( + [FACT_KEY_VERSION, chunk_key, source_index, global_index, extractor_local_index, text, fact_type], + ensure_ascii=False, + separators=(",", ":"), + ) + return hashlib.sha256(payload.encode("utf-8")).hexdigest() + + +def _validate_datetime(value: datetime | None, *, field_name: str) -> None: + if value is None: + return + if not isinstance(value, datetime): + raise TypeError(f"{field_name} must be a datetime or None") + if value.tzinfo is None or value.utcoffset() is None: + raise ValueError(f"{field_name} must be timezone-aware") + + +@dataclass(frozen=True, slots=True) +class CausalFactRelation: + """Stable provider-neutral causal edge between two extracted facts.""" + + source_fact_key: str + target_fact_key: str + relation_type: str + + def __post_init__(self) -> None: + for field_name, value in ( + ("source_fact_key", self.source_fact_key), + ("target_fact_key", self.target_fact_key), + ): + if ( + not isinstance(value, str) + or len(value) != 64 + or any(character not in "0123456789abcdef" for character in value) + ): + raise ValueError(f"{field_name} must be a lowercase SHA-256 hex digest") + if self.source_fact_key == self.target_fact_key: + raise ValueError("causal relations cannot be self-referential") + if self.relation_type != "caused_by": + raise ValueError("relation_type must be 'caused_by'") + + +@dataclass(frozen=True, slots=True) +class FactCandidate: + """Immutable extraction result identified independently of database IDs.""" + + fact_key: str + chunk_key: str + source_index: int | None + global_index: int + extractor_local_index: int + text: str + fact_type: str + context: str + where: str | None + occurred_start: datetime | None + occurred_end: datetime | None + mentioned_at: datetime | None + metadata: FrozenJson + declared_entities: tuple[FrozenJson, ...] + tags: tuple[str, ...] + observation_scopes: ObservationScopes + # Provider-extracted names are distinct from caller-declared entity + # objects. Passthrough extraction deliberately leaves these empty, which + # preserves the existing projection manifest semantics. + entity_mentions: tuple[str, ...] + causal_relations: tuple[CausalFactRelation, ...] + + def __post_init__(self) -> None: + if len(self.fact_key) != 64 or any(character not in "0123456789abcdef" for character in self.fact_key): + raise ValueError("fact_key must be a lowercase SHA-256 hex digest") + if not isinstance(self.chunk_key, str) or not self.chunk_key: + raise ValueError("chunk_key must be a non-empty string") + if self.source_index is not None: + if isinstance(self.source_index, bool) or not isinstance(self.source_index, int): + raise TypeError("source_index must be an integer or None") + if self.source_index < 0: + raise ValueError("source_index must be non-negative") + if isinstance(self.global_index, bool) or not isinstance(self.global_index, int): + raise TypeError("global_index must be an integer") + if self.global_index < 0: + raise ValueError("global_index must be non-negative") + if isinstance(self.extractor_local_index, bool) or not isinstance(self.extractor_local_index, int): + raise TypeError("extractor_local_index must be an integer") + if self.extractor_local_index < 0: + raise ValueError("extractor_local_index must be non-negative") + if not isinstance(self.text, str): + raise TypeError("text must be a string") + if not isinstance(self.fact_type, str) or not self.fact_type: + raise ValueError("fact_type must be a non-empty string") + if not isinstance(self.context, str): + raise TypeError("context must be a string") + if self.where is not None and not isinstance(self.where, str): + raise TypeError("where must be a string or None") + _validate_datetime(self.occurred_start, field_name="occurred_start") + _validate_datetime(self.occurred_end, field_name="occurred_end") + _validate_datetime(self.mentioned_at, field_name="mentioned_at") + if not isinstance(self.declared_entities, tuple): + raise TypeError("declared_entities must be a tuple of frozen JSON values") + if not isinstance(self.tags, tuple) or any(not isinstance(tag, str) for tag in self.tags): + raise TypeError("tags must be a tuple of strings") + if not isinstance(self.entity_mentions, tuple) or any( + not isinstance(entity, str) for entity in self.entity_mentions + ): + raise TypeError("entity_mentions must be a tuple of strings") + if not isinstance(self.causal_relations, tuple) or any( + not isinstance(relation, CausalFactRelation) for relation in self.causal_relations + ): + raise TypeError("causal_relations must be a tuple of CausalFactRelation values") + if any(relation.source_fact_key != self.fact_key for relation in self.causal_relations): + raise ValueError("every causal relation on a candidate must use that candidate as its source") diff --git a/core/dataplane/hms_api/engine/ingestion/extraction/passthrough.py b/core/dataplane/hms_api/engine/ingestion/extraction/passthrough.py new file mode 100644 index 0000000..09ca19f --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/extraction/passthrough.py @@ -0,0 +1,151 @@ +"""LLM-free extraction in which every planned chunk is one fact.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from datetime import timedelta + +from ..domain import ChunkPlan, ContentItem, ContentOrigin +from .models import FactCandidate, compute_fact_key + +# Chunk passthrough offsets facts by absolute extraction order so equal +# timestamps still have a deterministic temporal ordering. +SECONDS_PER_FACT = 0.01 + + +def build_content_position_map(items: Sequence[ContentItem]) -> dict[int | None, int]: + """Map stable source indices to positional content indices. + + At most one ``None`` key is permitted, and it must represent the synthetic + existing-document item used by append planning. + """ + + positions: dict[int | None, int] = {} + for position, item in enumerate(items): + source_index = item.source_index + if source_index is None and item.origin is not ContentOrigin.EXISTING_DOCUMENT: + raise ValueError("source_index=None is reserved for one synthetic existing-document item") + if source_index in positions: + label = "synthetic source_index=None" if source_index is None else f"source_index={source_index}" + raise ValueError(f"Duplicate {label} in content items") + positions[source_index] = position + return positions + + +def _fact_type(override: str | None) -> str: + if override is None or override == "": + return "world" + if not isinstance(override, str): + raise TypeError("fact_type_override must be a string or None") + return override + + +def extract_passthrough( + chunks: Sequence[ChunkPlan], + items: Sequence[ContentItem], + *, + fact_type_override: str | None = None, + fact_position_offset: int = 0, +) -> tuple[FactCandidate, ...]: + """Produce exactly one immutable candidate for every chunk plan.""" + + if isinstance(fact_position_offset, bool) or not isinstance(fact_position_offset, int): + raise TypeError("fact_position_offset must be an integer") + if fact_position_offset < 0: + raise ValueError("fact_position_offset must be non-negative") + + content_positions = build_content_position_map(items) + fact_type = _fact_type(fact_type_override) + seen_chunk_keys: set[str] = set() + seen_global_indices: set[int] = set() + candidates: list[FactCandidate] = [] + + for fact_position, chunk in enumerate(chunks): + if chunk.chunk_key in seen_chunk_keys: + raise ValueError(f"Duplicate chunk_key: {chunk.chunk_key!r}") + if chunk.global_index in seen_global_indices: + raise ValueError(f"Duplicate chunk global_index: {chunk.global_index}") + seen_chunk_keys.add(chunk.chunk_key) + seen_global_indices.add(chunk.global_index) + + try: + item = items[content_positions[chunk.source_index]] + except KeyError as exc: + raise ValueError( + f"Chunk {chunk.chunk_key!r} references unknown source_index={chunk.source_index!r}" + ) from exc + + mentioned_at = item.event_date.value + absolute_fact_position = fact_position_offset + fact_position + if mentioned_at is not None and absolute_fact_position: + mentioned_at += timedelta(seconds=absolute_fact_position * SECONDS_PER_FACT) + + candidates.append( + FactCandidate( + fact_key=compute_fact_key( + chunk_key=chunk.chunk_key, + source_index=chunk.source_index, + global_index=chunk.global_index, + extractor_local_index=0, + text=chunk.text, + fact_type=fact_type, + ), + chunk_key=chunk.chunk_key, + source_index=chunk.source_index, + global_index=chunk.global_index, + extractor_local_index=0, + text=chunk.text, + fact_type=fact_type, + context=item.context, + where=None, + occurred_start=None, + occurred_end=None, + mentioned_at=mentioned_at, + metadata=item.metadata, + declared_entities=item.entities, + tags=item.tags, + observation_scopes=item.observation_scopes, + entity_mentions=(), + causal_relations=(), + ) + ) + + return tuple(candidates) + + +def to_chunk_metadata(chunk: ChunkPlan, *, content_index: int): + """Adapt a planned one-fact chunk to persistence metadata. + + ``content_index`` is explicit because source indices are stable request + identities, whereas the storage field is a position in the current + extraction window. + """ + + if isinstance(content_index, bool) or not isinstance(content_index, int) or content_index < 0: + raise ValueError("content_index must be a non-negative integer") + + from ...retain.types import ChunkMetadata + + return ChunkMetadata( + chunk_text=chunk.text, + fact_count=1, + content_index=content_index, + chunk_index=chunk.global_index, + ) + + +def to_chunk_metadata_batch( + chunks: Sequence[ChunkPlan], + *, + content_positions: Mapping[int | None, int], +): + """Adapt chunks using an explicit stable-source-to-position mapping.""" + + metadata = [] + for chunk in chunks: + try: + content_index = content_positions[chunk.source_index] + except KeyError as exc: + raise ValueError(f"Missing content position for source_index={chunk.source_index!r}") from exc + metadata.append(to_chunk_metadata(chunk, content_index=content_index)) + return metadata diff --git a/core/dataplane/hms_api/engine/ingestion/extraction/ports.py b/core/dataplane/hms_api/engine/ingestion/extraction/ports.py new file mode 100644 index 0000000..2b47e1d --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/extraction/ports.py @@ -0,0 +1,116 @@ +"""Provider-neutral contracts for Retain fact extraction.""" + +from __future__ import annotations + +from dataclasses import dataclass +from enum import StrEnum +from typing import Protocol, runtime_checkable + +from ...response_models import TokenUsage +from ..domain import ChunkPlan, ContentItem +from .models import CausalFactRelation, FactCandidate + + +class ExtractionMode(StrEnum): + CONCISE = "concise" + VERBOSE = "verbose" + CUSTOM = "custom" + VERBATIM = "verbatim" + CHUNKS = "chunks" + + +@dataclass(frozen=True, slots=True) +class ExtractionPolicy: + mode: ExtractionMode + fact_type_override: str | None = None + + def __post_init__(self) -> None: + if not isinstance(self.mode, ExtractionMode): + raise TypeError("mode must be an ExtractionMode") + if self.fact_type_override is not None and not isinstance(self.fact_type_override, str): + raise TypeError("fact_type_override must be a string or None") + + +@dataclass(frozen=True, slots=True) +class ExtractionRequest: + items: tuple[ContentItem, ...] + chunks: tuple[ChunkPlan, ...] + policy: ExtractionPolicy + + def __post_init__(self) -> None: + if not isinstance(self.items, tuple) or any(not isinstance(item, ContentItem) for item in self.items): + raise TypeError("items must be a tuple of ContentItem values") + if not isinstance(self.chunks, tuple) or any(not isinstance(chunk, ChunkPlan) for chunk in self.chunks): + raise TypeError("chunks must be a tuple of ChunkPlan values") + if not isinstance(self.policy, ExtractionPolicy): + raise TypeError("policy must be an ExtractionPolicy") + + +@dataclass(frozen=True, slots=True) +class ChunkFactCount: + chunk_key: str + fact_count: int + + def __post_init__(self) -> None: + if not isinstance(self.chunk_key, str) or not self.chunk_key: + raise ValueError("chunk_key must be a non-empty string") + if isinstance(self.fact_count, bool) or not isinstance(self.fact_count, int) or self.fact_count < 0: + raise ValueError("fact_count must be a non-negative integer") + + +@dataclass(frozen=True, slots=True) +class ExtractionResult: + candidates: tuple[FactCandidate, ...] + chunk_fact_counts: tuple[ChunkFactCount, ...] + usage: TokenUsage + causal_relations: tuple[CausalFactRelation, ...] = () + + def __post_init__(self) -> None: + if not isinstance(self.candidates, tuple) or any( + not isinstance(candidate, FactCandidate) for candidate in self.candidates + ): + raise TypeError("candidates must be a tuple of FactCandidate values") + if not isinstance(self.chunk_fact_counts, tuple) or any( + not isinstance(count, ChunkFactCount) for count in self.chunk_fact_counts + ): + raise TypeError("chunk_fact_counts must be a tuple of ChunkFactCount values") + if sum(count.fact_count for count in self.chunk_fact_counts) != len(self.candidates): + raise ValueError("chunk fact counts must equal the number of candidates") + if not isinstance(self.usage, TokenUsage): + raise TypeError("usage must be a TokenUsage") + if not isinstance(self.causal_relations, tuple) or any( + not isinstance(relation, CausalFactRelation) for relation in self.causal_relations + ): + raise TypeError("causal_relations must be a tuple of CausalFactRelation values") + candidate_keys = {candidate.fact_key for candidate in self.candidates} + if any( + relation.source_fact_key not in candidate_keys or relation.target_fact_key not in candidate_keys + for relation in self.causal_relations + ): + raise ValueError("causal relation endpoints must reference candidates in this extraction result") + candidate_relations = tuple( + relation for candidate in self.candidates for relation in candidate.causal_relations + ) + if candidate_relations != self.causal_relations: + raise ValueError("result causal relations must exactly match the relations attached to candidates") + + +@runtime_checkable +class FactExtractor(Protocol): + async def extract(self, request: ExtractionRequest) -> ExtractionResult: ... + + +class ExtractionAdapterError(RuntimeError): + """Base error for an extractor boundary failure.""" + + +class ExtractionContractError(ExtractionAdapterError): + """The extraction backend returned an internally inconsistent result.""" + + +class ExtractionModeMismatchError(ExtractionAdapterError): + """The provider-neutral policy and resolved configuration disagree.""" + + +class BatchExtractionUnsupportedError(ExtractionAdapterError): + """Batch extraction was requested but cannot be safely adapted.""" diff --git a/core/dataplane/hms_api/engine/ingestion/normalization.py b/core/dataplane/hms_api/engine/ingestion/normalization.py new file mode 100644 index 0000000..e0a5b27 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/normalization.py @@ -0,0 +1,276 @@ +"""Pure normalization of caller-owned Retain inputs. + +This module is deliberately independent from the database, provider clients, +and process configuration. It converts the raw compatibility envelope into +immutable domain values without retaining references to caller-owned +containers. +""" + +from __future__ import annotations + +from collections.abc import Callable, Iterable, Mapping, Sequence +from copy import deepcopy +from datetime import UTC, datetime +from typing import Any, Final + +from .domain import ( + ContentItem, + EventDateState, + EventDateValue, + FrozenJson, + ObservationScopes, + UpdateMode, + freeze_json, +) + +Clock = Callable[[], datetime] + +_MISSING: Final = object() +_OBSERVATION_SCOPE_NAMES: Final = frozenset({"per_tag", "combined", "all_combinations"}) + + +def _utcnow() -> datetime: + return datetime.now(UTC) + + +def _as_aware(value: datetime, *, field_name: str) -> datetime: + """Return an aware datetime while preserving an explicit timezone. + + Retain interprets naive values as UTC but passes aware values through + unchanged. Keeping that rule matters before extraction because + converting to UTC can change the calendar date and weekday near a timezone + boundary. + """ + + if not isinstance(value, datetime): + raise TypeError(f"{field_name} must produce a datetime, got {type(value).__name__}") + if value.tzinfo is None or value.utcoffset() is None: + return value.replace(tzinfo=UTC) + return value + + +def parse_event_date( + value: Any = _MISSING, + *, + clock: Clock = _utcnow, +) -> EventDateValue: + """Normalize Retain event-date semantics into an explicit state. + + A missing value, or any falsey value other than ``None``, retains the + established behavior of defaulting to the injected clock. An explicit ``None`` + is timeless. Truthy values must be a :class:`datetime` or an ISO-8601 + string. Naive values are interpreted as UTC; explicit timezones are kept. + """ + + if value is None: + return EventDateValue(EventDateState.TIMELESS, None) + + if value is _MISSING or not value: + return EventDateValue(EventDateState.DEFAULTED, _as_aware(clock(), field_name="clock")) + + if isinstance(value, datetime): + parsed = value + elif isinstance(value, str): + try: + parsed = datetime.fromisoformat(value.replace("Z", "+00:00")) + except ValueError as exc: + raise ValueError(f"event_date must be a valid ISO-8601 datetime, got {value!r}") from exc + else: + raise TypeError(f"event_date must be a datetime or ISO-8601 string, got {type(value).__name__}") + + return EventDateValue(EventDateState.EXPLICIT, _as_aware(parsed, field_name="event_date")) + + +def _normalize_tags(value: Any, *, field_name: str) -> tuple[str, ...]: + if value is None: + return () + if isinstance(value, (str, bytes)) or not isinstance(value, Sequence): + raise TypeError(f"{field_name} must be a sequence of strings or None") + + normalized: list[str] = [] + for index, tag in enumerate(value): + if not isinstance(tag, str): + raise TypeError(f"{field_name}[{index}] must be a string, got {type(tag).__name__}") + normalized.append(tag) + return tuple(normalized) + + +def merge_tags(item_tags: Any = None, document_tags: Any = None) -> tuple[str, ...]: + """Merge item and batch tags with stable, first-occurrence deduplication.""" + + candidates = ( + *_normalize_tags(item_tags, field_name="tags"), + *_normalize_tags(document_tags, field_name="document_tags"), + ) + seen: set[str] = set() + merged: list[str] = [] + for tag in candidates: + if tag not in seen: + seen.add(tag) + merged.append(tag) + return tuple(merged) + + +def _freeze_mapping(value: Any, *, field_name: str) -> FrozenJson: + if value is None: + value = {} + if not isinstance(value, Mapping): + raise TypeError(f"{field_name} must be a mapping or None, got {type(value).__name__}") + + # ``dict`` snapshots arbitrary Mapping implementations; ``deepcopy`` + # severs nested references before the JSON-like tree is frozen. + snapshot = deepcopy(dict(value)) + try: + return freeze_json(snapshot) + except TypeError as exc: + raise TypeError(f"{field_name} must contain only JSON-compatible values: {exc}") from exc + + +def _normalize_entities(value: Any) -> tuple[FrozenJson, ...]: + if value is None: + return () + if isinstance(value, (str, bytes)) or not isinstance(value, Sequence): + raise TypeError("entities must be a sequence of mappings or None") + + normalized: list[FrozenJson] = [] + for index, entity in enumerate(value): + if not isinstance(entity, Mapping): + raise TypeError(f"entities[{index}] must be a mapping, got {type(entity).__name__}") + + snapshot = deepcopy(dict(entity)) + text = snapshot.get("text", _MISSING) + if text is _MISSING: + raise ValueError(f"entities[{index}] must contain a 'text' field") + if not isinstance(text, str): + raise TypeError(f"entities[{index}].text must be a string, got {type(text).__name__}") + entity_type = snapshot.get("type") + if entity_type is not None and not isinstance(entity_type, str): + raise TypeError(f"entities[{index}].type must be a string or None, got {type(entity_type).__name__}") + + try: + normalized.append(freeze_json(snapshot)) + except TypeError as exc: + raise TypeError(f"entities[{index}] must contain only JSON-compatible values: {exc}") from exc + return tuple(normalized) + + +def _normalize_observation_scopes(value: Any) -> ObservationScopes: + if value is None: + return None + if isinstance(value, str): + if value not in _OBSERVATION_SCOPE_NAMES: + choices = ", ".join(sorted(_OBSERVATION_SCOPE_NAMES)) + raise ValueError(f"observation_scopes must be one of {choices}, or a sequence of tag sequences") + return value + if isinstance(value, bytes) or not isinstance(value, Sequence): + raise TypeError("observation_scopes must be a supported string, a sequence of tag sequences, or None") + + scopes: list[tuple[str, ...]] = [] + for scope_index, scope in enumerate(value): + if isinstance(scope, (str, bytes)) or not isinstance(scope, Sequence): + raise TypeError(f"observation_scopes[{scope_index}] must be a sequence of strings") + tags: list[str] = [] + for tag_index, tag in enumerate(scope): + if not isinstance(tag, str): + raise TypeError( + f"observation_scopes[{scope_index}][{tag_index}] must be a string, got {type(tag).__name__}" + ) + tags.append(tag) + scopes.append(tuple(tags)) + return tuple(scopes) + + +def _normalize_document_id(value: Any) -> str | None: + if value is None or value == "": + return None + if not isinstance(value, str): + raise TypeError(f"document_id must be a string or None, got {type(value).__name__}") + return value + + +def _normalize_update_mode(value: Any) -> UpdateMode: + if value is None: + return UpdateMode.REPLACE + if not isinstance(value, str): + raise TypeError(f"update_mode must be a string or None, got {type(value).__name__}") + try: + return UpdateMode(value) + except ValueError as exc: + choices = ", ".join(mode.value for mode in UpdateMode) + raise ValueError(f"update_mode must be one of: {choices}; got {value!r}") from exc + + +def _validate_source_index(source_index: Any) -> int: + if isinstance(source_index, bool) or not isinstance(source_index, int): + raise TypeError(f"source_index must be an integer, got {type(source_index).__name__}") + if source_index < 0: + raise ValueError("source_index must be non-negative for submitted content") + return source_index + + +def normalize_content_item( + raw_item: Mapping[str, Any], + *, + source_index: int = 0, + document_tags: Sequence[str] | None = None, + clock: Clock = _utcnow, +) -> ContentItem: + """Return one immutable submitted item without retaining caller state.""" + + if not isinstance(raw_item, Mapping): + raise TypeError(f"content item must be a mapping, got {type(raw_item).__name__}") + normalized_source_index = _validate_source_index(source_index) + item = deepcopy(dict(raw_item)) + + if "content" not in item: + raise ValueError("content is required") + content = item["content"] + if not isinstance(content, str): + raise TypeError(f"content must be a string, got {type(content).__name__}") + + context = item.get("context", "") + if context is None: + context = "" + if not isinstance(context, str): + raise TypeError(f"context must be a string or None, got {type(context).__name__}") + + document_id = _normalize_document_id(item.get("document_id")) + update_mode = _normalize_update_mode(item.get("update_mode")) + if update_mode is UpdateMode.APPEND and document_id is None: + raise ValueError("update_mode='append' requires a document_id") + + raw_event_date = item["event_date"] if "event_date" in item else _MISSING + return ContentItem( + content=content, + context=context, + event_date=parse_event_date(raw_event_date, clock=clock), + metadata=_freeze_mapping(item.get("metadata"), field_name="metadata"), + entities=_normalize_entities(item.get("entities")), + tags=merge_tags(item.get("tags"), document_tags), + observation_scopes=_normalize_observation_scopes(item.get("observation_scopes")), + document_id=document_id, + update_mode=update_mode, + source_index=normalized_source_index, + ) + + +def normalize_contents( + raw_contents: Iterable[Mapping[str, Any]], + *, + document_tags: Sequence[str] | None = None, + clock: Clock = _utcnow, +) -> tuple[ContentItem, ...]: + """Normalize a Retain batch and assign stable zero-based source indices.""" + + # Validate and snapshot batch tags once; the tuple is already immutable and + # can safely be shared while each item applies stable deduplication. + normalized_document_tags = _normalize_tags(document_tags, field_name="document_tags") + return tuple( + normalize_content_item( + raw_item, + source_index=source_index, + document_tags=normalized_document_tags, + clock=clock, + ) + for source_index, raw_item in enumerate(raw_contents) + ) diff --git a/core/dataplane/hms_api/engine/ingestion/persistence/__init__.py b/core/dataplane/hms_api/engine/ingestion/persistence/__init__.py new file mode 100644 index 0000000..d5c661a --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/persistence/__init__.py @@ -0,0 +1,14 @@ +"""Persistence ports and database adapters for Retain.""" + +from .models import CommittedUnitBinding, ExistingDocument +from .ports import CheckpointStore, PlanningRepository +from .postgres import PostgresCheckpointStore, PostgresPlanningRepository + +__all__ = [ + "CheckpointStore", + "CommittedUnitBinding", + "ExistingDocument", + "PlanningRepository", + "PostgresCheckpointStore", + "PostgresPlanningRepository", +] diff --git a/core/dataplane/hms_api/engine/ingestion/persistence/backend.py b/core/dataplane/hms_api/engine/ingestion/persistence/backend.py new file mode 100644 index 0000000..63468e9 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/persistence/backend.py @@ -0,0 +1,87 @@ +"""Database adapter selection for the Retain ingestion service.""" + +from __future__ import annotations + +from contextlib import asynccontextmanager +from dataclasses import dataclass +from typing import Any + +from ..adapters.postgres_fresh_ownership import FreshPostgresDocumentOwnership +from .operation_fence import OperationActivityFence +from .oracle import ( + FreshOracleDocumentOwnership, + OracleCheckpointStore, + OracleDocumentOwnership, + OraclePlanningRepository, +) +from .postgres import ( + PostgresCheckpointStore, + PostgresDocumentOwnership, + PostgresPlanningRepository, +) + + +@dataclass(frozen=True, slots=True) +class RetainBackendAdapters: + """Factories and transaction semantics for one supported database.""" + + backend_type: str + + def planning_repository(self, connection: Any, *, schema: str | None = None) -> Any: + if self.backend_type == "oracle": + return OraclePlanningRepository(connection, schema=schema) + return PostgresPlanningRepository(connection, schema=schema) + + def checkpoint_store(self, connection: Any, *, schema: str | None = None) -> Any: + if self.backend_type == "oracle": + return OracleCheckpointStore(connection, schema=schema) + return PostgresCheckpointStore(connection, schema=schema) + + def document_ownership(self, *, schema: str | None = None, fresh: bool = False) -> Any: + if self.backend_type == "oracle": + if fresh: + return FreshOracleDocumentOwnership(schema=schema) + return OracleDocumentOwnership(schema=schema) + if fresh: + return FreshPostgresDocumentOwnership(schema=schema) + return PostgresDocumentOwnership(schema=schema) + + def operation_activity_fence( + self, + operation_id: str | None, + *, + schema: str | None = None, + ) -> OperationActivityFence | None: + """Build a database-neutral fence for a tracked core write.""" + + if operation_id is None: + return None + return OperationActivityFence(operation_id, schema=schema) + + @asynccontextmanager + async def planning_snapshot(self, connection: Any): + """Open a backend-native read-only snapshot for Retain planning.""" + + if self.backend_type == "oracle": + # Oracle requires SET TRANSACTION to be the first transaction + # statement. OracleConnection.transaction() starts with SAVEPOINT, + # so the outer acquired connection owns this read-only transaction. + await connection.execute("SET TRANSACTION READ ONLY") + yield + return + + async with connection.transaction(): + await connection.execute("SET TRANSACTION ISOLATION LEVEL REPEATABLE READ READ ONLY") + yield + + +def retain_backend_adapters(backend_type: str) -> RetainBackendAdapters: + """Return adapters for an explicitly supported database backend.""" + + normalized = backend_type.strip().lower() if isinstance(backend_type, str) else "" + if normalized not in {"postgresql", "oracle"}: + raise ValueError(f"Unsupported Retain database backend: {backend_type!r}") + return RetainBackendAdapters(normalized) + + +__all__ = ["RetainBackendAdapters", "retain_backend_adapters"] diff --git a/core/dataplane/hms_api/engine/ingestion/persistence/models.py b/core/dataplane/hms_api/engine/ingestion/persistence/models.py new file mode 100644 index 0000000..a07bd0c --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/persistence/models.py @@ -0,0 +1,112 @@ +"""Database-neutral records returned by Retain persistence ports.""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime +from typing import Any + + +@dataclass(frozen=True, slots=True) +class CommittedUnitBinding: + """A committed memory unit and the source chunk position it belongs to. + + ``chunk_index`` is optional because durable memory units can have no chunk + association. Recovery must preserve those units instead of silently + dropping them from the public result. + """ + + unit_id: str + chunk_index: int | None + + def __post_init__(self) -> None: + if not isinstance(self.unit_id, str) or not self.unit_id: + raise ValueError("unit_id must be a non-empty string") + if self.chunk_index is None: + return + if isinstance(self.chunk_index, bool) or not isinstance(self.chunk_index, int): + raise TypeError("chunk_index must be an integer or None") + if self.chunk_index < 0: + raise ValueError("chunk_index must be non-negative") + + +@dataclass(frozen=True, slots=True) +class ExistingDocument: + document_id: str + bank_id: str + original_text: str + content_hash: str | None + retain_params: dict[str, Any] + tags: tuple[str, ...] + created_at: datetime | None + updated_at: datetime | None + + +@dataclass(frozen=True, slots=True) +class OperationCheckpoint: + """Durable async-operation state needed to resume Retain after a crash.""" + + document_ids: tuple[str, ...] = () + core_committed_document_ids: tuple[str, ...] = () + final_ann_pending_document_ids: tuple[str, ...] = () + committed_unit_ids_by_document: tuple[tuple[str, tuple[str, ...]], ...] = () + unscoped_facts_committed: bool = False + + def __post_init__(self) -> None: + for field_name, values in ( + ("document_ids", self.document_ids), + ("core_committed_document_ids", self.core_committed_document_ids), + ("final_ann_pending_document_ids", self.final_ann_pending_document_ids), + ): + if not isinstance(values, tuple) or any(not isinstance(value, str) or not value for value in values): + raise TypeError(f"{field_name} must be a tuple of non-empty strings") + if len(values) != len(set(values)): + raise ValueError(f"{field_name} must not contain duplicates") + + if not isinstance(self.committed_unit_ids_by_document, tuple): + raise TypeError("committed_unit_ids_by_document must be a tuple") + seen_document_ids: set[str] = set() + for binding in self.committed_unit_ids_by_document: + if not isinstance(binding, tuple) or len(binding) != 2: + raise TypeError("committed_unit_ids_by_document entries must be (document_id, unit_ids) tuples") + document_id, unit_ids = binding + if not isinstance(document_id, str) or not document_id: + raise TypeError("committed unit document IDs must be non-empty strings") + if document_id in seen_document_ids: + raise ValueError("committed_unit_ids_by_document must not contain duplicate document IDs") + seen_document_ids.add(document_id) + if not isinstance(unit_ids, tuple) or any( + not isinstance(unit_id, str) or not unit_id for unit_id in unit_ids + ): + raise TypeError("committed unit IDs must be tuples of non-empty strings") + if len(unit_ids) != len(set(unit_ids)): + raise ValueError(f"committed unit IDs for document {document_id!r} must not contain duplicates") + if not isinstance(self.unscoped_facts_committed, bool): + raise TypeError("unscoped_facts_committed must be a bool") + + def is_core_committed(self, document_id: str) -> bool: + """Apply the single-document fallback checkpoint rule.""" + + if document_id in self.core_committed_document_ids: + return True + return ( + self.unscoped_facts_committed + and not self.core_committed_document_ids + and (not self.document_ids or self.document_ids == (document_id,)) + ) + + def unit_ids_for_document(self, document_id: str) -> tuple[str, ...] | None: + """Return exact operation-local IDs, or ``None`` for an unscoped checkpoint. + + An empty tuple is a known value: that document committed successfully + but this operation produced no units. ``None`` means the versioned + mapping was absent or did not contain the document, so callers must use + a conservative document-level fallback. + """ + + if not isinstance(document_id, str) or not document_id: + raise ValueError("document_id must be a non-empty string") + for committed_document_id, unit_ids in self.committed_unit_ids_by_document: + if committed_document_id == document_id: + return unit_ids + return None diff --git a/core/dataplane/hms_api/engine/ingestion/persistence/operation_fence.py b/core/dataplane/hms_api/engine/ingestion/persistence/operation_fence.py new file mode 100644 index 0000000..f66a073 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/persistence/operation_fence.py @@ -0,0 +1,92 @@ +"""Transaction-level activity fence for tracked Retain writes.""" + +from __future__ import annotations + +import json +import uuid +from collections.abc import Mapping +from typing import Any + +from ...schema import fq_table_explicit +from ..contracts import RetainOperationInactiveError + +_ACTIVE_STATUSES = frozenset({"pending", "processing"}) + + +def _metadata_object(connection: Any, value: Any) -> dict[str, Any]: + if value is None: + return {} + parser = getattr(connection, "parse_json", None) + if callable(parser): + value = parser(value) + elif isinstance(value, str): + value = json.loads(value) + if not isinstance(value, Mapping): + raise RetainOperationInactiveError("Tracked Retain operation metadata is invalid") + return dict(value) + + +class OperationActivityFence: + """Serialize a core write against cancellation of its child and parent. + + The child row is locked first and the optional parent row second, matching + the cancellation and completion aggregation order. Holding both locks until + the core transaction exits gives cancellation a precise linearization + point: either cancellation commits first and this write is rejected, or + this write commits before cancellation can be accepted. + """ + + def __init__(self, operation_id: str, *, schema: str | None = None) -> None: + if not isinstance(operation_id, str) or not operation_id: + raise ValueError("operation_id must be a non-empty string") + try: + self._operation_id = uuid.UUID(operation_id) + except (AttributeError, ValueError) as exc: + raise ValueError("operation_id must be a UUID string") from exc + self._schema = schema + + async def assert_active(self, connection: Any, *, bank_id: str) -> None: + """Lock the operation chain and fail closed unless every row is active.""" + + if not isinstance(bank_id, str) or not bank_id: + raise ValueError("bank_id must be a non-empty string") + operations = fq_table_explicit("async_operations", self._schema) + child = await connection.fetchrow( + f""" + SELECT status, result_metadata + FROM {operations} + WHERE operation_id = $1 AND bank_id = $2 + FOR UPDATE + """, + self._operation_id, + bank_id, + ) + if child is None or child["status"] not in _ACTIVE_STATUSES: + raise RetainOperationInactiveError("Tracked Retain operation is no longer active") + + metadata = _metadata_object(connection, child["result_metadata"]) + parent_value = metadata.get("parent_operation_id") + if parent_value is None: + return + try: + parent_id = uuid.UUID(str(parent_value)) + except (AttributeError, ValueError) as exc: + raise RetainOperationInactiveError("Tracked Retain parent operation metadata is invalid") from exc + if parent_id == self._operation_id: + raise RetainOperationInactiveError("Tracked Retain operation cannot be its own parent") + + parent = await connection.fetchrow( + f""" + SELECT status + FROM {operations} + WHERE operation_id = $1 AND bank_id = $2 + FOR UPDATE + """, + parent_id, + bank_id, + ) + if parent is None or parent["status"] not in _ACTIVE_STATUSES: + raise RetainOperationInactiveError("Tracked Retain parent operation is no longer active") + + +__all__ = ["OperationActivityFence"] diff --git a/core/dataplane/hms_api/engine/ingestion/persistence/oracle.py b/core/dataplane/hms_api/engine/ingestion/persistence/oracle.py new file mode 100644 index 0000000..17369ca --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/persistence/oracle.py @@ -0,0 +1,480 @@ +"""Oracle persistence adapters for the Retain ingestion pipeline.""" + +from __future__ import annotations + +import json +import uuid +from collections.abc import Mapping +from typing import Any + +from ...schema import fq_table_explicit +from ..adapters.postgres_fresh_ownership import FreshDocumentOwnershipConflict +from ..domain import ExistingChunkFingerprint +from .models import CommittedUnitBinding, ExistingDocument, OperationCheckpoint +from .postgres import ( + _COMMITTED_UNIT_IDS_FIELD, + _COMMITTED_UNIT_IDS_VERSION, + _committed_unit_ids_by_document, + _json_object, + _require_identifier, + _string_array, + _tags, + _unit_id_tuple, +) + + +def _affected_rows(status: str) -> int: + """Return the row count from an asyncpg-compatible command status.""" + + if not isinstance(status, str): + return 0 + try: + return int(status.rsplit(maxsplit=1)[-1]) + except (IndexError, ValueError): + return 0 + + +class OracleDocumentOwnership: + """Oracle row-lock and document-hash ownership adapter.""" + + def __init__(self, *, schema: str | None = None) -> None: + self._schema = schema + + async def prepare_first_window(self, connection: Any, *, bank_id: str, document_id: str) -> None: + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + documents = fq_table_explicit("documents", self._schema) + await connection.execute( + f""" + INSERT INTO {documents} (id, bank_id, original_text, content_hash) + VALUES ($1, $2, '', '__pending__') + ON CONFLICT (id, bank_id) DO NOTHING + """, + document_id, + bank_id, + ) + await connection.fetchval( + f""" + SELECT content_hash + FROM {documents} + WHERE id = $1 AND bank_id = $2 + FOR UPDATE + """, + document_id, + bank_id, + ) + + async def validate_later_window( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + expected_content_hash: str, + ) -> bool: + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + _require_identifier(expected_content_hash, field_name="expected_content_hash") + current_hash = await connection.fetchval( + f""" + SELECT content_hash + FROM {fq_table_explicit("documents", self._schema)} + WHERE id = $1 AND bank_id = $2 + FOR UPDATE + """, + document_id, + bank_id, + ) + return current_hash is not None and current_hash == expected_content_hash + + async def validate_unhashed_window( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + ) -> bool: + """Lock an upgraded row and confirm that Oracle still sees no hash.""" + + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + row = await connection.fetchrow( + f""" + SELECT content_hash + FROM {fq_table_explicit("documents", self._schema)} + WHERE id = $1 AND bank_id = $2 + FOR UPDATE + """, + document_id, + bank_id, + ) + return row is not None and not row["content_hash"] + + async def transition_content_hash( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + expected_content_hash: str, + new_content_hash: str, + ) -> bool: + for field_name, value in ( + ("bank_id", bank_id), + ("document_id", document_id), + ("expected_content_hash", expected_content_hash), + ("new_content_hash", new_content_hash), + ): + _require_identifier(value, field_name=field_name) + status = await connection.execute( + f""" + UPDATE {fq_table_explicit("documents", self._schema)} + SET content_hash = $1, updated_at = now() + WHERE id = $2 AND bank_id = $3 AND content_hash = $4 + """, + new_content_hash, + document_id, + bank_id, + expected_content_hash, + ) + return _affected_rows(status) == 1 + + +class FreshOracleDocumentOwnership(OracleDocumentOwnership): + """Atomically claim a document that was absent during Oracle preflight.""" + + async def prepare_first_window(self, connection: Any, *, bank_id: str, document_id: str) -> None: + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + documents = fq_table_explicit("documents", self._schema) + status = await connection.execute( + f""" + INSERT INTO {documents} (id, bank_id, original_text, content_hash) + VALUES ($1, $2, '', '__pending__') + ON CONFLICT (id, bank_id) DO NOTHING + """, + document_id, + bank_id, + ) + if _affected_rows(status) != 1: + raise FreshDocumentOwnershipConflict("The document is no longer fresh") + + locked = await connection.fetchval( + f""" + SELECT id + FROM {documents} + WHERE id = $1 AND bank_id = $2 + FOR UPDATE + """, + document_id, + bank_id, + ) + if locked is None: # pragma: no cover - same transaction inserted it + raise RuntimeError("Fresh document ownership row disappeared inside its transaction") + + async def validate_later_window( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + expected_content_hash: str, + ) -> bool: + del connection, bank_id, document_id, expected_content_hash + raise RuntimeError("Fresh-document ownership does not support later full-write windows") + + async def validate_unhashed_window( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + ) -> bool: + del connection, bank_id, document_id + raise RuntimeError("Fresh-document ownership does not support existing unhashed rows") + + +class OraclePlanningRepository: + """Load the durable Oracle state required for Retain planning.""" + + def __init__(self, connection: Any, *, schema: str | None = None) -> None: + self._connection = connection + self._schema = schema + + async def load_document(self, bank_id: str, document_id: str) -> ExistingDocument | None: + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + row = await self._connection.fetchrow( + f""" + SELECT id, bank_id, original_text, content_hash, retain_params, + tags, created_at, updated_at + FROM {fq_table_explicit("documents", self._schema)} + WHERE id = $1 AND bank_id = $2 + """, + document_id, + bank_id, + ) + if row is None: + return None + return ExistingDocument( + document_id=str(row["id"]), + bank_id=str(row["bank_id"]), + original_text=row["original_text"] or "", + content_hash=row["content_hash"], + retain_params=_json_object(row["retain_params"]), + tags=_tags(row["tags"]), + created_at=row["created_at"], + updated_at=row["updated_at"], + ) + + async def load_chunks( + self, + bank_id: str, + document_id: str, + ) -> tuple[ExistingChunkFingerprint, ...]: + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + rows = await self._connection.fetch( + f""" + SELECT chunk_id, chunk_index, content_hash + FROM {fq_table_explicit("chunks", self._schema)} + WHERE document_id = $1 AND bank_id = $2 + ORDER BY chunk_index + """, + document_id, + bank_id, + ) + return tuple( + ExistingChunkFingerprint( + chunk_id=str(row["chunk_id"]), + chunk_index=int(row["chunk_index"]), + content_hash=row["content_hash"], + ) + for row in rows + ) + + async def load_document_unit_ids(self, bank_id: str, document_id: str) -> tuple[str, ...]: + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + rows = await self._connection.fetch( + f""" + SELECT id + FROM {fq_table_explicit("memory_units", self._schema)} + WHERE bank_id = $1 AND document_id = $2 + ORDER BY created_at, id + """, + bank_id, + document_id, + ) + unit_ids = tuple(str(row["id"]) for row in rows) + if any(not unit_id for unit_id in unit_ids): # pragma: no cover - database PK invariant + raise ValueError("memory_units recovery query returned an empty ID") + return unit_ids + + async def load_document_unit_bindings( + self, + bank_id: str, + document_id: str, + *, + expected_unit_ids: tuple[str, ...] | None = None, + ) -> tuple[CommittedUnitBinding, ...]: + """Load bindings without PostgreSQL array-position ordering.""" + + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + if expected_unit_ids is not None: + expected_unit_ids = _unit_id_tuple(expected_unit_ids, field_name="expected_unit_ids") + if not expected_unit_ids: + return () + + rows = await self._connection.fetch( + f""" + SELECT mu.id AS unit_id, c.chunk_index + FROM {fq_table_explicit("memory_units", self._schema)} mu + LEFT JOIN {fq_table_explicit("chunks", self._schema)} c + ON c.chunk_id = mu.chunk_id + AND c.bank_id = mu.bank_id + AND c.document_id = mu.document_id + WHERE mu.bank_id = $1 AND mu.document_id = $2 + ORDER BY c.chunk_index NULLS LAST, mu.created_at, mu.id + """, + bank_id, + document_id, + ) + + bindings_by_unit_id: dict[str, CommittedUnitBinding] = {} + for row in rows: + chunk_index = row["chunk_index"] + binding = CommittedUnitBinding( + unit_id=str(row["unit_id"]), + chunk_index=(None if chunk_index is None else int(chunk_index)), + ) + if binding.unit_id in bindings_by_unit_id: + raise ValueError(f"memory_units recovery query returned duplicate unit ID: {binding.unit_id!r}") + bindings_by_unit_id[binding.unit_id] = binding + + if expected_unit_ids is None: + return tuple(bindings_by_unit_id.values()) + missing = tuple(unit_id for unit_id in expected_unit_ids if unit_id not in bindings_by_unit_id) + if missing: + raise ValueError( + f"checkpoint unit IDs are missing or do not belong to the requested bank/document: {missing!r}" + ) + return tuple(bindings_by_unit_id[unit_id] for unit_id in expected_unit_ids) + + +class OracleCheckpointStore: + """Maintain Retain checkpoints in an Oracle JSON CLOB under a row lock.""" + + def __init__(self, connection: Any, *, schema: str | None = None) -> None: + self._connection = connection + self._schema = schema + + @property + def _table(self) -> str: + return fq_table_explicit("async_operations", self._schema) + + async def recover_document_ids(self, operation_id: str) -> tuple[str, ...]: + return (await self.recover(operation_id)).document_ids + + async def recover(self, operation_id: str) -> OperationCheckpoint: + parsed_operation_id = uuid.UUID(operation_id) + row = await self._connection.fetchrow( + f"SELECT result_metadata FROM {self._table} WHERE operation_id = $1", + parsed_operation_id, + ) + if not row or not row["result_metadata"]: + return OperationCheckpoint() + metadata = _json_object(row["result_metadata"]) + unscoped_facts_committed = metadata.get("facts_committed", False) + if not isinstance(unscoped_facts_committed, bool): + raise ValueError("async operation facts_committed must be a bool") + return OperationCheckpoint( + document_ids=_string_array(metadata.get("document_ids"), field_name="document_ids"), + core_committed_document_ids=_string_array( + metadata.get("facts_committed_document_ids"), + field_name="facts_committed_document_ids", + ), + final_ann_pending_document_ids=_string_array( + metadata.get("final_ann_pending_document_ids"), + field_name="final_ann_pending_document_ids", + ), + committed_unit_ids_by_document=_committed_unit_ids_by_document(metadata), + unscoped_facts_committed=unscoped_facts_committed, + ) + + async def _locked_metadata(self, operation_id: uuid.UUID) -> dict[str, Any]: + row = await self._connection.fetchrow( + f""" + SELECT result_metadata + FROM {self._table} + WHERE operation_id = $1 + FOR UPDATE + """, + operation_id, + ) + if row is None: + raise RuntimeError("Async operation disappeared before checkpoint update") + return _json_object(row["result_metadata"]) + + async def _write_metadata(self, operation_id: uuid.UUID, metadata: Mapping[str, Any]) -> None: + status = await self._connection.execute( + f""" + UPDATE {self._table} + SET result_metadata = $1, updated_at = now() + WHERE operation_id = $2 + """, + json.dumps(metadata, sort_keys=True, separators=(",", ":")), + operation_id, + ) + if _affected_rows(status) != 1: + raise RuntimeError("Async operation disappeared before checkpoint update") + + @staticmethod + def _append_unique(metadata: dict[str, Any], field_name: str, value: str) -> None: + values = list(_string_array(metadata.get(field_name), field_name=field_name)) + if value not in values: + values.append(value) + metadata[field_name] = values + + async def record_document_id(self, operation_id: str, document_id: str) -> None: + parsed_operation_id = uuid.UUID(operation_id) + _require_identifier(document_id, field_name="document_id") + async with self._connection.transaction(): + metadata = await self._locked_metadata(parsed_operation_id) + self._append_unique(metadata, "document_ids", document_id) + await self._write_metadata(parsed_operation_id, metadata) + + async def record_core_committed( + self, + operation_id: str, + document_id: str, + *, + unit_ids: tuple[str, ...], + requires_final_ann: bool, + ) -> None: + parsed_operation_id = uuid.UUID(operation_id) + _require_identifier(document_id, field_name="document_id") + unit_ids = _unit_id_tuple(unit_ids, field_name="unit_ids") + if not isinstance(requires_final_ann, bool): + raise TypeError("requires_final_ann must be a bool") + + async with self._connection.transaction(): + metadata = await self._locked_metadata(parsed_operation_id) + # Validate an existing versioned payload before changing any field. + _committed_unit_ids_by_document(metadata) + self._append_unique(metadata, "document_ids", document_id) + self._append_unique(metadata, "facts_committed_document_ids", document_id) + metadata["facts_committed"] = True + metadata["unit_ids_count"] = len(unit_ids) + + existing_payload = metadata.get(_COMMITTED_UNIT_IDS_FIELD) + if existing_payload is None: + payload: dict[str, Any] = { + "version": _COMMITTED_UNIT_IDS_VERSION, + "documents": {}, + } + else: + payload = dict(existing_payload) + payload["documents"] = dict(payload["documents"]) + payload["documents"][document_id] = list(unit_ids) + metadata[_COMMITTED_UNIT_IDS_FIELD] = payload + + pending = list( + _string_array( + metadata.get("final_ann_pending_document_ids"), + field_name="final_ann_pending_document_ids", + ) + ) + if requires_final_ann and unit_ids and document_id not in pending: + pending.append(document_id) + metadata["final_ann_pending_document_ids"] = pending + await self._write_metadata(parsed_operation_id, metadata) + + async def record_final_ann_completed(self, operation_id: str, document_id: str) -> None: + parsed_operation_id = uuid.UUID(operation_id) + _require_identifier(document_id, field_name="document_id") + async with self._connection.transaction(): + metadata = await self._locked_metadata(parsed_operation_id) + pending = _string_array( + metadata.get("final_ann_pending_document_ids"), + field_name="final_ann_pending_document_ids", + ) + metadata["final_ann_pending_document_ids"] = [value for value in pending if value != document_id] + await self._write_metadata(parsed_operation_id, metadata) + + async def clear_provider_batch(self, operation_id: str) -> None: + parsed_operation_id = uuid.UUID(operation_id) + async with self._connection.transaction(): + metadata = await self._locked_metadata(parsed_operation_id) + for field_name in ("batch_id", "batch_provider", "chunk_count"): + metadata.pop(field_name, None) + await self._write_metadata(parsed_operation_id, metadata) + + +__all__ = [ + "FreshOracleDocumentOwnership", + "OracleCheckpointStore", + "OracleDocumentOwnership", + "OraclePlanningRepository", +] diff --git a/core/dataplane/hms_api/engine/ingestion/persistence/ports.py b/core/dataplane/hms_api/engine/ingestion/persistence/ports.py new file mode 100644 index 0000000..d186239 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/persistence/ports.py @@ -0,0 +1,53 @@ +"""Semantic persistence interfaces consumed by Retain planning.""" + +from __future__ import annotations + +from typing import Protocol + +from ..domain import ExistingChunkFingerprint +from .models import CommittedUnitBinding, ExistingDocument, OperationCheckpoint + + +class PlanningRepository(Protocol): + """Read the minimum durable state required to build a change plan.""" + + async def load_document(self, bank_id: str, document_id: str) -> ExistingDocument | None: ... + + async def load_chunks( + self, + bank_id: str, + document_id: str, + ) -> tuple[ExistingChunkFingerprint, ...]: ... + + async def load_document_unit_ids(self, bank_id: str, document_id: str) -> tuple[str, ...]: ... + + async def load_document_unit_bindings( + self, + bank_id: str, + document_id: str, + *, + expected_unit_ids: tuple[str, ...] | None = None, + ) -> tuple[CommittedUnitBinding, ...]: ... + + +class CheckpointStore(Protocol): + """Own async-operation metadata outside the core memory UoW.""" + + async def recover_document_ids(self, operation_id: str) -> tuple[str, ...]: ... + + async def recover(self, operation_id: str) -> OperationCheckpoint: ... + + async def record_document_id(self, operation_id: str, document_id: str) -> None: ... + + async def record_core_committed( + self, + operation_id: str, + document_id: str, + *, + unit_ids: tuple[str, ...], + requires_final_ann: bool, + ) -> None: ... + + async def record_final_ann_completed(self, operation_id: str, document_id: str) -> None: ... + + async def clear_provider_batch(self, operation_id: str) -> None: ... diff --git a/core/dataplane/hms_api/engine/ingestion/persistence/postgres.py b/core/dataplane/hms_api/engine/ingestion/persistence/postgres.py new file mode 100644 index 0000000..515bfde --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/persistence/postgres.py @@ -0,0 +1,566 @@ +"""PostgreSQL planning and checkpoint adapters for Retain.""" + +from __future__ import annotations + +import json +import uuid +from collections.abc import Mapping +from typing import Any + +from ...schema import fq_table_explicit +from ..domain import ExistingChunkFingerprint +from .models import CommittedUnitBinding, ExistingDocument, OperationCheckpoint + +_COMMITTED_UNIT_IDS_FIELD = "committed_unit_ids_v1" +_COMMITTED_UNIT_IDS_VERSION = 1 + + +def _require_identifier(value: str, *, field_name: str) -> str: + if not isinstance(value, str) or not value: + raise ValueError(f"{field_name} must be a non-empty string") + return value + + +def _json_object(value: Any) -> dict[str, Any]: + if value is None: + return {} + if isinstance(value, Mapping): + return dict(value) + if isinstance(value, str): + parsed = json.loads(value) + if isinstance(parsed, dict): + return parsed + raise ValueError(f"Expected a JSON object from PostgreSQL, got {type(value).__name__}") + + +def _tags(value: Any) -> tuple[str, ...]: + if value is None: + return () + if isinstance(value, str): + value = json.loads(value) + if not isinstance(value, (list, tuple)) or not all(isinstance(item, str) for item in value): + raise ValueError("Expected document tags to be an array of strings") + return tuple(value) + + +def _string_array(value: Any, *, field_name: str) -> tuple[str, ...]: + if value is None: + return () + if not isinstance(value, list) or any(not isinstance(item, str) or not item for item in value): + raise ValueError(f"async operation {field_name} must be an array of non-empty strings") + if len(value) != len(set(value)): + raise ValueError(f"async operation {field_name} must not contain duplicates") + return tuple(value) + + +def _unit_id_tuple(value: Any, *, field_name: str) -> tuple[str, ...]: + if not isinstance(value, tuple): + raise TypeError(f"{field_name} must be a tuple") + if any(not isinstance(unit_id, str) or not unit_id for unit_id in value): + raise ValueError(f"{field_name} must contain only non-empty strings") + if len(value) != len(set(value)): + raise ValueError(f"{field_name} must not contain duplicates") + return value + + +def _committed_unit_ids_by_document(metadata: Mapping[str, Any]) -> tuple[tuple[str, tuple[str, ...]], ...]: + if _COMMITTED_UNIT_IDS_FIELD not in metadata: + return () + + payload = metadata[_COMMITTED_UNIT_IDS_FIELD] + if not isinstance(payload, Mapping): + raise ValueError(f"async operation {_COMMITTED_UNIT_IDS_FIELD} must be an object") + if set(payload) != {"version", "documents"}: + raise ValueError(f"async operation {_COMMITTED_UNIT_IDS_FIELD} must contain exactly version and documents") + version = payload["version"] + if isinstance(version, bool) or not isinstance(version, int) or version != _COMMITTED_UNIT_IDS_VERSION: + raise ValueError(f"async operation {_COMMITTED_UNIT_IDS_FIELD}.version must be {_COMMITTED_UNIT_IDS_VERSION}") + documents = payload["documents"] + if not isinstance(documents, Mapping): + raise ValueError(f"async operation {_COMMITTED_UNIT_IDS_FIELD}.documents must be an object") + + parsed: list[tuple[str, tuple[str, ...]]] = [] + for document_id, unit_ids in documents.items(): + _require_identifier(document_id, field_name=f"{_COMMITTED_UNIT_IDS_FIELD} document_id") + parsed.append( + ( + document_id, + _string_array( + unit_ids, + field_name=f"{_COMMITTED_UNIT_IDS_FIELD}.documents[{document_id!r}]", + ), + ) + ) + return tuple(sorted(parsed, key=lambda item: item[0])) + + +class PostgresDocumentOwnership: + """PostgreSQL row-lock adapter for full-retain write windows. + + The adapter is intentionally connection-agnostic: the caller supplies the + connection already enlisted in the core write transaction. Every table + reference is built from the explicit request schema rather than ambient + schema context. + """ + + def __init__(self, *, schema: str | None = None) -> None: + self._schema = schema + + async def prepare_first_window(self, connection: Any, *, bank_id: str, document_id: str) -> None: + """Ensure a lockable row exists, then lock it for first-window tracking.""" + + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + documents = fq_table_explicit("documents", self._schema) + await connection.execute( + f""" + INSERT INTO {documents} (id, bank_id, original_text, content_hash) + VALUES ($1, $2, '', '__pending__') + ON CONFLICT (id, bank_id) DO NOTHING + """, + document_id, + bank_id, + ) + await connection.fetchval( + f""" + SELECT content_hash + FROM {documents} + WHERE id = $1 AND bank_id = $2 + FOR UPDATE + """, + document_id, + bank_id, + ) + + async def validate_later_window( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + expected_content_hash: str, + ) -> bool: + """Lock an existing row and report whether this request still owns it. + + A missing row, a NULL hash, or a different hash means ownership was + lost. Later windows never recreate the row: doing so could let an old + producer resume after another request deleted or replaced the document. + """ + + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + _require_identifier(expected_content_hash, field_name="expected_content_hash") + current_hash = await connection.fetchval( + f""" + SELECT content_hash + FROM {fq_table_explicit("documents", self._schema)} + WHERE id = $1 AND bank_id = $2 + FOR UPDATE + """, + document_id, + bank_id, + ) + return current_hash is not None and current_hash == expected_content_hash + + async def validate_unhashed_window( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + ) -> bool: + """Lock an upgraded row only while its content hash is absent.""" + + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + row = await connection.fetchrow( + f""" + SELECT content_hash + FROM {fq_table_explicit("documents", self._schema)} + WHERE id = $1 AND bank_id = $2 + FOR UPDATE + """, + document_id, + bank_id, + ) + return row is not None and not row["content_hash"] + + async def transition_content_hash( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + expected_content_hash: str, + new_content_hash: str, + ) -> bool: + """Atomically move a locked document between ownership hash states.""" + + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + _require_identifier(expected_content_hash, field_name="expected_content_hash") + _require_identifier(new_content_hash, field_name="new_content_hash") + updated = await connection.fetchval( + f""" + UPDATE {fq_table_explicit("documents", self._schema)} + SET content_hash = $1, updated_at = now() + WHERE id = $2 AND bank_id = $3 AND content_hash = $4 + RETURNING id + """, + new_content_hash, + document_id, + bank_id, + expected_content_hash, + ) + return updated is not None + + +class PostgresPlanningRepository: + """Read document fingerprints through an already-scoped connection.""" + + def __init__(self, connection: Any, *, schema: str | None = None) -> None: + self._connection = connection + self._schema = schema + + async def load_document(self, bank_id: str, document_id: str) -> ExistingDocument | None: + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + row = await self._connection.fetchrow( + f""" + SELECT id, bank_id, original_text, content_hash, retain_params, + tags, created_at, updated_at + FROM {fq_table_explicit("documents", self._schema)} + WHERE id = $1 AND bank_id = $2 + """, + document_id, + bank_id, + ) + if row is None: + return None + return ExistingDocument( + document_id=str(row["id"]), + bank_id=str(row["bank_id"]), + original_text=row["original_text"] or "", + content_hash=row["content_hash"], + retain_params=_json_object(row["retain_params"]), + tags=_tags(row["tags"]), + created_at=row["created_at"], + updated_at=row["updated_at"], + ) + + async def load_chunks( + self, + bank_id: str, + document_id: str, + ) -> tuple[ExistingChunkFingerprint, ...]: + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + rows = await self._connection.fetch( + f""" + SELECT chunk_id, chunk_index, content_hash + FROM {fq_table_explicit("chunks", self._schema)} + WHERE document_id = $1 AND bank_id = $2 + ORDER BY chunk_index + """, + document_id, + bank_id, + ) + return tuple( + ExistingChunkFingerprint( + chunk_id=str(row["chunk_id"]), + chunk_index=int(row["chunk_index"]), + content_hash=row["content_hash"], + ) + for row in rows + ) + + async def load_document_unit_ids(self, bank_id: str, document_id: str) -> tuple[str, ...]: + """Load committed unit IDs deterministically for crash recovery.""" + + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + rows = await self._connection.fetch( + f""" + SELECT id::text AS id + FROM {fq_table_explicit("memory_units", self._schema)} + WHERE bank_id = $1 AND document_id = $2 + ORDER BY created_at, id + """, + bank_id, + document_id, + ) + unit_ids = tuple(str(row["id"]) for row in rows) + if any(not unit_id for unit_id in unit_ids): # pragma: no cover - database PK invariant + raise ValueError("memory_units recovery query returned an empty ID") + return unit_ids + + async def load_document_unit_bindings( + self, + bank_id: str, + document_id: str, + *, + expected_unit_ids: tuple[str, ...] | None = None, + ) -> tuple[CommittedUnitBinding, ...]: + """Load committed unit IDs together with their durable chunk positions. + + The left join is deliberate: memory units may have no + ``chunk_id`` (or may reference a chunk that no longer exists), and such + units still belong to the recovered document result. + """ + + _require_identifier(bank_id, field_name="bank_id") + _require_identifier(document_id, field_name="document_id") + if expected_unit_ids is not None: + expected_unit_ids = _unit_id_tuple(expected_unit_ids, field_name="expected_unit_ids") + if not expected_unit_ids: + return () + + expected_filter = "" + order_by = "c.chunk_index NULLS LAST, mu.created_at, mu.id" + query_args: tuple[Any, ...] = (bank_id, document_id) + if expected_unit_ids is not None: + expected_filter = "AND mu.id::text = ANY($3::text[])" + order_by = "array_position($3::text[], mu.id::text)" + query_args = (bank_id, document_id, expected_unit_ids) + rows = await self._connection.fetch( + f""" + SELECT mu.id::text AS unit_id, c.chunk_index + FROM {fq_table_explicit("memory_units", self._schema)} AS mu + LEFT JOIN {fq_table_explicit("chunks", self._schema)} AS c + ON c.chunk_id = mu.chunk_id + AND c.bank_id = mu.bank_id + AND c.document_id = mu.document_id + WHERE mu.bank_id = $1 AND mu.document_id = $2 + {expected_filter} + ORDER BY {order_by} + """, + *query_args, + ) + + bindings_by_unit_id: dict[str, CommittedUnitBinding] = {} + for row in rows: + binding = CommittedUnitBinding( + unit_id=row["unit_id"], + chunk_index=row["chunk_index"], + ) + if binding.unit_id in bindings_by_unit_id: + raise ValueError(f"memory_units recovery query returned duplicate unit ID: {binding.unit_id!r}") + bindings_by_unit_id[binding.unit_id] = binding + + if expected_unit_ids is None: + return tuple(bindings_by_unit_id.values()) + + unexpected_unit_ids = set(bindings_by_unit_id).difference(expected_unit_ids) + if unexpected_unit_ids: # pragma: no cover - SQL predicate invariant + raise ValueError( + f"memory_units recovery query returned unexpected unit IDs: {sorted(unexpected_unit_ids)!r}" + ) + missing_unit_ids = tuple(unit_id for unit_id in expected_unit_ids if unit_id not in bindings_by_unit_id) + if missing_unit_ids: + raise ValueError( + f"checkpoint unit IDs are missing or do not belong to the requested bank/document: {missing_unit_ids!r}" + ) + return tuple(bindings_by_unit_id[unit_id] for unit_id in expected_unit_ids) + + +class PostgresCheckpointStore: + """Read and idempotently update async operation document checkpoints.""" + + def __init__(self, connection: Any, *, schema: str | None = None) -> None: + self._connection = connection + self._schema = schema + + async def recover_document_ids(self, operation_id: str) -> tuple[str, ...]: + return (await self.recover(operation_id)).document_ids + + async def recover(self, operation_id: str) -> OperationCheckpoint: + parsed_operation_id = uuid.UUID(operation_id) + row = await self._connection.fetchrow( + f""" + SELECT result_metadata + FROM {fq_table_explicit("async_operations", self._schema)} + WHERE operation_id = $1 + """, + parsed_operation_id, + ) + if not row or not row["result_metadata"]: + return OperationCheckpoint() + metadata = _json_object(row["result_metadata"]) + unscoped_facts_committed = metadata.get("facts_committed", False) + if not isinstance(unscoped_facts_committed, bool): + raise ValueError("async operation facts_committed must be a bool") + return OperationCheckpoint( + document_ids=_string_array(metadata.get("document_ids"), field_name="document_ids"), + core_committed_document_ids=_string_array( + metadata.get("facts_committed_document_ids"), + field_name="facts_committed_document_ids", + ), + final_ann_pending_document_ids=_string_array( + metadata.get("final_ann_pending_document_ids"), + field_name="final_ann_pending_document_ids", + ), + committed_unit_ids_by_document=_committed_unit_ids_by_document(metadata), + unscoped_facts_committed=unscoped_facts_committed, + ) + + async def record_document_id(self, operation_id: str, document_id: str) -> None: + parsed_operation_id = uuid.UUID(operation_id) + _require_identifier(document_id, field_name="document_id") + await self._connection.execute( + f""" + UPDATE {fq_table_explicit("async_operations", self._schema)} + SET result_metadata = jsonb_set( + COALESCE(result_metadata, '{{}}'::jsonb), + '{{document_ids}}', + CASE + WHEN COALESCE(result_metadata->'document_ids', '[]'::jsonb) @> $1::jsonb + THEN result_metadata->'document_ids' + ELSE COALESCE(result_metadata->'document_ids', '[]'::jsonb) || $1::jsonb + END, + true + ), + updated_at = now() + WHERE operation_id = $2 + """, + json.dumps([document_id]), + parsed_operation_id, + ) + + async def record_core_committed( + self, + operation_id: str, + document_id: str, + *, + unit_ids: tuple[str, ...], + requires_final_ann: bool, + ) -> None: + """Atomically checkpoint a document from inside its core write transaction.""" + + parsed_operation_id = uuid.UUID(operation_id) + _require_identifier(document_id, field_name="document_id") + unit_ids = _unit_id_tuple(unit_ids, field_name="unit_ids") + if not isinstance(requires_final_ann, bool): + raise TypeError("requires_final_ann must be a bool") + updated_operation_id = await self._connection.fetchval( + f""" + UPDATE {fq_table_explicit("async_operations", self._schema)} + SET result_metadata = jsonb_set( + jsonb_set( + jsonb_set( + jsonb_set( + COALESCE(result_metadata, '{{}}'::jsonb) || $1::jsonb, + '{{document_ids}}', + CASE + WHEN COALESCE(result_metadata->'document_ids', '[]'::jsonb) @> $2::jsonb + THEN result_metadata->'document_ids' + ELSE COALESCE(result_metadata->'document_ids', '[]'::jsonb) || $2::jsonb + END, + true + ), + '{{facts_committed_document_ids}}', + CASE + WHEN COALESCE(result_metadata->'facts_committed_document_ids', '[]'::jsonb) @> $2::jsonb + THEN result_metadata->'facts_committed_document_ids' + ELSE COALESCE(result_metadata->'facts_committed_document_ids', '[]'::jsonb) || $2::jsonb + END, + true + ), + '{{{_COMMITTED_UNIT_IDS_FIELD}}}', + jsonb_set( + COALESCE( + result_metadata->'{_COMMITTED_UNIT_IDS_FIELD}', + '{{"version": {_COMMITTED_UNIT_IDS_VERSION}, "documents": {{}}}}'::jsonb + ), + ARRAY['documents', $3::text], + $4::jsonb, + true + ), + true + ), + '{{final_ann_pending_document_ids}}', + CASE + WHEN NOT $5::boolean + THEN COALESCE(result_metadata->'final_ann_pending_document_ids', '[]'::jsonb) + WHEN COALESCE(result_metadata->'final_ann_pending_document_ids', '[]'::jsonb) @> $2::jsonb + THEN result_metadata->'final_ann_pending_document_ids' + ELSE COALESCE(result_metadata->'final_ann_pending_document_ids', '[]'::jsonb) || $2::jsonb + END, + true + ), + updated_at = now() + WHERE operation_id = $6 + AND ( + result_metadata->'{_COMMITTED_UNIT_IDS_FIELD}' IS NULL + OR ( + jsonb_typeof(result_metadata->'{_COMMITTED_UNIT_IDS_FIELD}') = 'object' + AND result_metadata->'{_COMMITTED_UNIT_IDS_FIELD}'->'version' = + '{_COMMITTED_UNIT_IDS_VERSION}'::jsonb + AND jsonb_typeof(result_metadata->'{_COMMITTED_UNIT_IDS_FIELD}'->'documents') = 'object' + ) + ) + RETURNING operation_id + """, + json.dumps({"facts_committed": True, "unit_ids_count": len(unit_ids)}), + json.dumps([document_id]), + document_id, + json.dumps(unit_ids), + requires_final_ann and bool(unit_ids), + parsed_operation_id, + ) + if updated_operation_id is None: + raise RuntimeError( + f"Async operation {operation_id} disappeared or had an invalid unit-ID checkpoint before core commit" + ) + + async def record_final_ann_completed(self, operation_id: str, document_id: str) -> None: + """Clear an idempotent post-commit ANN recovery marker after an attempt.""" + + parsed_operation_id = uuid.UUID(operation_id) + _require_identifier(document_id, field_name="document_id") + updated_operation_id = await self._connection.fetchval( + f""" + UPDATE {fq_table_explicit("async_operations", self._schema)} + SET result_metadata = jsonb_set( + COALESCE(result_metadata, '{{}}'::jsonb), + '{{final_ann_pending_document_ids}}', + COALESCE( + ( + SELECT jsonb_agg(value) + FROM jsonb_array_elements( + COALESCE(result_metadata->'final_ann_pending_document_ids', '[]'::jsonb) + ) AS pending(value) + WHERE value <> to_jsonb($1::text) + ), + '[]'::jsonb + ), + true + ), + updated_at = now() + WHERE operation_id = $2 + RETURNING operation_id + """, + document_id, + parsed_operation_id, + ) + if updated_operation_id is None: + raise RuntimeError(f"Async operation {operation_id} disappeared before final ANN checkpoint") + + async def clear_provider_batch(self, operation_id: str) -> None: + """Retire one completed provider Batch job before another extraction window.""" + + parsed_operation_id = uuid.UUID(operation_id) + updated_operation_id = await self._connection.fetchval( + f""" + UPDATE {fq_table_explicit("async_operations", self._schema)} + SET result_metadata = COALESCE(result_metadata, '{{}}'::jsonb) + - 'batch_id' + - 'batch_provider' + - 'chunk_count', + updated_at = now() + WHERE operation_id = $1 + RETURNING operation_id + """, + parsed_operation_id, + ) + if updated_operation_id is None: + raise RuntimeError(f"Async operation {operation_id} disappeared before provider Batch cleanup") diff --git a/core/dataplane/hms_api/engine/ingestion/persistence/unit_of_work.py b/core/dataplane/hms_api/engine/ingestion/persistence/unit_of_work.py new file mode 100644 index 0000000..18255a6 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/persistence/unit_of_work.py @@ -0,0 +1,631 @@ +"""Semantic write-window contracts and transaction orchestration for Retain. + +The application layer submits one :class:`WriteWindowRequest`; a persistence +adapter owns all backend-specific work performed inside the transaction, while +this unit-of-work wrapper owns the transaction and post-commit failure boundary. +""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping, Sequence +from contextlib import AbstractAsyncContextManager +from dataclasses import dataclass, field +from enum import StrEnum +from typing import Any, Protocol, TypeAlias + +from ...entity_resolution_contracts import EntityResolutionReadPlan +from ...retain.types import ChunkMetadata, ExtractedFact, ProcessedFact, RetainContent +from ..contracts import CoreCommitCallback, OutboxCallback +from ..domain import FrozenObject, freeze_json + + +class PersistenceContractError(ValueError): + """A planned write cannot be mapped one-to-one to durable records.""" + + +@dataclass(frozen=True, slots=True) +class FirstFullWriteWindow: + """Document work that must occur in the first full-retain write window.""" + + combined_content: str + is_first_batch: bool = True + retain_params: Mapping[str, Any] | None = None + document_tags: tuple[str, ...] = () + recovery: bool = False + # A full replacement planned from an existing document must lock the row + # and prove that this exact hash is still current before deleting anything. + # ``None`` is reserved for a genuinely fresh document claim. + expected_existing_content_hash: str | None = None + # Databases upgraded from releases that did not populate ``content_hash`` + # need a distinct ownership state. The writer must lock the existing row + # and prove that its hash is still absent before replacing any data. + expects_unhashed_existing_document: bool = False + # Multi-window FULL writes replace the final document hash with a unique + # in-flight ownership token after tracking the first window. A single + # window leaves this unset and publishes the final hash immediately. + continuation_content_hash: str | None = None + + def __post_init__(self) -> None: + if not isinstance(self.combined_content, str): + raise TypeError("combined_content must be a string") + if not isinstance(self.is_first_batch, bool): + raise TypeError("is_first_batch must be a bool") + if self.retain_params is not None and not isinstance(self.retain_params, Mapping): + raise TypeError("retain_params must be a mapping or None") + if not isinstance(self.document_tags, tuple) or any(not isinstance(tag, str) for tag in self.document_tags): + raise TypeError("document_tags must be a tuple of strings") + if not isinstance(self.recovery, bool): + raise TypeError("recovery must be a bool") + if self.expected_existing_content_hash is not None and ( + not isinstance(self.expected_existing_content_hash, str) or not self.expected_existing_content_hash + ): + raise ValueError("expected_existing_content_hash must be a non-empty string or None") + if not isinstance(self.expects_unhashed_existing_document, bool): + raise TypeError("expects_unhashed_existing_document must be a bool") + if self.expected_existing_content_hash is not None and self.expects_unhashed_existing_document: + raise ValueError("hashed and unhashed existing-document ownership are mutually exclusive") + if self.recovery and ( + self.expected_existing_content_hash is not None or self.expects_unhashed_existing_document + ): + raise ValueError("recovery and existing-document full replacement are mutually exclusive") + if self.continuation_content_hash is not None and ( + not isinstance(self.continuation_content_hash, str) or not self.continuation_content_hash + ): + raise ValueError("continuation_content_hash must be a non-empty string or None") + + +@dataclass(frozen=True, slots=True) +class LaterFullWriteWindow: + """A later full-retain window that may write only while it owns the document.""" + + expected_content_hash: str + # Only the last window sets this to the final document content hash. An + # intermediate window keeps the in-flight token unchanged. + completed_content_hash: str | None = None + + def __post_init__(self) -> None: + if not isinstance(self.expected_content_hash, str) or not self.expected_content_hash: + raise ValueError("expected_content_hash must be a non-empty string") + if self.completed_content_hash is not None and ( + not isinstance(self.completed_content_hash, str) or not self.completed_content_hash + ): + raise ValueError("completed_content_hash must be a non-empty string or None") + + +DocumentWindow: TypeAlias = FirstFullWriteWindow | LaterFullWriteWindow + + +@dataclass(frozen=True, slots=True) +class ChunkWrite: + """Stable chunk identity paired with its storage DTO.""" + + chunk_key: str + metadata: ChunkMetadata + + def __post_init__(self) -> None: + if not isinstance(self.chunk_key, str) or not self.chunk_key: + raise ValueError("chunk_key must be a non-empty string") + if not isinstance(self.metadata, ChunkMetadata): + raise TypeError("metadata must be ChunkMetadata") + + +@dataclass(frozen=True, slots=True) +class FactWrite: + """Stable fact identity and its one-to-one storage representations.""" + + fact_key: str + chunk_key: str + extracted: ExtractedFact + processed: ProcessedFact + + def __post_init__(self) -> None: + if not isinstance(self.fact_key, str) or not self.fact_key: + raise ValueError("fact_key must be a non-empty string") + if not isinstance(self.chunk_key, str) or not self.chunk_key: + raise ValueError("chunk_key must be a non-empty string") + if not isinstance(self.extracted, ExtractedFact): + raise TypeError("extracted must be ExtractedFact") + if not isinstance(self.processed, ProcessedFact): + raise TypeError("processed must be ProcessedFact") + + +@dataclass(frozen=True, slots=True) +class ExistingChunkWrite: + """Durable identity for an existing chunk affected by a delta write.""" + + chunk_id: str + chunk_index: int + + def __post_init__(self) -> None: + if not isinstance(self.chunk_id, str) or not self.chunk_id: + raise ValueError("chunk_id must be a non-empty string") + if isinstance(self.chunk_index, bool) or not isinstance(self.chunk_index, int): + raise TypeError("chunk_index must be an integer") + if self.chunk_index < 0: + raise ValueError("chunk_index must be non-negative") + + +@dataclass(frozen=True, slots=True) +class CoreGraphWrite: + """Entity-resolution and precomputed-link inputs for the core write.""" + + resolved_entity_ids: tuple[str, ...] = () + entity_to_unit: tuple[tuple[Any, ...], ...] = () + unit_to_entity_ids: tuple[tuple[str, tuple[str, ...]], ...] = () + semantic_ann_links: tuple[tuple[Any, ...], ...] = () + entity_read_plan: EntityResolutionReadPlan | None = None + + def __post_init__(self) -> None: + unit_keys = [unit_id for unit_id, _entity_ids in self.unit_to_entity_ids] + if len(unit_keys) != len(set(unit_keys)): + raise PersistenceContractError("unit_to_entity_ids contains duplicate unit keys") + if self.entity_read_plan is not None and ( + self.resolved_entity_ids or self.entity_to_unit or self.unit_to_entity_ids + ): + raise PersistenceContractError( + "entity_read_plan cannot be combined with already-finalized entity graph data" + ) + + +@dataclass(frozen=True, slots=True) +class WriteWindowRequest: + """All semantic writes that belong to one full-retain transaction.""" + + bank_id: str + document_id: str + document_window: DocumentWindow + contents: tuple[RetainContent, ...] + chunks: tuple[ChunkWrite, ...] = () + facts: tuple[FactWrite, ...] = () + graph: CoreGraphWrite = field(default_factory=CoreGraphWrite) + skip_semantic_links: bool = False + checkpoint_callback: CoreCommitCallback | None = None + outbox_callback: OutboxCallback | None = None + log_buffer: list[str] = field(default_factory=list) + + def __post_init__(self) -> None: + if not isinstance(self.bank_id, str) or not self.bank_id: + raise ValueError("bank_id must be a non-empty string") + if not isinstance(self.document_id, str) or not self.document_id: + raise ValueError("document_id must be a non-empty string") + if not isinstance(self.document_window, (FirstFullWriteWindow, LaterFullWriteWindow)): + raise TypeError("document_window must describe a first or later full write window") + if not isinstance(self.contents, tuple) or any(not isinstance(item, RetainContent) for item in self.contents): + raise TypeError("contents must be a tuple of RetainContent values") + if not self.contents: + raise PersistenceContractError("a write window must contain at least one content item") + if not isinstance(self.chunks, tuple) or any(not isinstance(item, ChunkWrite) for item in self.chunks): + raise TypeError("chunks must be a tuple of ChunkWrite values") + if not isinstance(self.facts, tuple) or any(not isinstance(item, FactWrite) for item in self.facts): + raise TypeError("facts must be a tuple of FactWrite values") + if not isinstance(self.graph, CoreGraphWrite): + raise TypeError("graph must be CoreGraphWrite") + if not isinstance(self.skip_semantic_links, bool): + raise TypeError("skip_semantic_links must be a bool") + if self.checkpoint_callback is not None and not callable(self.checkpoint_callback): + raise TypeError("checkpoint_callback must be callable or None") + if self.outbox_callback is not None and not callable(self.outbox_callback): + raise TypeError("outbox_callback must be callable or None") + if not isinstance(self.log_buffer, list): + raise TypeError("log_buffer must be a list") + _validate_identity_bindings(self.contents, self.chunks, self.facts) + + +def _freeze_retain_params(value: FrozenObject | Mapping[str, Any] | None) -> FrozenObject | None: + if value is None or isinstance(value, FrozenObject): + return value + if not isinstance(value, Mapping): + raise TypeError("retain_params must be a mapping, FrozenObject, or None") + frozen = freeze_json(dict(value)) + if not isinstance(frozen, FrozenObject): # pragma: no cover - dicts freeze to objects + raise AssertionError("retain_params must freeze to an object") + return frozen + + +def _validate_identity_bindings( + contents: Sequence[RetainContent], + chunks: Sequence[ChunkWrite], + facts: Sequence[FactWrite], +) -> None: + chunk_keys = [item.chunk_key for item in chunks] + if len(chunk_keys) != len(set(chunk_keys)): + raise PersistenceContractError("chunks contain duplicate chunk_key values") + + chunk_indices = [item.metadata.chunk_index for item in chunks] + if len(chunk_indices) != len(set(chunk_indices)): + raise PersistenceContractError("chunks contain duplicate chunk_index values") + + fact_keys = [item.fact_key for item in facts] + if len(fact_keys) != len(set(fact_keys)): + raise PersistenceContractError("facts contain duplicate fact_key values") + + chunks_by_key = {item.chunk_key: item for item in chunks} + fact_count_by_chunk = dict.fromkeys(chunk_keys, 0) + content_count = len(contents) + for fact in facts: + chunk = chunks_by_key.get(fact.chunk_key) + if chunk is None: + raise PersistenceContractError(f"fact_key={fact.fact_key!r} refers to unknown chunk_key={fact.chunk_key!r}") + fact_count_by_chunk[fact.chunk_key] += 1 + if fact.extracted.chunk_index != chunk.metadata.chunk_index: + raise PersistenceContractError( + f"fact_key={fact.fact_key!r} chunk_index does not match chunk_key={fact.chunk_key!r}" + ) + if fact.extracted.content_index != fact.processed.content_index: + raise PersistenceContractError(f"fact_key={fact.fact_key!r} extracted/processed content_index mismatch") + if not 0 <= fact.processed.content_index < content_count: + raise PersistenceContractError(f"fact_key={fact.fact_key!r} has an out-of-range content_index") + + for chunk in chunks: + if not 0 <= chunk.metadata.content_index < content_count: + raise PersistenceContractError(f"chunk_key={chunk.chunk_key!r} has an out-of-range content_index") + actual_count = fact_count_by_chunk[chunk.chunk_key] + if chunk.metadata.fact_count != actual_count: + raise PersistenceContractError( + f"chunk_key={chunk.chunk_key!r} declares {chunk.metadata.fact_count} facts, " + f"but {actual_count} fact bindings were supplied" + ) + + +def _validate_request_header(*, bank_id: str, document_id: str, expected_content_hash: str) -> None: + if not isinstance(bank_id, str) or not bank_id: + raise ValueError("bank_id must be a non-empty string") + if not isinstance(document_id, str) or not document_id: + raise ValueError("document_id must be a non-empty string") + if not isinstance(expected_content_hash, str) or not expected_content_hash: + raise ValueError("expected_content_hash must be a non-empty string") + + +@dataclass(frozen=True, slots=True) +class MetadataOnlyWriteRequest: + """Immutable metadata-only delta operation. + + ``input_slot_count`` preserves the public one-bucket-per-input outcome even + though no content enters extraction and the processed-token outcome is + therefore exactly zero. + """ + + bank_id: str + document_id: str + expected_content_hash: str + combined_content: str + input_slot_count: int + retain_params: FrozenObject | Mapping[str, Any] | None = None + document_tags: tuple[str, ...] = () + checkpoint_callback: CoreCommitCallback | None = None + outbox_callback: OutboxCallback | None = None + + def __post_init__(self) -> None: + _validate_request_header( + bank_id=self.bank_id, + document_id=self.document_id, + expected_content_hash=self.expected_content_hash, + ) + if not isinstance(self.combined_content, str): + raise TypeError("combined_content must be a string") + if isinstance(self.input_slot_count, bool) or not isinstance(self.input_slot_count, int): + raise TypeError("input_slot_count must be an integer") + if self.input_slot_count < 0: + raise ValueError("input_slot_count must be non-negative") + if not isinstance(self.document_tags, tuple) or any(not isinstance(tag, str) for tag in self.document_tags): + raise TypeError("document_tags must be a tuple of strings") + if self.checkpoint_callback is not None and not callable(self.checkpoint_callback): + raise TypeError("checkpoint_callback must be callable or None") + if self.outbox_callback is not None and not callable(self.outbox_callback): + raise TypeError("outbox_callback must be callable or None") + object.__setattr__(self, "retain_params", _freeze_retain_params(self.retain_params)) + + +@dataclass(frozen=True, slots=True) +class DeltaWriteRequest: + """Immutable partial-document transaction planned from a hash snapshot.""" + + bank_id: str + document_id: str + expected_content_hash: str + combined_content: str + contents: tuple[RetainContent, ...] + unchanged_chunk_indices: tuple[int, ...] + changed_chunks: tuple[ExistingChunkWrite, ...] + added_chunk_indices: tuple[int, ...] + removed_chunks: tuple[ExistingChunkWrite, ...] + chunks: tuple[ChunkWrite, ...] = () + facts: tuple[FactWrite, ...] = () + graph: CoreGraphWrite = field(default_factory=CoreGraphWrite) + processed_tokens: int = 0 + retain_params: FrozenObject | Mapping[str, Any] | None = None + document_tags: tuple[str, ...] = () + skip_semantic_links: bool = False + checkpoint_callback: CoreCommitCallback | None = None + outbox_callback: OutboxCallback | None = None + log_buffer: tuple[str, ...] = () + + def __post_init__(self) -> None: + _validate_request_header( + bank_id=self.bank_id, + document_id=self.document_id, + expected_content_hash=self.expected_content_hash, + ) + if not isinstance(self.combined_content, str): + raise TypeError("combined_content must be a string") + if not isinstance(self.contents, tuple) or any(not isinstance(item, RetainContent) for item in self.contents): + raise TypeError("contents must be a tuple of RetainContent values") + if not isinstance(self.chunks, tuple) or any(not isinstance(item, ChunkWrite) for item in self.chunks): + raise TypeError("chunks must be a tuple of ChunkWrite values") + if not isinstance(self.facts, tuple) or any(not isinstance(item, FactWrite) for item in self.facts): + raise TypeError("facts must be a tuple of FactWrite values") + if not isinstance(self.graph, CoreGraphWrite): + raise TypeError("graph must be CoreGraphWrite") + if isinstance(self.processed_tokens, bool) or not isinstance(self.processed_tokens, int): + raise TypeError("processed_tokens must be an integer") + if self.processed_tokens < 0: + raise ValueError("processed_tokens must be non-negative") + if not isinstance(self.document_tags, tuple) or any(not isinstance(tag, str) for tag in self.document_tags): + raise TypeError("document_tags must be a tuple of strings") + if not isinstance(self.skip_semantic_links, bool): + raise TypeError("skip_semantic_links must be a bool") + if self.checkpoint_callback is not None and not callable(self.checkpoint_callback): + raise TypeError("checkpoint_callback must be callable or None") + if self.outbox_callback is not None and not callable(self.outbox_callback): + raise TypeError("outbox_callback must be callable or None") + if not isinstance(self.log_buffer, tuple) or any(not isinstance(line, str) for line in self.log_buffer): + raise TypeError("log_buffer must be a tuple of strings") + object.__setattr__(self, "retain_params", _freeze_retain_params(self.retain_params)) + self._validate_chunk_sets() + _validate_identity_bindings(self.contents, self.chunks, self.facts) + + def _validate_chunk_sets(self) -> None: + def validate_indices(values: tuple[int, ...], *, field_name: str) -> set[int]: + if not isinstance(values, tuple): + raise TypeError(f"{field_name} must be a tuple of integers") + if any(isinstance(value, bool) or not isinstance(value, int) for value in values): + raise TypeError(f"{field_name} must be a tuple of integers") + if any(value < 0 for value in values): + raise ValueError(f"{field_name} must contain only non-negative integers") + if len(values) != len(set(values)): + raise PersistenceContractError(f"{field_name} contains duplicate chunk indices") + return set(values) + + unchanged = validate_indices(self.unchanged_chunk_indices, field_name="unchanged_chunk_indices") + added = validate_indices(self.added_chunk_indices, field_name="added_chunk_indices") + if not unchanged: + raise PersistenceContractError("delta writes require at least one unchanged chunk") + + for field_name, values in (("changed_chunks", self.changed_chunks), ("removed_chunks", self.removed_chunks)): + if not isinstance(values, tuple) or any(not isinstance(value, ExistingChunkWrite) for value in values): + raise TypeError(f"{field_name} must be a tuple of ExistingChunkWrite values") + changed_indices = [chunk.chunk_index for chunk in self.changed_chunks] + removed_indices = [chunk.chunk_index for chunk in self.removed_chunks] + if len(changed_indices) != len(set(changed_indices)): + raise PersistenceContractError("changed_chunks contains duplicate chunk indices") + if len(removed_indices) != len(set(removed_indices)): + raise PersistenceContractError("removed_chunks contains duplicate chunk indices") + affected_ids = [chunk.chunk_id for chunk in (*self.changed_chunks, *self.removed_chunks)] + if len(affected_ids) != len(set(affected_ids)): + raise PersistenceContractError("changed/removed chunks contain duplicate chunk IDs") + + named_sets = { + "unchanged": unchanged, + "changed": set(changed_indices), + "added": added, + "removed": set(removed_indices), + } + names = tuple(named_sets) + for position, left_name in enumerate(names): + for right_name in names[position + 1 :]: + overlap = named_sets[left_name] & named_sets[right_name] + if overlap: + raise PersistenceContractError( + f"{left_name}/{right_name} chunk sets overlap at indices {sorted(overlap)}" + ) + + if not (changed_indices or added or removed_indices): + raise PersistenceContractError("delta writes require at least one changed, added, or removed chunk") + supplied_indices = {chunk.metadata.chunk_index for chunk in self.chunks} + expected_indices = set(changed_indices) | added + if supplied_indices != expected_indices: + missing = sorted(expected_indices - supplied_indices) + unexpected = sorted(supplied_indices - expected_indices) + raise PersistenceContractError( + f"delta chunk bindings do not match changed/added sets (missing={missing}, unexpected={unexpected})" + ) + + +RetainWriteRequest: TypeAlias = WriteWindowRequest | DeltaWriteRequest | MetadataOnlyWriteRequest + + +class OperationActivityPort(Protocol): + """Lock and validate the tracked operation chain for one core transaction.""" + + async def assert_active(self, connection: Any, *, bank_id: str) -> None: ... + + +class DocumentOwnershipPort(Protocol): + """Own row-lock and document-hash transition semantics.""" + + async def prepare_first_window(self, connection: Any, *, bank_id: str, document_id: str) -> None: ... + + async def validate_unhashed_window( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + ) -> bool: ... + + async def validate_later_window( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + expected_content_hash: str, + ) -> bool: ... + + async def transition_content_hash( + self, + connection: Any, + *, + bank_id: str, + document_id: str, + expected_content_hash: str, + new_content_hash: str, + ) -> bool: ... + + +class OwnershipDisposition(StrEnum): + OWNED = "owned" + LOST = "lost" + + +@dataclass(frozen=True, slots=True) +class CoreWriteResult: + """Result produced inside the transaction before commit is attempted.""" + + ownership: OwnershipDisposition + unit_ids_by_content: tuple[tuple[str, ...], ...] = () + unit_ids_by_fact_key: tuple[tuple[str, str], ...] = () + phase3_payload: Any = None + post_commit_required: bool = False + processed_tokens: int | None = None + + def __post_init__(self) -> None: + fact_keys = [fact_key for fact_key, _unit_id in self.unit_ids_by_fact_key] + if len(fact_keys) != len(set(fact_keys)): + raise PersistenceContractError("unit_ids_by_fact_key contains duplicate fact keys") + if any(not fact_key or not unit_id for fact_key, unit_id in self.unit_ids_by_fact_key): + raise PersistenceContractError("unit_ids_by_fact_key contains an empty identity") + if self.processed_tokens is not None: + if isinstance(self.processed_tokens, bool) or not isinstance(self.processed_tokens, int): + raise TypeError("processed_tokens must be an integer or None") + if self.processed_tokens < 0: + raise ValueError("processed_tokens must be non-negative") + + +class PostCommitStage(StrEnum): + ENTITY_STATS = "entity_stats" + DISPLAY_ENTITY_LINKS = "display_entity_links" + + +class PostCommitStatus(StrEnum): + COMPLETED = "completed" + SKIPPED = "skipped" + FAILED = "failed" + + +class PostCommitSkipReason(StrEnum): + OWNERSHIP_LOST = "ownership_lost" + NO_FACTS = "no_facts" + + +@dataclass(frozen=True, slots=True) +class PostCommitFailure: + """A best-effort failure after the core transaction has committed.""" + + stage: PostCommitStage + exception: Exception + + +@dataclass(frozen=True, slots=True) +class PostCommitReport: + status: PostCommitStatus + failure: PostCommitFailure | None = None + skip_reason: PostCommitSkipReason | None = None + + def __post_init__(self) -> None: + if self.status is PostCommitStatus.FAILED and self.failure is None: + raise ValueError("a failed post-commit report requires failure details") + if self.status is not PostCommitStatus.FAILED and self.failure is not None: + raise ValueError("only a failed post-commit report may contain failure details") + if self.status is PostCommitStatus.SKIPPED and self.skip_reason is None: + raise ValueError("a skipped post-commit report requires a reason") + if self.status is not PostCommitStatus.SKIPPED and self.skip_reason is not None: + raise ValueError("only a skipped post-commit report may contain a reason") + + +@dataclass(frozen=True, slots=True) +class UnitOfWorkResult: + """Committed core result plus independently classified best-effort work.""" + + core: CoreWriteResult + post_commit: PostCommitReport + + +class PersistenceAdapter(Protocol): + """Backend adapter invoked by :class:`RetainUnitOfWork`.""" + + async def write_core(self, connection: Any, request: RetainWriteRequest) -> CoreWriteResult: ... + + async def flush_entity_stats(self) -> None: ... + + async def write_display_entity_links(self, request: RetainWriteRequest, phase3_payload: Any) -> None: ... + + +ConnectionScope: TypeAlias = Callable[[], AbstractAsyncContextManager[Any]] + + +class RetainUnitOfWork: + """Run one semantic Retain write window.""" + + def __init__(self, *, connection_scope: ConnectionScope, adapter: PersistenceAdapter) -> None: + if not callable(connection_scope): + raise TypeError("connection_scope must be callable") + self._connection_scope = connection_scope + self._adapter = adapter + + async def execute(self, request: RetainWriteRequest) -> UnitOfWorkResult: + """Commit core writes, then classify rather than raise best-effort failures. + + Any exception raised while acquiring a connection, running the core + transaction, invoking the outbox callback, or committing is propagated + unchanged. Only work that begins after a successful transaction exit is + converted into a :class:`PostCommitFailure`. + """ + + async with self._connection_scope() as connection: + async with connection.transaction(): + core = await self._adapter.write_core(connection, request) + + if core.ownership is OwnershipDisposition.LOST: + return UnitOfWorkResult( + core=core, + post_commit=PostCommitReport( + status=PostCommitStatus.SKIPPED, + skip_reason=PostCommitSkipReason.OWNERSHIP_LOST, + ), + ) + if not core.post_commit_required: + return UnitOfWorkResult( + core=core, + post_commit=PostCommitReport( + status=PostCommitStatus.SKIPPED, + skip_reason=PostCommitSkipReason.NO_FACTS, + ), + ) + + try: + await self._adapter.flush_entity_stats() + except Exception as exc: + return UnitOfWorkResult( + core=core, + post_commit=PostCommitReport( + status=PostCommitStatus.FAILED, + failure=PostCommitFailure(stage=PostCommitStage.ENTITY_STATS, exception=exc), + ), + ) + + try: + await self._adapter.write_display_entity_links(request, core.phase3_payload) + except Exception as exc: + return UnitOfWorkResult( + core=core, + post_commit=PostCommitReport( + status=PostCommitStatus.FAILED, + failure=PostCommitFailure(stage=PostCommitStage.DISPLAY_ENTITY_LINKS, exception=exc), + ), + ) + + return UnitOfWorkResult( + core=core, + post_commit=PostCommitReport(status=PostCommitStatus.COMPLETED), + ) diff --git a/core/dataplane/hms_api/engine/ingestion/persistence/writer.py b/core/dataplane/hms_api/engine/ingestion/persistence/writer.py new file mode 100644 index 0000000..5fb8ffc --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/persistence/writer.py @@ -0,0 +1,477 @@ +"""Durable writer for semantic ingestion records. + +This module coordinates the storage and graph helpers behind one persistence +boundary. +""" + +from __future__ import annotations + +from typing import Any + +from ...embedding_fingerprint import EmbeddingFingerprintError, ensure_bank_embedding_fingerprint +from ...retain import chunk_storage, fact_storage +from ...schema import fq_table_explicit +from .. import runtime +from ..adapters.storage_records import compute_document_hash +from ..domain import FrozenObject, thaw_json +from ..redaction import IdentifierSanitizer +from .unit_of_work import ( + CoreGraphWrite, + CoreWriteResult, + DeltaWriteRequest, + DocumentOwnershipPort, + FirstFullWriteWindow, + LaterFullWriteWindow, + MetadataOnlyWriteRequest, + OperationActivityPort, + OwnershipDisposition, + PersistenceContractError, + RetainWriteRequest, + WriteWindowRequest, +) + + +class PersistenceWriter: + """Coordinate semantic persistence operations.""" + + def __init__( + self, + *, + pool: Any, + embeddings_model: Any, + entity_resolver: Any, + config: Any, + ownership: DocumentOwnershipPort, + operation_activity: OperationActivityPort | None = None, + ops: Any = None, + schema: str | None = None, + sanitize_log_identifiers: bool = False, + ) -> None: + self._pool = pool + self._embeddings_model = embeddings_model + self._entity_resolver = entity_resolver + self._config = config + self._ownership = ownership + self._operation_activity = operation_activity + self._schema = schema + self._sanitize_log_identifiers = sanitize_log_identifiers + self._ops = ops if ops is not None else getattr(pool, "ops", None) + if self._ops is None: + raise ValueError("PersistenceWriter requires backend data-access ops") + + async def write_core(self, connection: Any, request: RetainWriteRequest) -> CoreWriteResult: + """Write one Retain transaction through the storage helpers. + + Full and delta writes still invoke ``_insert_facts_and_links`` when + ``facts`` is empty. That lets an outbox callback run in the same durable + transaction as document/chunk tracking. Metadata-only writes invoke the + callback directly after their final core mutation. Post-commit entity + work remains skipped when no facts were inserted. + """ + + if self._operation_activity is not None: + await self._operation_activity.assert_active(connection, bank_id=request.bank_id) + + sanitizer = IdentifierSanitizer.from_values( + enabled=self._sanitize_log_identifiers, + values=(request.bank_id, request.document_id), + ) + try: + await ensure_bank_embedding_fingerprint( + connection, + request.bank_id, + self._embeddings_model, + policy=getattr(self._config, "embedding_fingerprint_policy", "strict"), + for_write=True, + legacy_attestation=getattr( + self._config, + "embedding_fingerprint_legacy_attestation", + None, + ), + log_sanitizer=sanitizer, + ) + except EmbeddingFingerprintError as exc: + if sanitizer.enabled: + exc.args = (sanitizer.text(exc),) + raise + + if isinstance(request, MetadataOnlyWriteRequest): + return await self._write_metadata_only_core(connection, request) + if isinstance(request, DeltaWriteRequest): + return await self._write_delta_core(connection, request) + if not isinstance(request, WriteWindowRequest): # pragma: no cover - closed request union + raise TypeError(f"Unsupported Retain write request: {type(request).__name__}") + + document_window = request.document_window + if isinstance(document_window, FirstFullWriteWindow): + if document_window.expects_unhashed_existing_document: + owns_document = await self._ownership.validate_unhashed_window( + connection, + bank_id=request.bank_id, + document_id=request.document_id, + ) + if not owns_document: + return CoreWriteResult(ownership=OwnershipDisposition.LOST) + elif document_window.expected_existing_content_hash is None: + await self._ownership.prepare_first_window( + connection, + bank_id=request.bank_id, + document_id=request.document_id, + ) + else: + owns_document = await self._ownership.validate_later_window( + connection, + bank_id=request.bank_id, + document_id=request.document_id, + expected_content_hash=document_window.expected_existing_content_hash, + ) + if not owns_document: + return CoreWriteResult(ownership=OwnershipDisposition.LOST) + if document_window.recovery: + await fact_storage.upsert_document_metadata( + connection, + request.bank_id, + request.document_id, + document_window.combined_content, + dict(document_window.retain_params) if document_window.retain_params is not None else None, + list(document_window.document_tags), + ) + else: + await fact_storage.handle_document_tracking( + connection, + request.bank_id, + request.document_id, + document_window.combined_content, + document_window.is_first_batch, + dict(document_window.retain_params) if document_window.retain_params is not None else None, + list(document_window.document_tags), + ops=self._ops, + log_sanitizer=sanitizer, + ) + if document_window.continuation_content_hash is not None: + transitioned = await self._ownership.transition_content_hash( + connection, + bank_id=request.bank_id, + document_id=request.document_id, + expected_content_hash=compute_document_hash(document_window.combined_content), + new_content_hash=document_window.continuation_content_hash, + ) + if not transitioned: # pragma: no cover - row is locked in this transaction + raise PersistenceContractError( + "first FULL window could not publish its continuation ownership hash" + ) + elif isinstance(document_window, LaterFullWriteWindow): + owns_document = await self._ownership.validate_later_window( + connection, + bank_id=request.bank_id, + document_id=request.document_id, + expected_content_hash=document_window.expected_content_hash, + ) + if not owns_document: + return CoreWriteResult(ownership=OwnershipDisposition.LOST) + if document_window.completed_content_hash is not None: + transitioned = await self._ownership.transition_content_hash( + connection, + bank_id=request.bank_id, + document_id=request.document_id, + expected_content_hash=document_window.expected_content_hash, + new_content_hash=document_window.completed_content_hash, + ) + if not transitioned: # pragma: no cover - validate holds the row lock + raise PersistenceContractError("final FULL window could not publish its completed content hash") + else: # pragma: no cover - WriteWindowRequest validates this boundary + raise TypeError(f"Unsupported document window: {type(document_window).__name__}") + + graph = await self._finalize_entity_graph(connection, request) + chunk_ids_by_index: dict[int, str] = {} + if request.chunks: + chunk_ids_by_index = await chunk_storage.store_chunks_batch( + connection, + request.bank_id, + request.document_id, + [chunk.metadata for chunk in request.chunks], + ops=self._ops, + ) + self._assert_complete_chunk_result(request, chunk_ids_by_index) + + chunk_ids_by_key = {chunk.chunk_key: chunk_ids_by_index[chunk.metadata.chunk_index] for chunk in request.chunks} + for fact in request.facts: + # Stable fact/chunk keys were validated before the transaction and + # are used here instead of relying on zip/list completion order. + fact.processed.document_id = request.document_id + fact.processed.chunk_id = chunk_ids_by_key[fact.chunk_key] + + unit_ids_by_content, phase3_payload = await runtime.insert_facts_and_links( + connection, + self._entity_resolver, + request.bank_id, + list(request.contents), + [fact.processed for fact in request.facts], + self._config, + request.log_buffer, + resolved_entity_ids=list(graph.resolved_entity_ids), + entity_to_unit=list(graph.entity_to_unit), + unit_to_entity_ids={unit_id: list(entity_ids) for unit_id, entity_ids in graph.unit_to_entity_ids}, + semantic_ann_links=list(graph.semantic_ann_links), + skip_semantic_links=request.skip_semantic_links, + # When a strict checkpoint is present it must observe the exact + # IDs returned by this core write before the external outbox runs. + # Delay both callbacks until the result contract is validated. + outbox_callback=(request.outbox_callback if request.checkpoint_callback is None else None), + ops=self._ops, + ) + self._assert_complete_fact_result(request, unit_ids_by_content) + if request.checkpoint_callback is not None: + immutable_buckets = tuple(tuple(ids) for ids in unit_ids_by_content) + await request.checkpoint_callback(connection, immutable_buckets) + if request.outbox_callback is not None: + await request.outbox_callback(connection) + return CoreWriteResult( + ownership=OwnershipDisposition.OWNED, + unit_ids_by_content=tuple(tuple(ids) for ids in unit_ids_by_content), + unit_ids_by_fact_key=self._map_unit_ids_by_fact_key(request, unit_ids_by_content), + phase3_payload=phase3_payload, + post_commit_required=bool(request.facts), + ) + + async def _write_metadata_only_core( + self, + connection: Any, + request: MetadataOnlyWriteRequest, + ) -> CoreWriteResult: + owns_document = await self._ownership.validate_later_window( + connection, + bank_id=request.bank_id, + document_id=request.document_id, + expected_content_hash=request.expected_content_hash, + ) + if not owns_document: + return CoreWriteResult( + ownership=OwnershipDisposition.LOST, + processed_tokens=0, + ) + + await fact_storage.upsert_document_metadata( + connection, + request.bank_id, + request.document_id, + request.combined_content, + self._thaw_retain_params(request.retain_params), + list(request.document_tags), + ) + await fact_storage.update_memory_units_tags( + connection, + request.bank_id, + request.document_id, + list(request.document_tags), + ) + if request.checkpoint_callback is not None: + await request.checkpoint_callback( + connection, + tuple(() for _ in range(request.input_slot_count)), + ) + if request.outbox_callback is not None: + await request.outbox_callback(connection) + return CoreWriteResult( + ownership=OwnershipDisposition.OWNED, + unit_ids_by_content=tuple(() for _ in range(request.input_slot_count)), + processed_tokens=0, + ) + + async def _write_delta_core(self, connection: Any, request: DeltaWriteRequest) -> CoreWriteResult: + owns_document = await self._ownership.validate_later_window( + connection, + bank_id=request.bank_id, + document_id=request.document_id, + expected_content_hash=request.expected_content_hash, + ) + if not owns_document: + return CoreWriteResult( + ownership=OwnershipDisposition.LOST, + processed_tokens=request.processed_tokens, + ) + + graph = await self._finalize_entity_graph(connection, request) + await fact_storage.upsert_document_metadata( + connection, + request.bank_id, + request.document_id, + request.combined_content, + self._thaw_retain_params(request.retain_params), + list(request.document_tags), + ) + + chunk_ids_to_delete = [chunk.chunk_id for chunk in (*request.changed_chunks, *request.removed_chunks)] + await chunk_storage.delete_chunks_by_ids(connection, chunk_ids_to_delete) + await fact_storage.update_memory_units_tags( + connection, + request.bank_id, + request.document_id, + list(request.document_tags), + ) + + chunk_ids_by_index: dict[int, str] = {} + if request.chunks: + chunk_ids_by_index = await chunk_storage.store_chunks_batch( + connection, + request.bank_id, + request.document_id, + [chunk.metadata for chunk in request.chunks], + ops=self._ops, + ) + self._assert_complete_chunk_result(request, chunk_ids_by_index) + + chunk_ids_by_key = {chunk.chunk_key: chunk_ids_by_index[chunk.metadata.chunk_index] for chunk in request.chunks} + for fact in request.facts: + fact.processed.document_id = request.document_id + fact.processed.chunk_id = chunk_ids_by_key[fact.chunk_key] + + log_buffer = list(request.log_buffer) + unit_ids_by_content, phase3_payload = await runtime.insert_facts_and_links( + connection, + self._entity_resolver, + request.bank_id, + list(request.contents), + [fact.processed for fact in request.facts], + self._config, + log_buffer, + resolved_entity_ids=list(graph.resolved_entity_ids), + entity_to_unit=list(graph.entity_to_unit), + unit_to_entity_ids={unit_id: list(entity_ids) for unit_id, entity_ids in graph.unit_to_entity_ids}, + semantic_ann_links=list(graph.semantic_ann_links), + skip_semantic_links=request.skip_semantic_links, + outbox_callback=(request.outbox_callback if request.checkpoint_callback is None else None), + ops=self._ops, + ) + self._assert_complete_fact_result(request, unit_ids_by_content) + if request.checkpoint_callback is not None: + immutable_buckets = tuple(tuple(ids) for ids in unit_ids_by_content) + await request.checkpoint_callback(connection, immutable_buckets) + if request.outbox_callback is not None: + await request.outbox_callback(connection) + return CoreWriteResult( + ownership=OwnershipDisposition.OWNED, + unit_ids_by_content=tuple(tuple(ids) for ids in unit_ids_by_content), + unit_ids_by_fact_key=self._map_unit_ids_by_fact_key(request, unit_ids_by_content), + phase3_payload=phase3_payload, + post_commit_required=bool(request.facts), + processed_tokens=request.processed_tokens, + ) + + async def flush_entity_stats(self) -> None: + """Flush resolver statistics after, and never inside, the core commit.""" + + await self._entity_resolver.flush_pending_stats() + + async def _finalize_entity_graph( + self, + connection: Any, + request: WriteWindowRequest | DeltaWriteRequest, + ) -> CoreGraphWrite: + """Finalize missing canonical rows after ownership, inside the core UoW.""" + + graph = request.graph + if graph.entity_read_plan is None: + return graph + finalized = await self._entity_resolver.finalize_entity_read_plan( + connection, + request.bank_id, + graph.entity_read_plan, + entities_table=fq_table_explicit("entities", self._schema), + ) + fact_position = {fact.fact_key: str(index) for index, fact in enumerate(request.facts)} + planned_unit_keys = {occurrence.unit_key for occurrence in graph.entity_read_plan.occurrences} + if not planned_unit_keys <= set(fact_position): + unexpected = sorted(planned_unit_keys - set(fact_position)) + raise PersistenceContractError(f"entity read plan references unknown stable fact keys: {unexpected!r}") + entity_to_unit = tuple( + (fact_position[unit_key], local_index, event_date) + for unit_key, local_index, event_date in finalized.entity_to_unit + ) + unit_to_entity_ids = tuple( + (fact_position[unit_key], tuple(entity_ids)) for unit_key, entity_ids in finalized.unit_to_entity_ids + ) + return CoreGraphWrite( + resolved_entity_ids=tuple(finalized.resolved_entity_ids), + entity_to_unit=entity_to_unit, + unit_to_entity_ids=unit_to_entity_ids, + semantic_ann_links=graph.semantic_ann_links, + ) + + async def write_display_entity_links(self, request: RetainWriteRequest, phase3_payload: Any) -> None: + """Delegate best-effort Phase 3 display-link work after commit.""" + + if isinstance(request, MetadataOnlyWriteRequest): # pragma: no cover - metadata never requests phase 3 + raise PersistenceContractError("metadata-only writes cannot require display entity links") + log_buffer = request.log_buffer if isinstance(request, WriteWindowRequest) else list(request.log_buffer) + + await runtime.build_and_insert_entity_links( + self._pool, + self._entity_resolver, + request.bank_id, + phase3_payload, + self._config, + log_buffer, + ) + + @staticmethod + def _assert_complete_chunk_result( + request: WriteWindowRequest | DeltaWriteRequest, + chunk_ids_by_index: dict[int, str], + ) -> None: + expected_indices = {chunk.metadata.chunk_index for chunk in request.chunks} + actual_indices = set(chunk_ids_by_index) + if actual_indices != expected_indices: + missing = sorted(expected_indices - actual_indices) + unexpected = sorted(actual_indices - expected_indices) + raise PersistenceContractError( + f"chunk upsert returned an invalid key set (missing={missing}, unexpected={unexpected})" + ) + if any(not isinstance(chunk_id, str) or not chunk_id for chunk_id in chunk_ids_by_index.values()): + raise PersistenceContractError("chunk upsert returned an empty or non-string chunk_id") + + @staticmethod + def _assert_complete_fact_result( + request: WriteWindowRequest | DeltaWriteRequest, + unit_ids_by_content: Any, + ) -> None: + if not isinstance(unit_ids_by_content, list) or len(unit_ids_by_content) != len(request.contents): + raise PersistenceContractError("core write did not return one unit-id bucket per content item") + if any(not isinstance(bucket, list) for bucket in unit_ids_by_content): + raise PersistenceContractError("core write returned a non-list unit-id bucket") + flattened = [unit_id for bucket in unit_ids_by_content for unit_id in bucket] + if len(flattened) != len(request.facts): + raise PersistenceContractError( + f"core write returned {len(flattened)} unit IDs for {len(request.facts)} fact bindings" + ) + expected_per_content = [0] * len(request.contents) + for fact in request.facts: + expected_per_content[fact.processed.content_index] += 1 + actual_per_content = [len(bucket) for bucket in unit_ids_by_content] + if actual_per_content != expected_per_content: + raise PersistenceContractError( + "core write returned invalid per-content fact cardinality " + f"(expected={expected_per_content}, actual={actual_per_content})" + ) + if any(not isinstance(unit_id, str) or not unit_id for unit_id in flattened): + raise PersistenceContractError("core write returned an empty or non-string unit_id") + + @staticmethod + def _map_unit_ids_by_fact_key( + request: WriteWindowRequest | DeltaWriteRequest, + unit_ids_by_content: list[list[str]], + ) -> tuple[tuple[str, str], ...]: + mappings: list[tuple[str, str]] = [] + for content_index, unit_ids in enumerate(unit_ids_by_content): + facts = [fact for fact in request.facts if fact.processed.content_index == content_index] + mappings.extend((fact.fact_key, unit_id) for fact, unit_id in zip(facts, unit_ids, strict=True)) + return tuple(mappings) + + @staticmethod + def _thaw_retain_params(retain_params: FrozenObject | None) -> dict[str, Any] | None: + if retain_params is None: + return None + value = thaw_json(retain_params) + if not isinstance(value, dict): # pragma: no cover - request normalization guarantees an object + raise PersistenceContractError("retain_params did not thaw to an object") + return value diff --git a/core/dataplane/hms_api/engine/ingestion/projection/__init__.py b/core/dataplane/hms_api/engine/ingestion/projection/__init__.py new file mode 100644 index 0000000..09c2df7 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/projection/__init__.py @@ -0,0 +1,27 @@ +"""Projection of extracted facts into persistence-ready memory records.""" + +from .embeddings import ( + AsyncEmbeddingPort, + EmbeddingCardinalityError, + EmbeddingFailurePolicy, + build_embedding_text, + project_embeddings, +) +from .records import ( + MemoryRecord, + build_projection_manifest, + thaw_declared_entities, + to_processed_fact, +) + +__all__ = [ + "AsyncEmbeddingPort", + "EmbeddingCardinalityError", + "EmbeddingFailurePolicy", + "MemoryRecord", + "build_embedding_text", + "build_projection_manifest", + "project_embeddings", + "thaw_declared_entities", + "to_processed_fact", +] diff --git a/core/dataplane/hms_api/engine/ingestion/projection/embeddings.py b/core/dataplane/hms_api/engine/ingestion/projection/embeddings.py new file mode 100644 index 0000000..54a2f79 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/projection/embeddings.py @@ -0,0 +1,126 @@ +"""Async embedding projection with explicit whole-batch failure policy.""" + +from __future__ import annotations + +from collections.abc import Callable, Sequence +from datetime import datetime +from enum import StrEnum +from typing import Protocol + +from ..extraction.models import FactCandidate +from .records import MemoryRecord, build_projection_manifest + +EmbeddingVector = Sequence[float] | None +FormatDate = Callable[[datetime], str] + + +class AsyncEmbeddingPort(Protocol): + """Provider-independent asynchronous batch embedding boundary.""" + + async def embed_batch(self, texts: tuple[str, ...]) -> Sequence[EmbeddingVector]: ... + + +class EmbeddingFailurePolicy(StrEnum): + STORE_WITHOUT_EMBEDDING = "store_without_embedding" + RAISE = "raise" + + +class EmbeddingCardinalityError(RuntimeError): + """Raised when a backend violates one-output-per-fact alignment.""" + + +def build_embedding_text(candidate: FactCandidate, format_date: FormatDate) -> str: + """Build temporal and entity context for embedding.""" + + fact_date = candidate.occurred_start or candidate.mentioned_at + if fact_date is not None: + readable_date = format_date(fact_date) + if not isinstance(readable_date, str): + raise TypeError("format_date must return a string") + if candidate.occurred_end is not None and candidate.occurred_end != candidate.occurred_start: + readable_end = format_date(candidate.occurred_end) + if not isinstance(readable_end, str): + raise TypeError("format_date must return a string") + text = f"{candidate.text} (happened from {readable_date} to {readable_end})" + else: + text = f"{candidate.text} (happened in {readable_date})" + else: + text = candidate.text + + # Declared entities stay on their dedicated resolution path. Only names + # produced by an extraction strategy augment embeddings, matching the + # ProcessedFact manifest/entity semantics. + if candidate.entity_mentions: + text = f"{text} [{', '.join(candidate.entity_mentions)}]" + return text + + +def _freeze_vector(vector: EmbeddingVector, *, index: int) -> tuple[float, ...] | None: + if vector is None: + return None + if isinstance(vector, (str, bytes)): + raise TypeError(f"embedding[{index}] must be a numeric sequence or None") + try: + return tuple(float(value) for value in vector) + except (TypeError, ValueError) as exc: + raise TypeError(f"embedding[{index}] must be a numeric sequence or None") from exc + + +async def project_embeddings( + candidates: Sequence[FactCandidate], + *, + embedder: AsyncEmbeddingPort, + format_date: FormatDate, + embedding_model_version: str = "unknown", + extraction_version: str = "5w-v1", + failure_policy: EmbeddingFailurePolicy | str = EmbeddingFailurePolicy.STORE_WITHOUT_EMBEDDING, +) -> tuple[MemoryRecord, ...]: + """Project candidates 1:1, degrading only genuine backend exceptions. + + A backend exception defaults to the existing Retain behavior: every fact + remains persistable with a NULL embedding and ``embedding.ok=false``. + Returning the wrong number of vectors is instead a port contract violation + and always raises; no policy may silently truncate or positionally shift + facts. + """ + + candidate_batch = tuple(candidates) + if not candidate_batch: + return () + try: + policy = EmbeddingFailurePolicy(failure_policy) + except ValueError as exc: + choices = ", ".join(value.value for value in EmbeddingFailurePolicy) + raise ValueError(f"failure_policy must be one of: {choices}") from exc + + embedding_texts = tuple(build_embedding_text(candidate, format_date) for candidate in candidate_batch) + try: + raw_embeddings = await embedder.embed_batch(embedding_texts) + except Exception: + if policy is EmbeddingFailurePolicy.RAISE: + raise + embeddings: tuple[tuple[float, ...] | None, ...] = (None,) * len(candidate_batch) + else: + if raw_embeddings is None: + raise TypeError("embedding backend must return a sequence, got None") + raw_batch = tuple(raw_embeddings) + if len(raw_batch) != len(candidate_batch): + raise EmbeddingCardinalityError( + f"Embedding backend returned {len(raw_batch)} vectors for {len(candidate_batch)} facts; " + "expected exact 1:1 alignment" + ) + embeddings = tuple(_freeze_vector(vector, index=index) for index, vector in enumerate(raw_batch)) + + return tuple( + MemoryRecord.from_candidate( + candidate, + embedding=embedding, + projection=build_projection_manifest( + candidate, + embedding=embedding, + embedding_model_version=embedding_model_version, + extraction_version=extraction_version, + ), + ) + for candidate, embedding in zip(candidate_batch, embeddings, strict=True) + ) diff --git a/core/dataplane/hms_api/engine/ingestion/projection/records.py b/core/dataplane/hms_api/engine/ingestion/projection/records.py new file mode 100644 index 0000000..9d67917 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/projection/records.py @@ -0,0 +1,178 @@ +"""Immutable projected records and their persistence conversion.""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any + +from ..domain import FrozenJson, FrozenObject, freeze_json, thaw_json +from ..extraction.models import FactCandidate + + +@dataclass(frozen=True, slots=True) +class MemoryRecord(FactCandidate): + """One fact candidate enriched for persistence and current Recall.""" + + embedding: tuple[float, ...] | None + projection: FrozenJson + + def __post_init__(self) -> None: + FactCandidate.__post_init__(self) + if self.embedding is not None: + if not isinstance(self.embedding, tuple) or any(not isinstance(value, float) for value in self.embedding): + raise TypeError("embedding must be a tuple of floats or None") + if not isinstance(self.projection, FrozenObject): + raise TypeError("projection must be a frozen JSON object") + + @classmethod + def from_candidate( + cls, + candidate: FactCandidate, + *, + embedding: tuple[float, ...] | None, + projection: FrozenJson, + ) -> "MemoryRecord": + return cls( + fact_key=candidate.fact_key, + chunk_key=candidate.chunk_key, + source_index=candidate.source_index, + global_index=candidate.global_index, + extractor_local_index=candidate.extractor_local_index, + text=candidate.text, + fact_type=candidate.fact_type, + context=candidate.context, + where=candidate.where, + occurred_start=candidate.occurred_start, + occurred_end=candidate.occurred_end, + mentioned_at=candidate.mentioned_at, + metadata=candidate.metadata, + declared_entities=candidate.declared_entities, + tags=candidate.tags, + observation_scopes=candidate.observation_scopes, + entity_mentions=candidate.entity_mentions, + causal_relations=candidate.causal_relations, + embedding=embedding, + projection=projection, + ) + + +def build_projection_manifest( + candidate: FactCandidate, + *, + embedding: tuple[float, ...] | None, + embedding_model_version: str, + extraction_version: str, +) -> FrozenObject: + """Build the read-side manifest currently understood by Recall.""" + + if not isinstance(embedding_model_version, str) or not embedding_model_version: + raise ValueError("embedding_model_version must be a non-empty string") + if not isinstance(extraction_version, str) or not extraction_version: + raise ValueError("extraction_version must be a non-empty string") + + temporal_grade = ( + "resolved" if candidate.occurred_start is not None or candidate.mentioned_at is not None else "unresolved" + ) + manifest = freeze_json( + { + "embedding": {"v": embedding_model_version, "ok": embedding is not None}, + "tsvector": {"v": 1, "ok": True}, + "temporal": {"v": 1, "grade": temporal_grade}, + "entities": {"v": 1, "ok": bool(candidate.entity_mentions)}, + "extraction": {"v": extraction_version}, + } + ) + if not isinstance(manifest, FrozenObject): # pragma: no cover - guaranteed by the literal above + raise AssertionError("projection manifest must be an object") + return manifest + + +def thaw_declared_entities(record: FactCandidate) -> list[dict[str, Any]]: + """Return caller-declared entities for entity resolution.""" + + entities: list[dict[str, Any]] = [] + for index, frozen_entity in enumerate(record.declared_entities): + entity = thaw_json(frozen_entity) + if not isinstance(entity, dict): + raise TypeError(f"declared_entities[{index}] must thaw to an object") + entities.append(entity) + return entities + + +def to_processed_fact( + record: MemoryRecord, + *, + document_id: str, + chunk_id: str, + content_index: int, + fact_positions: Mapping[str, int] | None = None, +): + """Adapt a record to ``ProcessedFact`` with explicit identity bindings. + + Stable ``fact_key``/``chunk_key`` remain on the projected record. The adapter is + intentionally passed the effective document ID, persisted chunk ID, and + current content position so none of those mappings are inferred from + mutable list order. + """ + + if not isinstance(document_id, str) or not document_id: + raise ValueError("document_id must be a non-empty string") + if not isinstance(chunk_id, str) or not chunk_id: + raise ValueError("chunk_id must be a non-empty string") + if isinstance(content_index, bool) or not isinstance(content_index, int) or content_index < 0: + raise ValueError("content_index must be a non-negative integer") + + metadata = thaw_json(record.metadata) + projection = thaw_json(record.projection) + if not isinstance(metadata, dict): + raise TypeError("record metadata must thaw to an object") + if not isinstance(projection, dict): # pragma: no cover - MemoryRecord enforces this + raise TypeError("record projection must thaw to an object") + + from ...retain.types import CausalRelation, EntityRef, ProcessedFact + + if isinstance(record.observation_scopes, tuple): + observation_scopes = [list(scope) for scope in record.observation_scopes] + else: + observation_scopes = record.observation_scopes + + causal_relations = [] + if record.causal_relations: + if fact_positions is None: + raise ValueError("fact_positions is required when a record has causal relations") + source_position = fact_positions.get(record.fact_key) + if isinstance(source_position, bool) or not isinstance(source_position, int) or source_position < 0: + raise ValueError("fact_positions must contain a non-negative position for the source fact") + for relation in record.causal_relations: + target_position = fact_positions.get(relation.target_fact_key) + if isinstance(target_position, bool) or not isinstance(target_position, int) or target_position < 0: + raise ValueError("fact_positions must contain a non-negative position for every causal target") + if target_position >= source_position: + raise ValueError("causal targets must precede their source fact") + causal_relations.append( + CausalRelation( + relation_type=relation.relation_type, + target_fact_index=target_position, + ) + ) + + return ProcessedFact( + fact_text=record.text, + fact_type=record.fact_type, + embedding=list(record.embedding) if record.embedding is not None else None, + occurred_start=record.occurred_start, + occurred_end=record.occurred_end, + mentioned_at=record.mentioned_at, + context=record.context, + metadata=metadata, + where=record.where, + entities=[EntityRef(name=name) for name in record.entity_mentions], + causal_relations=causal_relations, + chunk_id=chunk_id, + document_id=document_id, + content_index=content_index, + tags=list(record.tags), + observation_scopes=observation_scopes, + projection=projection, + ) diff --git a/core/dataplane/hms_api/engine/ingestion/redaction.py b/core/dataplane/hms_api/engine/ingestion/redaction.py new file mode 100644 index 0000000..ba2777c --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/redaction.py @@ -0,0 +1,50 @@ +"""Identifier redaction helpers for trusted Retain ingestion.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +REDACTED_IDENTIFIER = "" + + +@dataclass(frozen=True, slots=True) +class IdentifierSanitizer: + """Redact request identifiers from logs and persisted error messages.""" + + enabled: bool = False + identifiers: tuple[str, ...] = () + + @classmethod + def from_values( + cls, + *, + enabled: bool, + values: tuple[Any, ...] = (), + ) -> IdentifierSanitizer: + identifiers = tuple(value for item in values if item is not None for value in (str(item),) if value) + return cls(enabled=enabled, identifiers=identifiers) + + def identifier(self, value: Any) -> str: + """Return one identifier in a log-safe form.""" + + if value is None: + return "None" + text = str(value) + if self.enabled and text: + return REDACTED_IDENTIFIER + return text + + def text(self, value: Any, *, extra_identifiers: tuple[Any, ...] = ()) -> str: + """Redact every known identifier from arbitrary diagnostic text.""" + + message = str(value) + if not self.enabled: + return message + extras = tuple(text for item in extra_identifiers if item is not None for text in (str(item),) if text) + for identifier in sorted((*self.identifiers, *extras), key=len, reverse=True): + message = message.replace(identifier, REDACTED_IDENTIFIER) + return message + + +__all__ = ["IdentifierSanitizer", "REDACTED_IDENTIFIER"] diff --git a/core/dataplane/hms_api/engine/ingestion/runtime.py b/core/dataplane/hms_api/engine/ingestion/runtime.py new file mode 100644 index 0000000..e579766 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/runtime.py @@ -0,0 +1,419 @@ +"""Runtime operations shared by the final Retain application service. + +This module contains the database-facing entity, graph, and post-commit +operations used by the single Retain pipeline. Planning remains in the pure +domain modules, while write orchestration remains in ``persistence``. +""" + +from __future__ import annotations + +import asyncio +import logging +import time +from typing import Any + +from ...worker.stage import set_stage +from ..db_utils import acquire_with_retry +from ..embedding_fingerprint import embedding_model_version +from ..memory_engine import count_tokens, fq_table +from ..retain import entity_processing, fact_storage, link_creation +from ..retain.link_utils import _bulk_insert_links, compute_semantic_links_ann +from ..retain.types import ( + EntityReadPlanPhase1Result, + Phase3Context, + ProcessedFact, + RetainContent, +) +from .adapters.oracle_semantic import compute_oracle_semantic_links_ann +from .persistence.backend import retain_backend_adapters + +logger = logging.getLogger(__name__) + +_ANN_CHUNK_SIZE = 1000 +_ANN_PARALLELISM = 4 +_ORACLE_IN_CHUNK_SIZE = 900 + + +async def pre_resolve_entities( + pool: Any, + entity_resolver: Any, + bank_id: str, + contents: list[RetainContent], + fact_keys: list[str], + processed_facts: list[ProcessedFact], + config: Any, + log_buffer: list[str], + *, + skip_semantic_ann: bool = False, +) -> EntityReadPlanPhase1Result: + """Build a read-only entity and semantic-link plan before the write UoW.""" + + set_stage("retain.phase1.resolve") + if len(fact_keys) != len(processed_facts): + raise ValueError("Retain requires one stable fact key per processed fact") + + user_entities_per_content = {index: content.entities for index, content in enumerate(contents) if content.entities} + placeholder_unit_ids = [str(index) for index in range(len(processed_facts))] + embeddings = [fact.embedding for fact in processed_facts] + backend_adapters = retain_backend_adapters(getattr(pool, "backend_type", "postgresql")) + backend_type = backend_adapters.backend_type + + async with acquire_with_retry(pool) as resolve_conn, backend_adapters.planning_snapshot(resolve_conn): + entity_read_plan = await entity_processing.plan_entities( + entity_resolver, + resolve_conn, + bank_id, + fact_keys, + processed_facts, + log_buffer, + user_entities_per_content=user_entities_per_content, + entity_labels=getattr(config, "entity_labels", None), + ) + semantic_ann_links = [] + if not skip_semantic_ann and all(embedding is not None for embedding in embeddings): + fact_types = [fact.fact_type for fact in processed_facts] + if backend_type == "oracle": + semantic_ann_links = await compute_oracle_semantic_links_ann( + resolve_conn, + bank_id, + placeholder_unit_ids, + embeddings, + fact_types=fact_types, + ) + else: + semantic_ann_links = await compute_semantic_links_ann( + resolve_conn, + bank_id, + placeholder_unit_ids, + embeddings, + fact_types=fact_types, + log_buffer=log_buffer, + read_only=True, + ) + elif not skip_semantic_ann: + log_buffer.append(" Semantic ANN precompute: skipped (missing embeddings)") + + return EntityReadPlanPhase1Result( + entity_read_plan=entity_read_plan, + semantic_ann_links=semantic_ann_links, + ) + + +def _remap_phase1_results( + resolved_entity_ids: list[str], + entity_to_unit: list[tuple[Any, ...]], + unit_to_entity_ids: dict[str, list[str]], + semantic_ann_links: list[tuple[Any, ...]], + actual_unit_ids: list[str], +) -> tuple[list[tuple[Any, ...]], dict[str, list[str]], list[tuple[Any, ...]]]: + """Replace read-plan placeholder unit IDs with committed database IDs.""" + + placeholder_to_actual = {str(index): actual_id for index, actual_id in enumerate(actual_unit_ids)} + remapped_entity_to_unit = [ + ( + placeholder_to_actual.get(unit_id, unit_id), + local_index, + fact_date, + ) + for unit_id, local_index, fact_date in entity_to_unit + ] + remapped_unit_to_entity_ids = { + placeholder_to_actual.get(placeholder_id, placeholder_id): entity_ids + for placeholder_id, entity_ids in unit_to_entity_ids.items() + } + remapped_semantic = [ + ( + placeholder_to_actual.get(link[0], link[0]), + link[1], + link[2], + link[3], + link[4], + ) + for link in semantic_ann_links + ] + return ( + remapped_entity_to_unit, + remapped_unit_to_entity_ids, + remapped_semantic, + ) + + +async def insert_facts_and_links( + conn: Any, + entity_resolver: Any, + bank_id: str, + contents: list[RetainContent], + processed_facts: list[ProcessedFact], + config: Any, + log_buffer: list[str], + resolved_entity_ids: list[str], + entity_to_unit: list[tuple[Any, ...]], + unit_to_entity_ids: dict[str, list[str]], + semantic_ann_links: list[tuple[Any, ...]], + *, + skip_semantic_links: bool = False, + outbox_callback: Any = None, + ops: Any = None, +) -> tuple[list[list[str]], Phase3Context]: + """Insert facts and retrieval-critical graph edges in one transaction.""" + + set_stage("retain.phase2.insert_facts") + step_start = time.time() + unit_ids = await fact_storage.insert_facts_batch( + conn, + bank_id, + processed_facts, + ops=ops, + ) + log_buffer.append(f" Insert facts: {len(unit_ids)} units in {time.time() - step_start:.3f}s") + + phase3_context = Phase3Context() + if unit_ids: + step_start = time.time() + ( + remapped_entity_to_unit, + remapped_unit_to_entity_ids, + semantic_ann_links, + ) = _remap_phase1_results( + resolved_entity_ids, + entity_to_unit, + unit_to_entity_ids, + semantic_ann_links or [], + unit_ids, + ) + unit_entity_pairs = [ + (unit_id, resolved_entity_ids[index], fact_date) + for index, (unit_id, _local_index, fact_date) in enumerate(remapped_entity_to_unit) + ] + await entity_resolver.link_units_to_entities_batch( + unit_entity_pairs, + conn=conn, + ) + log_buffer.append(f" Insert unit_entities: {len(unit_entity_pairs)} pairs in {time.time() - step_start:.3f}s") + phase3_context = Phase3Context( + unit_ids=unit_ids, + resolved_entity_ids=resolved_entity_ids, + entity_to_unit=remapped_entity_to_unit, + unit_to_entity_ids=remapped_unit_to_entity_ids, + ) + + step_start = time.time() + temporal_link_count = await link_creation.create_temporal_links_batch( + conn, + bank_id, + unit_ids, + ops=ops, + write_temporal_links=getattr(config, "write_temporal_links", True), + ) + log_buffer.append(f" Temporal links: {temporal_link_count} links in {time.time() - step_start:.3f}s") + + if not getattr(config, "write_semantic_links", True): + log_buffer.append(" Semantic links: skipped (mode=ann)") + elif skip_semantic_links: + log_buffer.append(" Semantic links: skipped (deferred to final ANN pass)") + else: + step_start = time.time() + embeddings_for_links = [fact.embedding for fact in processed_facts] + if not all(embedding is not None for embedding in embeddings_for_links): + log_buffer.append(" Semantic links: skipped (missing embeddings)") + else: + semantic_link_count = await link_creation.create_semantic_links_batch( + conn, + bank_id, + unit_ids, + embeddings_for_links, + pre_computed_ann_links=semantic_ann_links, + ops=ops, + write_semantic_links=getattr( + config, + "write_semantic_links", + True, + ), + ) + log_buffer.append(f" Semantic links: {semantic_link_count} links in {time.time() - step_start:.3f}s") + + step_start = time.time() + causal_link_count = await link_creation.create_causal_links_batch( + conn, + bank_id, + unit_ids, + processed_facts, + ops=ops, + ) + log_buffer.append(f" Causal links: {causal_link_count} links in {time.time() - step_start:.3f}s") + + result_unit_ids: list[list[str]] = [[] for _ in contents] + for processed_fact, unit_id in zip(processed_facts, unit_ids, strict=True): + content_index = processed_fact.content_index + if content_index < 0 or content_index >= len(contents): + raise ValueError(f"Fact content index {content_index} is outside the request") + result_unit_ids[content_index].append(unit_id) + + if outbox_callback: + await outbox_callback(conn) + return result_unit_ids, phase3_context + + +async def build_and_insert_entity_links( + pool: Any, + entity_resolver: Any, + bank_id: str, + phase3_context: Phase3Context, + config: Any, + log_buffer: list[str], +) -> None: + """Build visualization-only entity links after the core transaction.""" + + set_stage("retain.phase3.entity_links") + if not getattr(config, "write_entity_links", True): + log_buffer.append(" Entity links (viz): skipped (write_entity_links=false)") + return + if not phase3_context.unit_ids or not phase3_context.resolved_entity_ids: + return + + async with acquire_with_retry(pool) as conn: + step_start = time.time() + entity_links = await entity_processing.build_entity_links( + entity_resolver, + conn, + bank_id, + phase3_context.unit_ids, + phase3_context.resolved_entity_ids, + phase3_context.entity_to_unit, + phase3_context.unit_to_entity_ids, + log_buffer, + skip_unit_entities_insert=True, + ops=pool.ops, + ) + if entity_links: + await entity_processing.insert_entity_links_batch( + conn, + entity_links, + bank_id, + ops=pool.ops, + ) + log_buffer.append(f" Entity links (viz): {len(entity_links)} links in {time.time() - step_start:.3f}s") + + +async def run_final_semantic_ann( + pool: Any, + bank_id: str, + unit_ids: list[str], + config: Any, + log_buffer: list[str], +) -> None: + """Create semantic links for all committed units in bounded ANN batches.""" + + if not getattr(config, "write_semantic_links", True): + log_buffer.append("[streaming] Final ANN: semantic links skipped") + return + if not unit_ids: + return + + backend_type = retain_backend_adapters(getattr(pool, "backend_type", "postgresql")).backend_type + load_start = time.time() + async with acquire_with_retry(pool) as conn: + if backend_type == "oracle": + rows = [] + for start in range(0, len(unit_ids), _ORACLE_IN_CHUNK_SIZE): + rows.extend( + await conn.fetch( + f""" + SELECT id, embedding, fact_type + FROM {fq_table("memory_units")} + WHERE bank_id = $1 AND id = ANY($2::uuid[]) + ORDER BY id + """, + bank_id, + unit_ids[start : start + _ORACLE_IN_CHUNK_SIZE], + ) + ) + else: + rows = await conn.fetch( + f""" + SELECT id::text, embedding::text, fact_type + FROM {fq_table("memory_units")} + WHERE bank_id = $1 AND id = ANY($2::uuid[]) + ORDER BY id + """, + bank_id, + unit_ids, + ) + if not rows: + log_buffer.append("[streaming] Final ANN: no committed units found") + return + + unit_map = {str(row["id"]): (row["embedding"], row["fact_type"]) for row in rows} + ann_unit_ids: list[str] = [] + ann_embeddings: list[Any] = [] + ann_fact_types: list[str] = [] + for unit_id in unit_ids: + stored = unit_map.get(unit_id) + if stored is not None and stored[0] is not None: + ann_unit_ids.append(unit_id) + ann_embeddings.append(stored[0]) + ann_fact_types.append(stored[1]) + + log_buffer.append( + f"[streaming] Final ANN: loaded {len(ann_unit_ids)} units with embeddings in {time.time() - load_start:.3f}s" + ) + if not ann_unit_ids: + return + + chunk_count = (len(ann_unit_ids) + _ANN_CHUNK_SIZE - 1) // _ANN_CHUNK_SIZE + semaphore = asyncio.Semaphore(_ANN_PARALLELISM) + link_counts = [0] * chunk_count + + async def process_chunk(chunk_index: int) -> None: + start = chunk_index * _ANN_CHUNK_SIZE + end = min(start + _ANN_CHUNK_SIZE, len(ann_unit_ids)) + async with semaphore: + started_at = time.time() + async with acquire_with_retry(pool) as conn: + if backend_type == "oracle": + ann_links = await compute_oracle_semantic_links_ann( + conn, + bank_id, + ann_unit_ids[start:end], + ann_embeddings[start:end], + fact_types=ann_fact_types[start:end], + top_k=20, + ) + else: + ann_links = await compute_semantic_links_ann( + conn, + bank_id, + ann_unit_ids[start:end], + ann_embeddings[start:end], + fact_types=ann_fact_types[start:end], + top_k=20, + log_buffer=log_buffer, + ) + if ann_links: + await _bulk_insert_links( + conn, + ann_links, + bank_id=bank_id, + ops=pool.ops, + ) + link_counts[chunk_index] = len(ann_links) + logger.info( + "Final ANN chunk %d/%d: %d links in %.3fs", + chunk_index + 1, + chunk_count, + len(ann_links), + time.time() - started_at, + ) + + await asyncio.gather(*(process_chunk(index) for index in range(chunk_count))) + log_buffer.append(f"[streaming] Final ANN: {sum(link_counts)} total semantic links") + + +__all__ = [ + "build_and_insert_entity_links", + "count_tokens", + "embedding_model_version", + "insert_facts_and_links", + "pre_resolve_entities", + "run_final_semantic_ann", +] diff --git a/core/dataplane/hms_api/engine/ingestion/service.py b/core/dataplane/hms_api/engine/ingestion/service.py new file mode 100644 index 0000000..1f8e921 --- /dev/null +++ b/core/dataplane/hms_api/engine/ingestion/service.py @@ -0,0 +1,1591 @@ +"""Retain application service for database-backed ingestion. + +All documents are planned from one read-only snapshot before any bank, +checkpoint, document, chunk, fact, link, or outbox write begins. Each document +then executes hash-guarded writes selected from the pure change plan: +fresh/existing full replacement (possibly split into bounded windows), delta, +or metadata-only. Stale plans perform no write at all. +""" + +from __future__ import annotations + +import asyncio +import logging +import sys +import time +import uuid +from collections.abc import Iterator, Sequence +from contextlib import asynccontextmanager +from dataclasses import dataclass, replace +from datetime import UTC, datetime +from typing import Any + +from ..db_utils import acquire_with_retry +from ..embedding_fingerprint import EmbeddingFingerprintError, ensure_bank_embedding_fingerprint +from ..response_models import TokenUsage +from ..retain import bank_utils +from .adapters.embedding_model import EmbeddingModelAdapter +from .adapters.postgres_fresh_ownership import FreshDocumentOwnershipConflict +from .adapters.storage_records import ( + chunks_to_storage, + compute_document_hash, + content_positions, + content_to_storage, + record_to_extracted_fact, + record_to_processed_fact, + retain_document_metadata, +) +from .change_detection import detect_document_change +from .chunking import build_chunk_plans +from .contracts import RetainExecutionContext, RetainInvocation, RetainOutcome +from .document_planner import plan_documents, prepend_existing_document +from .domain import ( + ChunkPlan, + ChunkPolicy, + DocumentChangeKind, + DocumentChangePlan, + DocumentIntent, + ExistingChunkFingerprint, + UpdateMode, +) +from .execution import merge_window_unit_ids, plan_full_write_windows +from .extraction import ( + ExtractionMode, + ExtractionPolicy, + FactExtractorAdapter, + build_prechunked_extraction_layout, + extract_passthrough, +) +from .normalization import normalize_contents +from .persistence.backend import RetainBackendAdapters, retain_backend_adapters +from .persistence.models import CommittedUnitBinding, ExistingDocument, OperationCheckpoint +from .persistence.unit_of_work import ( + ChunkWrite, + CoreGraphWrite, + DeltaWriteRequest, + ExistingChunkWrite, + FactWrite, + FirstFullWriteWindow, + LaterFullWriteWindow, + MetadataOnlyWriteRequest, + OwnershipDisposition, + PostCommitStatus, + RetainUnitOfWork, + RetainWriteRequest, + WriteWindowRequest, +) +from .persistence.writer import PersistenceWriter +from .projection import EmbeddingFailurePolicy, MemoryRecord, project_embeddings +from .redaction import IdentifierSanitizer +from .runtime import ( + count_tokens, + embedding_model_version, + pre_resolve_entities, + run_final_semantic_ann, +) + +logger = logging.getLogger(__name__) + +_INFLIGHT_CONTENT_HASH_PREFIX = "retain-inflight:" +_DEFAULT_PROJECTION_PIPELINE_CONCURRENCY = 4 + + +class RetainError(RuntimeError): + """Base class for Retain application-boundary failures.""" + + +class RetainUnsupportedError(RetainError): + """The Retain pipeline does not safely implement this request.""" + + +class RetainExtractionModeUnsupportedError(RetainUnsupportedError): + """The configured extraction strategy is not a supported Retain mode.""" + + +class RetainDatabaseUnsupportedError(RetainUnsupportedError): + """The configured database backend is not implemented by Retain.""" + + +class RetainResultMappingError(RetainError): + """A persistence result cannot be mapped one-to-one to request inputs.""" + + +class RetainPublicationAborted(RetainError): + """The request no longer owns a document state that it can publish. + + The message must remain identifier-free because asynchronous workers may + persist it as part of a terminal operation result. + """ + + +class RetainOwnershipLostError(RetainPublicationAborted): + """The document changed after planning and this attempt must be retried.""" + + +class RetainCheckpointRecoveryError(RetainError): + """A durable operation checkpoint cannot be reconciled with Retain rows.""" + + +def _request_sanitizer( + invocation: RetainInvocation, + *identifiers: Any, +) -> IdentifierSanitizer: + return IdentifierSanitizer.from_values( + enabled=invocation.sanitize_log_identifiers, + values=( + invocation.bank_id, + invocation.operation_id, + *identifiers, + ), + ) + + +def _log_identifier(invocation: RetainInvocation, value: Any) -> str: + return _request_sanitizer(invocation, value).identifier(value) + + +def _log_warning( + invocation: RetainInvocation, + message: str, + *args: Any, + identifiers: tuple[Any, ...] = (), + exc_info: bool = False, +) -> None: + """Log a warning without exposing trusted request identifiers.""" + + sanitizer = _request_sanitizer(invocation, *identifiers) + if not sanitizer.enabled: + logger.warning(message, *args, exc_info=exc_info) + return + + safe_args = tuple(sanitizer.text(value) for value in args) + if exc_info: + exception = sys.exc_info()[1] + if exception is not None: + message = f"{message}: %s" + safe_args = (*safe_args, sanitizer.text(exception)) + logger.warning(message, *safe_args) + + +@dataclass(frozen=True, slots=True) +class _DocumentExecutionPlan: + intent: DocumentIntent + chunks: tuple[ChunkPlan, ...] + combined_content: str + existing: ExistingDocument | None + existing_chunks: tuple[ExistingChunkFingerprint, ...] + change: DocumentChangePlan + recovered_unit_bindings: tuple[CommittedUnitBinding, ...] | None = None + recovered_chunk_sources: tuple[tuple[int, int | None], ...] | None = None + final_ann_pending: bool = False + + @property + def recovered_unit_ids(self) -> tuple[str, ...] | None: + if self.recovered_unit_bindings is None: + return None + return tuple(binding.unit_id for binding in self.recovered_unit_bindings) + + +@dataclass(frozen=True, slots=True) +class _FactPayload: + storage_contents: tuple[Any, ...] + chunks: tuple[ChunkWrite, ...] + facts: tuple[FactWrite, ...] + graph: CoreGraphWrite + + +@dataclass(frozen=True, slots=True) +class _DocumentOutcome: + unit_ids_by_content: tuple[tuple[str, ...], ...] + usage: TokenUsage + processed_tokens: int | None + + +@dataclass(frozen=True, slots=True) +class _ProjectedChunkOutcome: + records: tuple[MemoryRecord, ...] + usage: TokenUsage + extraction_seconds: float + embedding_seconds: float + + +def _projection_pipeline_concurrency(config: Any) -> int: + """Resolve a finite positive producer width without trusting bool-as-int.""" + + configured = getattr(config, "retain_llm_max_concurrent", None) + if configured is None: + configured = getattr(config, "llm_max_concurrent", None) + if configured is None: + return _DEFAULT_PROJECTION_PIPELINE_CONCURRENCY + if isinstance(configured, bool) or not isinstance(configured, int) or configured <= 0: + logger.warning( + "Invalid Retain projection pipeline concurrency %r; using fallback=%d", + configured, + _DEFAULT_PROJECTION_PIPELINE_CONCURRENCY, + ) + return _DEFAULT_PROJECTION_PIPELINE_CONCURRENCY + return configured + + +def _log_window_producer_summary( + *, + path: str, + chunks: int, + facts: int, + wall_seconds: float, + extraction_seconds: Sequence[float], + embedding_seconds: Sequence[float], + configured_concurrency: int, +) -> None: + """Emit payload-free timing telemetry for one extraction window.""" + + logger.info( + "Retain window producer: path=%s chunks=%d facts=%d wall_seconds=%.3f " + "extraction_sum_seconds=%.3f extraction_max_seconds=%.3f " + "embedding_sum_seconds=%.3f embedding_max_seconds=%.3f configured_concurrency=%d", + path, + chunks, + facts, + wall_seconds, + sum(extraction_seconds), + max(extraction_seconds, default=0.0), + sum(embedding_seconds), + max(embedding_seconds, default=0.0), + configured_concurrency, + ) + + +def _backend_type(execution: RetainExecutionContext) -> str: + backend_type = getattr(execution.pool, "backend_type", None) + if backend_type is None: + backend_type = getattr(execution.resolved_config, "database_backend", "postgresql") + if not isinstance(backend_type, str): + raise RetainDatabaseUnsupportedError("Retain could not determine the database backend") + return backend_type.lower() + + +def _require_supported_route(execution: RetainExecutionContext) -> None: + backend_type = _backend_type(execution) + if backend_type not in {"postgresql", "oracle"}: + raise RetainDatabaseUnsupportedError( + f"Retain supports PostgreSQL and Oracle; configured backend is {backend_type!r}" + ) + mode = getattr(execution.resolved_config, "retain_extraction_mode", None) + try: + ExtractionMode(mode) + except (TypeError, ValueError) as exc: + raise RetainExtractionModeUnsupportedError(f"Retain does not support retain_extraction_mode={mode!r}.") from exc + + +def _recovered_id_factory(recovered_ids: Sequence[str], explicit_ids: set[str]) -> Iterator[str]: + for document_id in recovered_ids: + if document_id not in explicit_ids: + yield document_id + while True: + yield str(uuid.uuid4()) + + +def _append_recovery_chunk_sources( + durable_chunks: Sequence[ExistingChunkFingerprint], + submitted_chunks: Sequence[ChunkPlan], + *, + document_id: str, +) -> tuple[tuple[int, int | None], ...]: + """Map a completed append's durable chunk suffix to submitted sources. + + The document row now contains the *post-append* combined text, so + re-splitting it as one synthetic item cannot reliably reproduce the chunk + boundaries used when the operation committed. Append planning always + places submitted chunks after the synthetic pre-append chunks. Recover + that stable suffix from the durable chunk table and verify every hash + before exposing any operation-local unit ID. + """ + + indices = tuple(chunk.chunk_index for chunk in durable_chunks) + if indices != tuple(range(len(durable_chunks))): + raise RetainCheckpointRecoveryError( + f"Committed append document {document_id!r} has non-contiguous chunk indices {indices!r}" + ) + if len(submitted_chunks) > len(durable_chunks): + raise RetainCheckpointRecoveryError( + f"Committed append document {document_id!r} has fewer durable chunks than the retry payload" + ) + + submitted_count = len(submitted_chunks) + suffix_start = len(durable_chunks) - submitted_count + if submitted_count: + durable_suffix = durable_chunks[suffix_start:] + for durable, submitted in zip(durable_suffix, submitted_chunks, strict=True): + if durable.content_hash != submitted.content_hash: + raise RetainCheckpointRecoveryError( + f"Committed append document {document_id!r} does not match the retry payload " + f"at durable chunk index {durable.chunk_index}" + ) + + sources: list[tuple[int, int | None]] = [] + for position, durable in enumerate(durable_chunks): + source_index = None + if position >= suffix_start: + source_index = submitted_chunks[position - suffix_start].source_index + sources.append((durable.chunk_index, source_index)) + return tuple(sources) + + +def _graph_write(phase1: Any) -> CoreGraphWrite: + if phase1 is None: + return CoreGraphWrite() + entity_read_plan = getattr(phase1, "entity_read_plan", None) + if entity_read_plan is not None: + return CoreGraphWrite( + semantic_ann_links=tuple(tuple(link) for link in phase1.semantic_ann_links), + entity_read_plan=entity_read_plan, + ) + return CoreGraphWrite( + resolved_entity_ids=tuple(phase1.entities.resolved_entity_ids), + entity_to_unit=tuple(tuple(binding) for binding in phase1.entities.entity_to_unit), + unit_to_entity_ids=tuple( + (unit_id, tuple(entity_ids)) for unit_id, entity_ids in phase1.entities.unit_to_entity_ids.items() + ), + semantic_ann_links=tuple(tuple(link) for link in phase1.semantic_ann_links), + ) + + +def _merge_processed_tokens(current: int | None, document: int | None) -> int | None: + """Preserve the public token-accounting contract across documents.""" + + if current is None or document is None: + return None + return current + document + + +@asynccontextmanager +async def _database_budget(semaphore: Any): + if semaphore is None: + yield + return + async with semaphore: + yield + + +def _backend_adapters(execution: RetainExecutionContext) -> RetainBackendAdapters: + """Select the persistence contracts for the request's configured backend.""" + + try: + return retain_backend_adapters(_backend_type(execution)) + except ValueError as exc: # Defensive parity with the route guard. + raise RetainDatabaseUnsupportedError(str(exc)) from exc + + +class RetainPipelineService: + """Hash-guarded chunks-mode Retain pipeline.""" + + async def retain( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + ) -> RetainOutcome: + """Run the pipeline with one schema shared by every Retain stage. + + Planning receives the schema explicitly, while low-level storage + helpers resolve table names through ``memory_engine``'s task-local + schema context. Scoping both to the same value prevents a request from + reading one tenant and writing another. + """ + + from ..memory_engine import _current_schema, get_current_schema + + effective_schema = execution.schema or get_current_schema() + scoped_execution = replace(execution, schema=effective_schema) + schema_token = _current_schema.set(effective_schema) + try: + return await self._retain_in_schema(invocation, scoped_execution) + finally: + _current_schema.reset(schema_token) + + async def _retain_in_schema( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + ) -> RetainOutcome: + _require_supported_route(execution) + request_started_at = datetime.now(UTC) + normalized = normalize_contents(invocation.raw_contents, document_tags=invocation.document_tags) + if not normalized: + return RetainOutcome([], TokenUsage(), None) + + checkpoint = await self._recover_checkpoint(invocation, execution) + # Core checkpoints atomically carry document IDs. The secondary field + # recovers checkpoints whose early document-ID write did not complete. + recovered_ids = checkpoint.document_ids or checkpoint.core_committed_document_ids + explicit_ids = {item.document_id for item in normalized if item.document_id is not None} + generated_ids = _recovered_id_factory(recovered_ids, explicit_ids) + intents = plan_documents( + normalized, + batch_document_id=invocation.batch_document_id, + recovered_document_id=recovered_ids[0] if not explicit_ids and recovered_ids else None, + id_factory=lambda: next(generated_ids), + ) + policy = ChunkPolicy( + version="retain-chunker-v1", + max_chars=getattr(execution.resolved_config, "retain_chunk_size", 3000), + conversation_mode=True, + overlap=0, + ) + + # This is the all-document read barrier. No bank auto-create, + # checkpoint update, semantic write, or outbox call occurs before it. + plans = await self._preflight_documents( + invocation, + execution, + intents, + policy, + checkpoint=checkpoint, + request_started_at=request_started_at, + ) + + unpublished_stale_plans = tuple( + plan + for plan in plans + if plan.change.kind is DocumentChangeKind.STALE_SKIP and plan.recovered_unit_ids is None + ) + if unpublished_stale_plans and invocation.outbox_callback is not None: + raise RetainPublicationAborted("Retain was superseded before publication") + + bank_profile = await bank_utils.get_bank_profile(execution.pool, invocation.bank_id) + agent_name = bank_profile["name"] + try: + async with acquire_with_retry(execution.pool) as fingerprint_connection: + await ensure_bank_embedding_fingerprint( + fingerprint_connection, + invocation.bank_id, + execution.embeddings_model, + policy=getattr( + execution.resolved_config, + "embedding_fingerprint_policy", + "strict", + ), + for_write=False, + legacy_attestation=getattr( + execution.resolved_config, + "embedding_fingerprint_legacy_attestation", + None, + ), + log_sanitizer=_request_sanitizer(invocation), + ) + except EmbeddingFingerprintError as exc: + sanitizer = _request_sanitizer(invocation) + if sanitizer.enabled: + exc.args = (sanitizer.text(exc),) + raise + await self._record_document_ids(invocation, execution, intents) + + unit_ids_by_input: list[list[str]] = [[] for _ in invocation.raw_contents] + total_usage = TokenUsage() + total_processed_tokens: int | None = 0 + commit_positions = [ + position + for position, plan in enumerate(plans) + if plan.change.kind is not DocumentChangeKind.STALE_SKIP and plan.recovered_unit_ids is None + ] + last_commit_position = commit_positions[-1] if commit_positions else None + + for position, plan in enumerate(plans): + if plan.recovered_unit_ids is not None: + document_outcome = await self._resume_committed_document( + invocation, + execution, + plan, + ) + elif plan.change.kind is DocumentChangeKind.STALE_SKIP: + document_outcome = _DocumentOutcome( + unit_ids_by_content=tuple(() for _ in plan.intent.items), + usage=TokenUsage(), + processed_tokens=0, + ) + else: + document_outcome = await self._execute_document( + invocation, + execution, + plan, + agent_name=agent_name, + outbox_callback=(invocation.outbox_callback if position == last_commit_position else None), + ) + total_usage = total_usage + document_outcome.usage + total_processed_tokens = _merge_processed_tokens( + total_processed_tokens, + document_outcome.processed_tokens, + ) + self._merge_document_result( + unit_ids_by_input, + plan.intent, + document_outcome.unit_ids_by_content, + ) + + return RetainOutcome(unit_ids_by_input, total_usage, total_processed_tokens) + + async def _recover_checkpoint( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + ) -> OperationCheckpoint: + if invocation.operation_id is None: + return OperationCheckpoint() + try: + async with acquire_with_retry(execution.pool) as connection: + return ( + await _backend_adapters(execution) + .checkpoint_store( + connection, + schema=execution.schema, + ) + .recover(invocation.operation_id) + ) + except Exception as exc: + # Once an operation ID exists, an unreadable checkpoint is an + # unknown commit state. Treating it as empty can allocate a new + # document ID and duplicate an already-committed retry. + operation_id = _log_identifier(invocation, invocation.operation_id) + raise RetainCheckpointRecoveryError( + f"Retain could not safely recover operation checkpoint {operation_id}" + ) from exc + + async def _record_document_ids( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + intents: Sequence[DocumentIntent], + ) -> None: + if invocation.operation_id is None: + return + try: + async with acquire_with_retry(execution.pool) as connection: + store = _backend_adapters(execution).checkpoint_store( + connection, + schema=execution.schema, + ) + for intent in intents: + await store.record_document_id(invocation.operation_id, intent.document_id) + except Exception: + # This early checkpoint is best effort. The core-commit checkpoint + # below is strict because it shares the memory write transaction. + _log_warning( + invocation, + "Retain could not record document IDs for operation %s", + _log_identifier(invocation, invocation.operation_id), + identifiers=tuple(intent.document_id for intent in intents), + exc_info=True, + ) + + async def _preflight_documents( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + intents: Sequence[DocumentIntent], + policy: ChunkPolicy, + *, + checkpoint: OperationCheckpoint, + request_started_at: datetime, + ) -> tuple[_DocumentExecutionPlan, ...]: + plans: list[_DocumentExecutionPlan] = [] + adapters = _backend_adapters(execution) + async with acquire_with_retry(execution.pool) as connection, adapters.planning_snapshot(connection): + repository = adapters.planning_repository(connection, schema=execution.schema) + for submitted_intent in intents: + existing = await repository.load_document( + invocation.bank_id, + submitted_intent.document_id, + ) + if checkpoint.is_core_committed(submitted_intent.document_id): + if existing is None: + raise RetainCheckpointRecoveryError( + "Operation checkpoint says document " + f"{submitted_intent.document_id!r} committed, but no document row exists" + ) + expected_unit_ids = checkpoint.unit_ids_for_document(submitted_intent.document_id) + intent = submitted_intent + recovered_chunk_sources = None + if submitted_intent.update_mode is UpdateMode.APPEND and expected_unit_ids is not None: + chunks = build_chunk_plans( + submitted_intent.document_id, + submitted_intent.items, + policy, + ) + durable_chunks = await repository.load_chunks( + invocation.bank_id, + submitted_intent.document_id, + ) + recovered_chunk_sources = _append_recovery_chunk_sources( + durable_chunks, + chunks, + document_id=submitted_intent.document_id, + ) + else: + if submitted_intent.update_mode is UpdateMode.APPEND and existing.original_text: + intent = prepend_existing_document( + submitted_intent, + existing.original_text, + ) + chunks = build_chunk_plans( + intent.document_id, + intent.items, + policy, + ) + try: + unit_bindings = await repository.load_document_unit_bindings( + invocation.bank_id, + submitted_intent.document_id, + expected_unit_ids=expected_unit_ids, + ) + except (TypeError, ValueError) as exc: + raise RetainCheckpointRecoveryError( + "Operation checkpoint unit IDs cannot be reconciled for document " + f"{submitted_intent.document_id!r}" + ) from exc + combined_content = "\n".join(item.content for item in intent.items) + fallback_checkpoint = ( + checkpoint.unscoped_facts_committed and not checkpoint.core_committed_document_ids + ) + plans.append( + _DocumentExecutionPlan( + intent=intent, + chunks=chunks, + combined_content=combined_content, + existing=existing, + existing_chunks=(), + change=DocumentChangePlan( + kind=DocumentChangeKind.METADATA_ONLY, + reason="operation core commit recovered", + ), + recovered_unit_bindings=unit_bindings, + recovered_chunk_sources=recovered_chunk_sources, + final_ann_pending=( + submitted_intent.document_id in checkpoint.final_ann_pending_document_ids + or fallback_checkpoint + ), + ) + ) + continue + intent = submitted_intent + if ( + submitted_intent.update_mode is UpdateMode.APPEND + and existing is not None + and existing.original_text + ): + intent = prepend_existing_document(submitted_intent, existing.original_text) + + combined_content = "\n".join(item.content for item in intent.items) + chunks = build_chunk_plans(intent.document_id, intent.items, policy) + existing_chunks = ( + await repository.load_chunks(invocation.bank_id, intent.document_id) if existing is not None else () + ) + change = detect_document_change( + chunks, + existing_chunks, + document_exists=existing is not None, + existing_document_content_hash=existing.content_hash if existing is not None else None, + new_document_content_hash=compute_document_hash(combined_content), + updated_at=existing.updated_at if existing is not None else None, + request_started_at=request_started_at, + policy_compatible=not ( + existing is not None + and isinstance(existing.content_hash, str) + and existing.content_hash.startswith(_INFLIGHT_CONTENT_HASH_PREFIX) + ), + ) + plans.append( + _DocumentExecutionPlan( + intent=intent, + chunks=chunks, + combined_content=combined_content, + existing=existing, + existing_chunks=existing_chunks, + change=change, + ) + ) + return tuple(plans) + + async def _execute_document( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + plan: _DocumentExecutionPlan, + *, + agent_name: str, + outbox_callback: Any, + ) -> _DocumentOutcome: + if plan.change.kind is DocumentChangeKind.FULL: + return await self._execute_full_document_windows( + invocation, + execution, + plan, + agent_name=agent_name, + outbox_callback=outbox_callback, + ) + + selected_chunks = self._chunks_for_extraction(plan) + processed_tokens = self._processed_chunk_tokens(plan, selected_chunks) + records, extraction_usage = await self._extract_and_project_selected_chunks( + invocation, + execution, + plan, + selected_chunks, + agent_name=agent_name, + ) + checkpoint_callback = self._compose_checkpoint_callback( + invocation, + execution, + plan, + expected_unit_ids_count=len(records), + ) + + ownership = _backend_adapters(execution).document_ownership( + schema=execution.schema, + fresh=plan.existing is None, + ) + adapter = PersistenceWriter( + pool=execution.pool, + embeddings_model=execution.embeddings_model, + entity_resolver=execution.entity_resolver, + config=execution.resolved_config, + ownership=ownership, + operation_activity=_backend_adapters(execution).operation_activity_fence( + invocation.operation_id, + schema=execution.schema, + ), + schema=execution.schema, + sanitize_log_identifiers=invocation.sanitize_log_identifiers, + ) + unit_of_work = RetainUnitOfWork( + connection_scope=lambda: acquire_with_retry(execution.pool), + adapter=adapter, + ) + + try: + async with _database_budget(execution.db_semaphore): + request = await self._build_write_request( + invocation, + execution, + plan, + selected_chunks, + records, + processed_tokens=processed_tokens, + checkpoint_callback=checkpoint_callback, + outbox_callback=outbox_callback, + ) + result = await unit_of_work.execute(request) + except FreshDocumentOwnershipConflict as exc: + execution.entity_resolver.discard_pending_stats() + raise RetainOwnershipLostError( + "Retain lost fresh-document ownership for " + f"{_log_identifier(invocation, plan.intent.document_id)!r}; retry" + ) from exc + except BaseException: + # Phase 1 accumulates resolver stats before the core transaction. + # Never retain those task-local entries after rollback, commit + # failure, cancellation, or an adapter contract error. + execution.entity_resolver.discard_pending_stats() + raise + + if result.core.ownership is OwnershipDisposition.LOST: + # Phase 1 may have accumulated resolver statistics even though the + # hash guard prevented every core write. They must never leak into + # a later successful request's post-commit flush. + execution.entity_resolver.discard_pending_stats() + raise RetainOwnershipLostError( + f"Retain lost document ownership for {_log_identifier(invocation, plan.intent.document_id)!r}; retry" + ) + self._log_post_commit_failure(invocation, plan, result) + return _DocumentOutcome( + unit_ids_by_content=result.core.unit_ids_by_content, + usage=extraction_usage, + processed_tokens=( + result.core.processed_tokens if result.core.processed_tokens is not None else processed_tokens + ), + ) + + async def _execute_full_document_windows( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + plan: _DocumentExecutionPlan, + *, + agent_name: str, + outbox_callback: Any, + ) -> _DocumentOutcome: + """Execute ordered, memory-bounded FULL windows with one finalizer.""" + + windows = plan_full_write_windows( + plan.chunks, + getattr(execution.resolved_config, "retain_chunk_batch_size", 100), + ) + final_content_hash = compute_document_hash(plan.combined_content) + inflight_content_hash = f"{_INFLIGHT_CONTENT_HASH_PREFIX}{uuid.uuid4()}" if len(windows) > 1 else None + window_results: list[tuple[Sequence[MemoryRecord], Sequence[tuple[str, str]]]] = [] + committed_unit_ids: list[str] = [] + total_usage = TokenUsage() + + for window in windows: + records, extraction_usage = await self._extract_and_project_selected_chunks( + invocation, + execution, + plan, + window.chunks, + agent_name=agent_name, + fact_position_offset=(window.global_indices[0] if window.global_indices else 0), + ) + total_usage = total_usage + extraction_usage + checkpoint_callback = None + if window.is_last: + checkpoint_callback = self._compose_checkpoint_callback( + invocation, + execution, + plan, + expected_unit_ids_count=len(committed_unit_ids) + len(records), + prior_unit_ids=tuple(committed_unit_ids), + ) + + ownership = _backend_adapters(execution).document_ownership( + schema=execution.schema, + fresh=window.is_first and plan.existing is None, + ) + unit_of_work = RetainUnitOfWork( + connection_scope=lambda: acquire_with_retry(execution.pool), + adapter=PersistenceWriter( + pool=execution.pool, + embeddings_model=execution.embeddings_model, + entity_resolver=execution.entity_resolver, + config=execution.resolved_config, + ownership=ownership, + operation_activity=_backend_adapters(execution).operation_activity_fence( + invocation.operation_id, + schema=execution.schema, + ), + schema=execution.schema, + sanitize_log_identifiers=invocation.sanitize_log_identifiers, + ), + ) + + try: + request = await self._build_full_window_request( + invocation, + execution, + plan, + window.chunks, + records, + is_first=window.is_first, + is_last=window.is_last, + inflight_content_hash=inflight_content_hash, + final_content_hash=final_content_hash, + checkpoint_callback=checkpoint_callback, + outbox_callback=(outbox_callback if window.is_last else None), + ) + async with _database_budget(execution.db_semaphore): + result = await unit_of_work.execute(request) + except FreshDocumentOwnershipConflict as exc: + execution.entity_resolver.discard_pending_stats() + raise RetainOwnershipLostError( + "Retain lost fresh-document ownership for " + f"{_log_identifier(invocation, plan.intent.document_id)!r}; retry" + ) from exc + except BaseException: + execution.entity_resolver.discard_pending_stats() + raise + + if result.core.ownership is OwnershipDisposition.LOST: + execution.entity_resolver.discard_pending_stats() + raise RetainOwnershipLostError( + "Retain lost document ownership for " + f"{_log_identifier(invocation, plan.intent.document_id)!r} " + f"at FULL window {window.window_index}; retry" + ) + self._log_post_commit_failure(invocation, plan, result) + bindings = result.core.unit_ids_by_fact_key + window_results.append((records, bindings)) + units_by_key = dict(bindings) + if len(units_by_key) != len(bindings): + raise RetainResultMappingError("FULL window returned duplicate fact-key bindings") + try: + committed_unit_ids.extend(units_by_key[record.fact_key] for record in records) + except KeyError as exc: + raise RetainResultMappingError("FULL window returned an incomplete fact-key binding") from exc + + public_buckets = merge_window_unit_ids( + tuple(item.source_index for item in plan.intent.items), + tuple(window_results), + ) + public_iterator = iter(public_buckets) + unit_ids_by_content = tuple( + () if item.source_index is None else next(public_iterator) for item in plan.intent.items + ) + if committed_unit_ids: + final_ann_completed = await self._run_full_semantic_ann_best_effort( + invocation, + execution, + plan, + committed_unit_ids, + ) + if final_ann_completed: + await self._record_final_ann_completed_best_effort( + invocation, + execution, + plan.intent.document_id, + ) + return _DocumentOutcome( + unit_ids_by_content=unit_ids_by_content, + usage=total_usage, + processed_tokens=None, + ) + + @staticmethod + def _log_post_commit_failure( + invocation: RetainInvocation, + plan: _DocumentExecutionPlan, + result: Any, + ) -> None: + if result.post_commit.status is not PostCommitStatus.FAILED: + return + failure = result.post_commit.failure + sanitizer = _request_sanitizer(invocation, plan.intent.document_id) + _log_warning( + invocation, + "Retain post-commit stage failed for document %s at %s: %s", + sanitizer.identifier(plan.intent.document_id), + failure.stage.value if failure is not None else "unknown", + (sanitizer.text(failure.exception) if failure is not None else "unknown failure"), + identifiers=(plan.intent.document_id,), + ) + + async def _resume_committed_document( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + plan: _DocumentExecutionPlan, + ) -> _DocumentOutcome: + """Resume only the idempotent post-commit work for a durable core write.""" + + if plan.recovered_unit_bindings is None: # pragma: no cover - caller invariant + raise RetainCheckpointRecoveryError("Recovery plan has no committed unit-ID snapshot") + recovered_unit_ids = tuple(binding.unit_id for binding in plan.recovered_unit_bindings) + # Validate the durable unit/chunk/source mapping before running ANN or + # clearing its retry marker. Corrupt recovery state must have no + # post-commit side effects. + buckets = self._recovery_result_buckets(plan) + if plan.final_ann_pending: + final_ann_completed = True + if recovered_unit_ids: + final_ann_completed = await self._run_full_semantic_ann_best_effort( + invocation, + execution, + plan, + list(recovered_unit_ids), + ) + if final_ann_completed: + await self._record_final_ann_completed_best_effort( + invocation, + execution, + plan.intent.document_id, + ) + return _DocumentOutcome( + unit_ids_by_content=buckets, + usage=TokenUsage(), + processed_tokens=0, + ) + + @staticmethod + def _recovery_result_buckets( + plan: _DocumentExecutionPlan, + ) -> tuple[tuple[str, ...], ...]: + """Restore committed unit IDs to their original public content buckets. + + Units with a durable chunk association can be mapped back to their + immutable input source. Rows without a chunk association cannot be + attributed exactly; return every such unit in the first document + bucket rather than silently dropping or guessing any unit. + """ + + bindings = plan.recovered_unit_bindings + if bindings is None: # pragma: no cover - caller invariant + raise RetainCheckpointRecoveryError("Recovery plan has no committed unit bindings") + if not plan.intent.items: + raise RetainCheckpointRecoveryError("Recovery plan has no content items") + if not bindings: + return tuple(() for _ in plan.intent.items) + + recovered_unit_ids = tuple(binding.unit_id for binding in bindings) + if any(binding.chunk_index is None for binding in bindings): + return (recovered_unit_ids,) + tuple(() for _ in plan.intent.items[1:]) + + sources_by_chunk_index: dict[int, int | None] = {} + if plan.recovered_chunk_sources is not None: + for chunk_index, source_index in plan.recovered_chunk_sources: + if chunk_index in sources_by_chunk_index: + raise RetainCheckpointRecoveryError(f"Recovery plan has duplicate chunk index {chunk_index}") + sources_by_chunk_index[chunk_index] = source_index + else: + for chunk in plan.chunks: + if chunk.global_index in sources_by_chunk_index: + raise RetainCheckpointRecoveryError(f"Recovery plan has duplicate chunk index {chunk.global_index}") + sources_by_chunk_index[chunk.global_index] = chunk.source_index + + try: + positions = content_positions(plan.intent.items) + except (TypeError, ValueError) as exc: + raise RetainCheckpointRecoveryError( + f"Recovery plan for document {plan.intent.document_id!r} has invalid content sources" + ) from exc + + buckets: list[list[str]] = [[] for _ in plan.intent.items] + for binding in bindings: + chunk_index = binding.chunk_index + if chunk_index is None: # pragma: no cover - handled by fallback above + raise AssertionError("unmapped recovery binding escaped fallback") + try: + source_index = sources_by_chunk_index[chunk_index] + except KeyError as exc: + raise RetainCheckpointRecoveryError( + f"Committed unit {binding.unit_id!r} references unknown chunk index " + f"{chunk_index} for document {plan.intent.document_id!r}" + ) from exc + if source_index is None: + # An append may regenerate units for changed synthetic + # pre-append chunks. They remain part of post-commit ANN but + # never occupy a caller-visible input bucket. + continue + try: + content_position = positions[source_index] + except KeyError as exc: # pragma: no cover - chunk planner owns this invariant + raise RetainCheckpointRecoveryError( + f"Recovery chunk index {chunk_index} references unknown source {source_index!r}" + ) from exc + buckets[content_position].append(binding.unit_id) + return tuple(tuple(bucket) for bucket in buckets) + + @staticmethod + def _compose_checkpoint_callback( + invocation: RetainInvocation, + execution: RetainExecutionContext, + plan: _DocumentExecutionPlan, + *, + expected_unit_ids_count: int, + prior_unit_ids: tuple[str, ...] = (), + ) -> Any: + """Build the strict in-transaction core checkpoint callback.""" + + if invocation.operation_id is None: + return None + + async def record_exact_core_commit( + connection: Any, + current_unit_ids_by_content: tuple[tuple[str, ...], ...], + ) -> None: + current_unit_ids = tuple(unit_id for bucket in current_unit_ids_by_content for unit_id in bucket) + unit_ids = (*prior_unit_ids, *current_unit_ids) + if len(unit_ids) != expected_unit_ids_count: + raise RetainCheckpointRecoveryError( + f"Core write for document {plan.intent.document_id!r} returned " + f"{len(unit_ids)} unit IDs; expected {expected_unit_ids_count}" + ) + if len(unit_ids) != len(set(unit_ids)): + raise RetainCheckpointRecoveryError( + f"Core write for document {plan.intent.document_id!r} returned duplicate unit IDs" + ) + await ( + _backend_adapters(execution) + .checkpoint_store( + connection, + schema=execution.schema, + ) + .record_core_committed( + invocation.operation_id, + plan.intent.document_id, + unit_ids=unit_ids, + requires_final_ann=(plan.change.kind is DocumentChangeKind.FULL), + ) + ) + + return record_exact_core_commit + + @staticmethod + async def _record_final_ann_completed_best_effort( + invocation: RetainInvocation, + execution: RetainExecutionContext, + document_id: str, + ) -> None: + if invocation.operation_id is None: + return + try: + async with acquire_with_retry(execution.pool) as connection: + await ( + _backend_adapters(execution) + .checkpoint_store( + connection, + schema=execution.schema, + ) + .record_final_ann_completed( + invocation.operation_id, + document_id, + ) + ) + except Exception: + # The final ANN pass is idempotent and best effort. Leaving its + # marker in place causes a safe retry rather than a duplicate core + # write or outbox event. + _log_warning( + invocation, + "Retain could not clear final ANN checkpoint for operation %s, document %s", + _log_identifier(invocation, invocation.operation_id), + _log_identifier(invocation, document_id), + identifiers=(document_id,), + exc_info=True, + ) + + @staticmethod + async def _run_full_semantic_ann_best_effort( + invocation: RetainInvocation, + execution: RetainExecutionContext, + plan: _DocumentExecutionPlan, + committed_unit_ids: list[str], + ) -> bool: + """Run the deferred semantic-link pass and report whether it completed.""" + + try: + await run_final_semantic_ann( + execution.pool, + invocation.bank_id, + committed_unit_ids, + execution.resolved_config, + [], + ) + return True + except Exception: + # Facts and the retrieval-critical core graph are already committed, + # so semantic ANN remains a post-commit best-effort operation. The + # caller must preserve the durable retry marker after a failure. + _log_warning( + invocation, + "Retain final semantic ANN failed for document %s", + _log_identifier(invocation, plan.intent.document_id), + identifiers=(plan.intent.document_id,), + exc_info=True, + ) + return False + + @staticmethod + def _chunks_for_extraction(plan: _DocumentExecutionPlan) -> tuple[ChunkPlan, ...]: + if plan.change.kind is DocumentChangeKind.FULL: + return plan.chunks + if plan.change.kind is DocumentChangeKind.DELTA: + selected = set(plan.change.chunks_to_process) + return tuple(chunk for chunk in plan.chunks if chunk.global_index in selected) + return () + + @staticmethod + def _processed_chunk_tokens( + plan: _DocumentExecutionPlan, + selected_chunks: Sequence[ChunkPlan], + ) -> int: + items_by_source = {item.source_index: item for item in plan.intent.items} + total = 0 + for chunk in selected_chunks: + try: + item = items_by_source[chunk.source_index] + except KeyError as exc: # pragma: no cover - chunk planning owns this invariant + raise RetainResultMappingError(f"Chunk {chunk.chunk_key!r} references an unknown source item") from exc + total += count_tokens(chunk.text) + total += count_tokens(item.context) + return total + + async def _extract_and_project_selected_chunks( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + plan: _DocumentExecutionPlan, + selected_chunks: Sequence[ChunkPlan], + *, + agent_name: str, + fact_position_offset: int = 0, + ) -> tuple[tuple[MemoryRecord, ...], TokenUsage]: + mode = ExtractionMode(execution.resolved_config.retain_extraction_mode) + if not selected_chunks: + return (), TokenUsage() + configured_concurrency = _projection_pipeline_concurrency(execution.resolved_config) + window_started = time.perf_counter() + extraction_seconds: list[float] = [] + embedding_seconds: list[float] = [] + embedder = EmbeddingModelAdapter(execution.embeddings_model) + embedding_model_version_value = embedding_model_version(execution.embeddings_model) + extraction_version = getattr(execution.resolved_config, "extraction_prompt_version", "5w-v1") + + async def project(candidates) -> tuple[tuple[MemoryRecord, ...], float]: + started = time.perf_counter() + records = await project_embeddings( + candidates, + embedder=embedder, + format_date=execution.format_date_fn, + embedding_model_version=embedding_model_version_value, + extraction_version=extraction_version, + failure_policy=getattr( + execution.resolved_config, + "retain_embedding_failure_policy", + EmbeddingFailurePolicy.STORE_WITHOUT_EMBEDDING, + ), + ) + return records, time.perf_counter() - started + + async def extract_structured(chunks: Sequence[ChunkPlan]): + layout = build_prechunked_extraction_layout( + plan.intent.items, + chunks, + ) + extraction = await FactExtractorAdapter( + llm_config=execution.llm_config, + config=execution.resolved_config, + agent_name=agent_name, + pool=execution.pool, + operation_id=invocation.operation_id, + schema=execution.schema, + batch_checkpoint_clearer=self._provider_batch_checkpoint_clearer( + invocation, + execution, + ), + ).extract( + layout.extraction_request( + ExtractionPolicy( + mode=mode, + fact_type_override=invocation.fact_type_override, + ) + ) + ) + return layout.remap_result(extraction) + + chunk_batch_size = getattr(execution.resolved_config, "retain_chunk_batch_size", 100) + pipeline_selected_chunks = ( + mode is not ExtractionMode.CHUNKS + and len(selected_chunks) > 1 + and not getattr(execution.resolved_config, "retain_batch_enabled", False) + and isinstance(chunk_batch_size, int) + and not isinstance(chunk_batch_size, bool) + and chunk_batch_size > 0 + and len(selected_chunks) <= chunk_batch_size + ) + + if pipeline_selected_chunks: + extraction_semaphore = asyncio.Semaphore(configured_concurrency) + embedding_semaphore = asyncio.Semaphore(configured_concurrency) + + async def extract_and_project_one(chunk: ChunkPlan) -> _ProjectedChunkOutcome: + async with extraction_semaphore: + extract_started = time.perf_counter() + extraction = await extract_structured((chunk,)) + extract_elapsed = time.perf_counter() - extract_started + + async with embedding_semaphore: + records, embed_elapsed = await project(extraction.candidates) + return _ProjectedChunkOutcome( + records=records, + usage=extraction.usage, + extraction_seconds=extract_elapsed, + embedding_seconds=embed_elapsed, + ) + + results = await asyncio.gather( + *(extract_and_project_one(chunk) for chunk in selected_chunks), + return_exceptions=True, + ) + for result in results: + if isinstance(result, BaseException): + # gather preserves input order and waits for every sibling. + # Raise the earliest source-chunk failure only after all + # provider work has reached a terminal state. + raise result + + records: list[MemoryRecord] = [] + extraction_usage = TokenUsage() + for result in results: + # Every BaseException was rejected in the loop above. + assert isinstance(result, _ProjectedChunkOutcome) + records.extend(result.records) + extraction_usage = extraction_usage + result.usage + extraction_seconds.append(result.extraction_seconds) + embedding_seconds.append(result.embedding_seconds) + frozen_records = tuple(records) + _log_window_producer_summary( + path="pipelined", + chunks=len(selected_chunks), + facts=len(frozen_records), + wall_seconds=time.perf_counter() - window_started, + extraction_seconds=extraction_seconds, + embedding_seconds=embedding_seconds, + configured_concurrency=configured_concurrency, + ) + return frozen_records, extraction_usage + + if mode is ExtractionMode.CHUNKS: + candidates = extract_passthrough( + selected_chunks, + plan.intent.items, + fact_type_override=invocation.fact_type_override, + fact_position_offset=fact_position_offset, + ) + extraction_usage = TokenUsage() + else: + extract_started = time.perf_counter() + extraction = await extract_structured(selected_chunks) + extraction_seconds.append(time.perf_counter() - extract_started) + candidates = extraction.candidates + extraction_usage = extraction.usage + + records, embed_elapsed = await project(candidates) + embedding_seconds.append(embed_elapsed) + _log_window_producer_summary( + path="window_batch", + chunks=len(selected_chunks), + facts=len(records), + wall_seconds=time.perf_counter() - window_started, + extraction_seconds=extraction_seconds, + embedding_seconds=embedding_seconds, + configured_concurrency=configured_concurrency, + ) + return records, extraction_usage + + @staticmethod + def _provider_batch_checkpoint_clearer( + invocation: RetainInvocation, + execution: RetainExecutionContext, + ) -> Any: + if invocation.operation_id is None: + return None + + async def clear_completed_provider_batch() -> None: + async with acquire_with_retry(execution.pool) as connection: + await ( + _backend_adapters(execution) + .checkpoint_store( + connection, + schema=execution.schema, + ) + .clear_provider_batch(invocation.operation_id) + ) + + return clear_completed_provider_batch + + async def _build_full_window_request( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + plan: _DocumentExecutionPlan, + selected_chunks: Sequence[ChunkPlan], + records: Sequence[MemoryRecord], + *, + is_first: bool, + is_last: bool, + inflight_content_hash: str | None, + final_content_hash: str, + checkpoint_callback: Any = None, + outbox_callback: Any = None, + ) -> WriteWindowRequest: + retain_params, document_tags = retain_document_metadata(plan.intent.items) + payload = await self._build_fact_payload( + invocation, + execution, + plan, + selected_chunks, + records, + ) + if is_first: + document_window = FirstFullWriteWindow( + combined_content=plan.combined_content, + is_first_batch=True, + retain_params=retain_params, + document_tags=document_tags, + recovery=False, + expected_existing_content_hash=(plan.existing.content_hash if plan.existing is not None else None), + expects_unhashed_existing_document=(plan.existing is not None and not plan.existing.content_hash), + continuation_content_hash=(inflight_content_hash if not is_last else None), + ) + else: + if inflight_content_hash is None: # pragma: no cover - window planner invariant + raise RetainError("Later FULL window requires an in-flight ownership hash") + document_window = LaterFullWriteWindow( + expected_content_hash=inflight_content_hash, + completed_content_hash=final_content_hash if is_last else None, + ) + return WriteWindowRequest( + bank_id=invocation.bank_id, + document_id=plan.intent.document_id, + document_window=document_window, + contents=payload.storage_contents, + chunks=payload.chunks, + facts=payload.facts, + graph=payload.graph, + skip_semantic_links=True, + checkpoint_callback=checkpoint_callback, + outbox_callback=outbox_callback, + log_buffer=[], + ) + + async def _build_write_request( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + plan: _DocumentExecutionPlan, + selected_chunks: Sequence[ChunkPlan], + records: Sequence[MemoryRecord], + *, + processed_tokens: int | None, + checkpoint_callback: Any = None, + outbox_callback: Any = None, + ) -> RetainWriteRequest: + retain_params, document_tags = retain_document_metadata(plan.intent.items) + if plan.change.kind is DocumentChangeKind.METADATA_ONLY: + if plan.existing is None or not plan.existing.content_hash: # pragma: no cover - classifier invariant + raise RetainError("Metadata-only change requires an existing hash snapshot") + return MetadataOnlyWriteRequest( + bank_id=invocation.bank_id, + document_id=plan.intent.document_id, + expected_content_hash=plan.existing.content_hash, + combined_content=plan.combined_content, + input_slot_count=len(plan.intent.items), + retain_params=retain_params, + document_tags=document_tags, + checkpoint_callback=checkpoint_callback, + outbox_callback=outbox_callback, + ) + + payload = await self._build_fact_payload( + invocation, + execution, + plan, + selected_chunks, + records, + ) + if plan.change.kind is not DocumentChangeKind.DELTA: + raise RetainError(f"Cannot build a write request for {plan.change.kind.value}") + if plan.existing is None or not plan.existing.content_hash: # pragma: no cover - classifier invariant + raise RetainError("Delta change requires an existing hash snapshot") + + existing_by_index = {chunk.chunk_index: chunk for chunk in plan.existing_chunks} + + def existing_write(index: int) -> ExistingChunkWrite: + try: + chunk = existing_by_index[index] + except KeyError as exc: + raise RetainError(f"Change plan references missing existing chunk index {index}") from exc + return ExistingChunkWrite(chunk_id=chunk.chunk_id, chunk_index=index) + + return DeltaWriteRequest( + bank_id=invocation.bank_id, + document_id=plan.intent.document_id, + expected_content_hash=plan.existing.content_hash, + combined_content=plan.combined_content, + contents=payload.storage_contents, + unchanged_chunk_indices=plan.change.unchanged, + changed_chunks=tuple(existing_write(index) for index in plan.change.changed), + added_chunk_indices=plan.change.added, + removed_chunks=tuple(existing_write(index) for index in plan.change.removed), + chunks=payload.chunks, + facts=payload.facts, + graph=payload.graph, + processed_tokens=processed_tokens or 0, + retain_params=retain_params, + document_tags=document_tags, + skip_semantic_links=False, + checkpoint_callback=checkpoint_callback, + outbox_callback=outbox_callback, + log_buffer=(), + ) + + async def _build_fact_payload( + self, + invocation: RetainInvocation, + execution: RetainExecutionContext, + plan: _DocumentExecutionPlan, + selected_chunks: Sequence[ChunkPlan], + records: Sequence[MemoryRecord], + ) -> _FactPayload: + storage_contents = tuple(content_to_storage(item) for item in plan.intent.items) + positions = content_positions(plan.intent.items) + fact_positions = {record.fact_key: position for position, record in enumerate(records)} + fact_writes: list[FactWrite] = [] + processed_facts = [] + for record in records: + try: + content_index = positions[record.source_index] + except KeyError as exc: # pragma: no cover - passthrough validates this + raise RetainResultMappingError( + f"Record {record.fact_key!r} references unknown source index {record.source_index!r}" + ) from exc + extracted = record_to_extracted_fact(record, content_index=content_index) + processed = record_to_processed_fact( + record, + document_id=plan.intent.document_id, + content_index=content_index, + fact_positions=fact_positions, + ) + fact_writes.append( + FactWrite( + fact_key=record.fact_key, + chunk_key=record.chunk_key, + extracted=extracted, + processed=processed, + ) + ) + processed_facts.append(processed) + + phase1 = None + if processed_facts: + execution.entity_resolver.discard_pending_stats() + phase1_kwargs = { + "skip_semantic_ann": ( + plan.change.kind is DocumentChangeKind.FULL + or not getattr(execution.resolved_config, "write_semantic_links", True) + ) + } + phase1 = await pre_resolve_entities( + execution.pool, + execution.entity_resolver, + invocation.bank_id, + list(storage_contents), + [fact.fact_key for fact in fact_writes], + processed_facts, + execution.resolved_config, + [], + **phase1_kwargs, + ) + chunk_metadata = chunks_to_storage(selected_chunks, plan.intent.items, records) + return _FactPayload( + storage_contents=storage_contents, + chunks=tuple( + ChunkWrite(chunk_key=chunk.chunk_key, metadata=metadata) + for chunk, metadata in zip(selected_chunks, chunk_metadata, strict=True) + ), + facts=tuple(fact_writes), + graph=_graph_write(phase1), + ) + + @staticmethod + def _merge_document_result( + result_by_input: list[list[str]], + intent: DocumentIntent, + document_result: Sequence[Sequence[str]], + ) -> None: + if len(document_result) != len(intent.items): + raise RetainResultMappingError( + f"Document {intent.document_id!r} returned {len(document_result)} result buckets " + f"for {len(intent.items)} content items" + ) + for item, unit_ids in zip(intent.items, document_result, strict=True): + if item.source_index is None: + continue + if not 0 <= item.source_index < len(result_by_input): + raise RetainResultMappingError( + f"Document {intent.document_id!r} references out-of-range input slot {item.source_index}" + ) + if result_by_input[item.source_index]: + raise RetainResultMappingError( + f"Input slot {item.source_index} received results from more than one document" + ) + bucket = list(unit_ids) + if any(not isinstance(unit_id, str) or not unit_id for unit_id in bucket): + raise RetainResultMappingError(f"Document {intent.document_id!r} returned an invalid unit ID") + result_by_input[item.source_index] = bucket diff --git a/core/dataplane/hms_api/engine/memory_engine.py b/core/dataplane/hms_api/engine/memory_engine.py index e88b9bd..16e2d90 100644 --- a/core/dataplane/hms_api/engine/memory_engine.py +++ b/core/dataplane/hms_api/engine/memory_engine.py @@ -61,12 +61,158 @@ class _MultimodalCommandSuperseded(RuntimeError): """Internal control flow used to roll back an obsolete child retain.""" +class _RetainOperationCancelled(RuntimeError): + """Internal control flow for a tracked Retain cancelled between commits.""" + + _MULTIMODAL_CACHE_PIPELINE_PREFIX = "__hms_pipeline__." _MULTIMODAL_UNCONSTRAINED_HINT = "unconstrained" _MULTIMODAL_FALLBACK_POLICY_VERSION = "typed-not-applicable-only-v1" +def _is_retain_document_id(value: Any) -> bool: + """Match the ingestion planner's non-empty document-ID boundary.""" + + return isinstance(value, str) and bool(value.strip()) + + +def _retain_group_payload_digest( + indexed_items: list[tuple[int, "RetainContentDict"]], +) -> str: + """Hash one logical document without depending on token split settings.""" + + payload = [ + { + "source_index": source_index, + "item": dict(item), + } + for source_index, item in indexed_items + ] + serialized = json.dumps( + payload, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + default=_json_default, + ) + return hashlib.sha256(serialized.encode("utf-8")).hexdigest() + + +def _derive_tracked_retain_document_id( + operation_id: str, + *, + indexed_items: list[tuple[int, "RetainContentDict"]], +) -> str: + """Derive a retry-stable ID for one anonymous logical document. + + The identity includes original source positions and the complete logical + payload, but not the mutable token threshold or shard ordinal. A retry + after a configuration change therefore addresses the same document rather + than aliasing a previously committed shard with different contents. + """ + + if not indexed_items: + raise ValueError("indexed_items must not be empty") + operation_namespace = uuid.UUID(operation_id) + identity = f"hms:retain:logical-document:source-stable-v1:{_retain_group_payload_digest(indexed_items)}" + return str(uuid.uuid5(operation_namespace, identity)) + + +def _prepare_retain_document_groups( + contents: list["RetainContentDict"], + *, + operation_id: str | None, +) -> tuple[list["RetainContentDict"], list[list[tuple[int, "RetainContentDict"]]]]: + """Resolve logical document groups before token batching. + + The grouping mirrors ``plan_documents`` over the complete request: + no explicit ID means one shared document; one explicit ID absorbs + anonymous siblings; multiple explicit IDs leave each anonymous item as an + independent generated document. Generated tracked IDs are payload-stable + across retries. Groups retain original source indices so outer batching can + restore result order even when repeated document IDs are interleaved. + """ + + prepared = cast(list[RetainContentDict], [dict(item) for item in contents]) + explicit_document_ids: set[str] = set() + for source_index, item in enumerate(prepared): + document_id = item.get("document_id") + if document_id is not None and not isinstance(document_id, str): + raise TypeError( + f"contents[{source_index}].document_id must be a string or None, got {type(document_id).__name__}" + ) + if _is_retain_document_id(document_id): + explicit_document_ids.add(document_id) + + if not explicit_document_ids: + indexed = list(enumerate(prepared)) + if operation_id is not None: + document_id = _derive_tracked_retain_document_id( + operation_id, + indexed_items=indexed, + ) + for item in prepared: + item["document_id"] = document_id + indexed = list(enumerate(prepared)) + return prepared, [indexed] + + if len(explicit_document_ids) == 1: + document_id = next(iter(explicit_document_ids)) + for item in prepared: + if not _is_retain_document_id(item.get("document_id")): + item["document_id"] = document_id + return prepared, [list(enumerate(prepared))] + + for source_index, item in enumerate(prepared): + if _is_retain_document_id(item.get("document_id")): + continue + indexed_item = [(source_index, item)] + item["document_id"] = ( + _derive_tracked_retain_document_id( + operation_id, + indexed_items=indexed_item, + ) + if operation_id is not None + else str(uuid.uuid4()) + ) + + groups_by_document: dict[str, list[tuple[int, RetainContentDict]]] = {} + for source_index, item in enumerate(prepared): + item_document_id = item.get("document_id") + if not _is_retain_document_id(item_document_id): + raise ValueError("Retain document grouping requires a non-empty document_id") + groups_by_document.setdefault(item_document_id, []).append((source_index, item)) + return prepared, list(groups_by_document.values()) + + +def _split_retain_document_groups( + groups: list[list[tuple[int, "RetainContentDict"]]], + *, + tokens_per_batch: int, +) -> list[list[tuple[int, "RetainContentDict"]]]: + """Pack complete logical documents into bounded token batches.""" + + if isinstance(tokens_per_batch, bool) or not isinstance(tokens_per_batch, int) or tokens_per_batch <= 0: + raise ValueError("retain_batch_tokens must be a positive integer") + + sub_batches: list[list[tuple[int, RetainContentDict]]] = [] + current_batch: list[tuple[int, RetainContentDict]] = [] + current_batch_tokens = 0 + for group in groups: + group_tokens = sum(count_tokens(item.get("content", "")) for _source_index, item in group) + if current_batch and current_batch_tokens + group_tokens > tokens_per_batch: + sub_batches.append(current_batch) + current_batch = list(group) + current_batch_tokens = group_tokens + else: + current_batch.extend(group) + current_batch_tokens += group_tokens + if current_batch: + sub_batches.append(current_batch) + return sub_batches + + def _derive_anonymous_multimodal_document_id( *, tenant_scope: str | None, @@ -984,7 +1130,7 @@ def __init__( self._search_semaphore = asyncio.Semaphore(get_config().recall_max_concurrent) # Backpressure for retain DB writes: limit concurrent transactions to prevent contention - # on entity/link tables. Acquired in the orchestrator *after* LLM extraction completes, + # on entity/link tables. Acquired in the pipeline after LLM extraction completes, # so LLM calls run in full parallelism while only the DB-heavy phase is throttled. # Configurable via HMS_API_RETAIN_MAX_CONCURRENT (default: 4). self._put_semaphore = asyncio.Semaphore(get_config().retain_max_concurrent) @@ -1148,6 +1294,7 @@ def record_multimodal_retain(*, success: bool, reason: str | None = None, cancel file_metadata=task_dict.get("_file_metadata"), after_publish=webhook_callback, ) + from .ingestion import RetainOperationInactiveError try: await self.retain_batch_async( @@ -1281,6 +1428,12 @@ def record_multimodal_retain(*, success: bool, reason: str | None = None, cancel terminal_command.bank_id, ) return False + except RetainOperationInactiveError as exc: + record_multimodal_retain(success=False, reason="operation.cancelled", cancelled=True) + raise _RetainOperationCancelled("Tracked Retain was cancelled before its core write") from exc + except _RetainOperationCancelled: + record_multimodal_retain(success=False, reason="operation.cancelled", cancelled=True) + raise except asyncio.CancelledError: record_multimodal_retain(success=False, reason="operation.cancelled", cancelled=True) raise @@ -1810,10 +1963,10 @@ async def _save_video_segment_checkpoint(checkpoint: Any) -> None: ) if ledgered_multimodal and not is_multimodal: - # A typed not-applicable outcome may legitimately select a legacy + # A typed not-applicable outcome may legitimately select another # parser from the explicit chain. Release the unused multimodal - # command/descriptor state; the winning legacy output then follows - # the unchanged legacy retain path. + # command/descriptor state; the selected output then follows the + # standard retain path. backend = await self._get_backend() async with acquire_with_retry(backend) as conn: ledger = MultimodalLedger.for_connection(conn, schema=get_current_schema()) @@ -2289,20 +2442,12 @@ async def execute_task(self, task_dict: dict[str, Any]): # Check if operation was cancelled (only for tasks with operation_id) if operation_id: - try: - backend = await self._get_backend() - async with acquire_with_retry(backend) as conn: - result = await conn.fetchrow( - f"SELECT status FROM {fq_table('async_operations')} WHERE operation_id = $1", - uuid.UUID(operation_id), - ) - if not result or result["status"] == "cancelled": - # Operation was cancelled, skip processing - logger.info(f"Skipping cancelled operation: {operation_id}") - return - except Exception as e: - logger.error(f"Failed to check operation status {operation_id}: {e}") - # Continue with processing if we can't check status + if not await self._check_op_alive(operation_id): + # Child liveness includes the parent batch state. This closes + # the claim/cancel race where a worker claims a child just + # before its parent is cancelled. + logger.info("Skipping cancelled or terminal operation") + return consolidation_result: dict | None = None bank_id = task_dict.get("bank_id") @@ -2359,6 +2504,12 @@ async def execute_task(self, task_dict: dict[str, Any]): # would convert a legitimate defer into a 60-second RetryTaskAt # and lose the "not a failure" semantics entirely. raise + except _RetainOperationCancelled: + # The operation row already records cancellation (or was + # removed with its bank). Never retry, fail, or complete a task + # that stopped at the between-document cancellation fence. + audit_entry.response = {"status": "cancelled", "operation_id": operation_id} + return except Exception as e: logger.error(f"Task execution failed: {task_type}, error: {e}") import traceback @@ -2412,8 +2563,8 @@ async def execute_task(self, task_dict: dict[str, Any]): raise RetryTaskAt(retry_at=datetime.now(UTC) + timedelta(seconds=60), message=str(e)) if task_type == "batch_retain" and task_dict.get("_multimodal_command") and operation_id: # Route the terminal multimodal child failure through - # the ledger-aware transaction. Legacy tasks still - # re-raise here and retain the poller's old failure + # the ledger-aware transaction. Standard tasks still + # re-raise here and retain the poller's existing failure # path. await self._mark_operation_failed(operation_id, str(e), error_traceback) return @@ -2606,19 +2757,43 @@ async def _delete_operation_record(self, operation_id: str): logger.error(f"Failed to delete async operation record {operation_id}: {e}") async def _check_op_alive(self, operation_id: str) -> bool: - """Return False if the operation was cancelled or no longer exists (e.g. bank deleted via CASCADE). + """Return whether an operation and its optional batch parent are active. Long-running operations should call this at natural checkpoints (e.g. after each - committed batch) to detect cancellation or bank deletion early and abort cleanly. + committed batch) to detect cancellation, terminal parent state, or bank deletion + early and abort cleanly. """ try: backend = await self._get_backend() async with acquire_with_retry(backend) as conn: row = await conn.fetchrow( - f"SELECT status FROM {fq_table('async_operations')} WHERE operation_id = $1", + f""" + SELECT status, bank_id, result_metadata + FROM {fq_table("async_operations")} + WHERE operation_id = $1 + """, uuid.UUID(operation_id), ) - return row is not None and row["status"] != "cancelled" + if row is None or row["status"] not in {"pending", "processing"}: + return False + + result_metadata = conn.parse_json(row["result_metadata"]) or {} + parent_operation_id = ( + result_metadata.get("parent_operation_id") if isinstance(result_metadata, dict) else None + ) + if not parent_operation_id: + return True + + parent_row = await conn.fetchrow( + f""" + SELECT status + FROM {fq_table("async_operations")} + WHERE operation_id = $1 AND bank_id = $2 + """, + uuid.UUID(parent_operation_id), + row["bank_id"], + ) + return parent_row is not None and parent_row["status"] in {"pending", "processing"} except Exception as e: logger.error(f"Failed to check operation liveness {operation_id}: {e}") return True # Assume alive on DB error to avoid false-positive aborts @@ -2762,8 +2937,8 @@ async def _mark_operation_failed(self, operation_id: str, error_message: str, er """Helper to mark an operation as failed in the database. Multimodal child retains use the same transaction to close their - durable document command and update the parent status metadata. Legacy - child and non-child operations retain the existing failure behavior. + durable document command and update the parent status metadata. + Non-multimodal operations retain the existing failure behavior. """ try: backend = await self._get_backend() @@ -2794,13 +2969,14 @@ async def _mark_operation_failed(self, operation_id: str, error_message: str, er logger.info("Marked multimodal child operation as failed: %s", operation_id) return - # Mark legacy operations exactly as before. The multimodal - # branch above is the only path that inspects task payloads. + # Mark standard operations exactly as before. The + # multimodal branch is the only path that inspects payloads. row = await conn.fetchrow( f""" UPDATE {fq_table("async_operations")} SET status = 'failed', error_message = $2, updated_at = NOW() - WHERE operation_id = $1 AND status <> 'cancelled' + WHERE operation_id = $1 + AND status IN ('pending', 'processing') RETURNING operation_id """, uuid.UUID(operation_id), @@ -2811,8 +2987,8 @@ async def _mark_operation_failed(self, operation_id: str, error_message: str, er return logger.info(f"Marked async operation as failed: {operation_id}") - # Check if this is a legacy child operation and update its - # parent if all siblings are done, preserving old semantics. + # Check whether this is a child operation and update its + # parent once all siblings are terminal. await self._maybe_update_parent_operation(operation_id, conn) except Exception as e: logger.error(f"Failed to mark operation as failed {operation_id}: {e}") @@ -2833,14 +3009,13 @@ async def _mark_operation_completed(self, operation_id: str): UPDATE {fq_table("async_operations")} SET status = 'completed', updated_at = NOW(), completed_at = NOW() WHERE operation_id = $1 + AND status IN ('pending', 'processing') RETURNING operation_id """, uuid.UUID(operation_id), ) if row is None: - logger.info( - f"Operation {operation_id} no longer exists (bank deleted), skipping mark-completed" - ) + logger.info(f"Operation {operation_id} is missing or terminal, skipping mark-completed") return logger.info(f"Marked async operation as completed: {operation_id}") @@ -2876,6 +3051,7 @@ async def _mark_operation_completed_and_fire_webhook( UPDATE {fq_table("async_operations")} SET status = 'completed', updated_at = NOW(), completed_at = NOW() WHERE operation_id = $1 + AND status IN ('pending', 'processing') RETURNING operation_id """, uuid.UUID(operation_id), @@ -2946,7 +3122,7 @@ async def _maybe_update_parent_operation(self, child_operation_id: str, conn): # Use FOR UPDATE to ensure only one child can update the parent at a time parent_row = await conn.fetchrow( f""" - SELECT operation_id + SELECT operation_id, status FROM {fq_table("async_operations")} WHERE operation_id = $1 AND bank_id = $2 FOR UPDATE @@ -2958,6 +3134,11 @@ async def _maybe_update_parent_operation(self, child_operation_id: str, conn): if not parent_row: # Parent doesn't exist (shouldn't happen) return + if parent_row["status"] not in {"pending", "processing"}: + # Terminal parent state is authoritative. In particular, an + # accepted cancellation must never be overwritten by a late + # child completion or failure. + return # Get all sibling operations (including this one). # This query runs in the same transaction, so it sees the current @@ -2971,19 +3152,22 @@ async def _maybe_update_parent_operation(self, child_operation_id: str, conn): SELECT status, error_message FROM {fq_table("async_operations")} WHERE bank_id = $1 - AND result_metadata::jsonb @> $2::jsonb + AND (result_metadata->>'parent_operation_id')::uuid = $2 """, bank_id, - json.dumps({"parent_operation_id": parent_operation_id}), + uuid.UUID(parent_operation_id), ) if not siblings: return - # Check if all siblings are done (completed or failed) + # Cancellation is terminal for aggregation. Otherwise a parent can + # remain pending forever when its last outstanding child is + # cancelled directly or through parent cancellation propagation. all_completed = all(sib["status"] == "completed" for sib in siblings) any_failed = any(sib["status"] == "failed" for sib in siblings) - all_done = all(sib["status"] in ("completed", "failed") for sib in siblings) + any_cancelled = any(sib["status"] == "cancelled" for sib in siblings) + all_done = all(sib["status"] in ("completed", "failed", "cancelled") for sib in siblings) if not all_done: # Some siblings still pending/processing @@ -3002,11 +3186,26 @@ async def _maybe_update_parent_operation(self, child_operation_id: str, conn): f""" UPDATE {fq_table("async_operations")} SET status = $2, error_message = $3, updated_at = NOW() - WHERE operation_id = $1 + WHERE operation_id = $1 AND bank_id = $4 + AND status IN ('pending', 'processing') """, uuid.UUID(parent_operation_id), new_status, _summarise_child_error_messages(siblings), + bank_id, + ) + elif any_cancelled: + new_status = "cancelled" + await conn.execute( + f""" + UPDATE {fq_table("async_operations")} + SET status = $2, updated_at = NOW() + WHERE operation_id = $1 AND bank_id = $3 + AND status IN ('pending', 'processing') + """, + uuid.UUID(parent_operation_id), + new_status, + bank_id, ) elif all_completed: new_status = "completed" @@ -3014,10 +3213,12 @@ async def _maybe_update_parent_operation(self, child_operation_id: str, conn): f""" UPDATE {fq_table("async_operations")} SET status = $2, updated_at = NOW(), completed_at = NOW() - WHERE operation_id = $1 + WHERE operation_id = $1 AND bank_id = $3 + AND status IN ('pending', 'processing') """, uuid.UUID(parent_operation_id), new_status, + bank_id, ) logger.info(f"Updated parent operation {parent_operation_id} to status '{new_status}' (all children done)") @@ -3863,8 +4064,9 @@ async def retain_batch_async( This is MUCH more efficient than calling retain_async multiple times: - Extracts facts from all contents in parallel - Generates ALL embeddings in ONE batch - - Does ALL database operations in ONE transaction - - Automatically chunks large batches to prevent timeouts + - Commits each logical document through ownership-guarded transactions + - Packs complete logical documents into bounded outer batches; large + individual documents use the Retain pipeline's internal write windows Args: bank_id: Unique identifier for the bank @@ -3877,7 +4079,6 @@ async def retain_batch_async( Applies the same document_id to ALL content items that don't specify their own. fact_type_override: Override fact type for all facts ('world', 'experience') return_usage: If True, returns tuple of (unit_ids, TokenUsage). Default False for backward compatibility. - Returns: If return_usage=False: List of lists of unit IDs (one list per content item) If return_usage=True: Tuple of (unit_ids, TokenUsage) @@ -3930,35 +4131,28 @@ async def retain_batch_async( if result and result.contents is not None: contents = result.contents - # Engine-owned copy: the orchestrator clears per-item "content" strings - # after building the document's combined text (memory pressure - # optimization, see retain/orchestrator.py). Without an internal copy - # those mutations leak back to the caller's dicts. + # Keep an engine-owned copy so validation, normalization, and future + # pipeline optimizations cannot mutate caller-owned dictionaries. contents = cast(list[RetainContentDict], [dict(c) for c in contents]) - # Apply batch-level document_id to contents that don't have their own (backwards compatibility) - if document_id: + # Apply the deprecated batch-level ID only to anonymous items. The + # service receives the resulting per-item IDs so explicit item IDs keep + # their documented precedence. + if _is_retain_document_id(document_id): for item in contents: - if "document_id" not in item: + if not _is_retain_document_id(item.get("document_id")): item["document_id"] = document_id - # Validate no duplicate document_ids in the batch - # Having duplicate document_ids causes race conditions in document upserts during parallel processing - doc_ids = [item.get("document_id") for item in contents if item.get("document_id")] - if len(doc_ids) != len(set(doc_ids)): - from collections import Counter - - duplicates = [doc_id for doc_id, count in Counter(doc_ids).items() if count > 1] - raise ValueError( - f"Batch contains duplicate document_ids: {duplicates}. " - f"Each content item in a batch must have a unique document_id to avoid race conditions." - ) - # Validate update_mode=append requires document_id for item in contents: - if item.get("update_mode") == "append" and not item.get("document_id"): + if item.get("update_mode") == "append" and not _is_retain_document_id(item.get("document_id")): raise ValueError("update_mode='append' requires a document_id") + contents, document_groups = _prepare_retain_document_groups( + contents, + operation_id=operation_id, + ) + # Auto-chunk large batches by token count to avoid timeouts and memory issues # Calculate total token count total_tokens = sum(count_tokens(item.get("content", "")) for item in contents) @@ -3974,43 +4168,28 @@ async def retain_batch_async( tokens_per_batch = config.retain_batch_tokens if total_tokens > tokens_per_batch: - # Split into smaller batches based on token count + # Split between complete logical documents. A single document may + # exceed this outer threshold; the Retain service bounds its work + # using internal chunk write windows without changing ownership. logger.info( f"Large batch detected ({total_tokens:,} tokens from {len(contents)} items). Splitting into sub-batches of ~{tokens_per_batch:,} tokens each..." ) - sub_batches = [] - current_batch = [] - current_batch_tokens = 0 - - for item in contents: - item_tokens = count_tokens(item.get("content", "")) - - # If adding this item would exceed the limit, start a new batch - # (unless current batch is empty - then we must include it even if it's large) - if current_batch and current_batch_tokens + item_tokens > tokens_per_batch: - sub_batches.append(current_batch) - current_batch = [item] - current_batch_tokens = item_tokens - else: - current_batch.append(item) - current_batch_tokens += item_tokens - - # Add the last batch - if current_batch: - sub_batches.append(current_batch) + sub_batches = _split_retain_document_groups( + document_groups, + tokens_per_batch=tokens_per_batch, + ) logger.info(f"Split into {len(sub_batches)} sub-batches: {[len(b) for b in sub_batches]} items each") # Process each sub-batch - all_results = [] - for i, sub_batch in enumerate(sub_batches, 1): + all_results: list[list[str] | None] = [None] * len(contents) + for i, indexed_sub_batch in enumerate(sub_batches, 1): # Checkpoint: abort if the operation was deleted (bank was deleted) between sub-batches. if operation_id and not await self._check_op_alive(operation_id): if _retain_extraction_mode == "chunks": logger.info( - "[BATCH_RETAIN] multimodal operation %s cancelled; stopping after %s/%s sub-batches", - operation_id, + "[BATCH_RETAIN] multimodal operation cancelled; stopping after %s/%s sub-batches", i - 1, len(sub_batches), ) @@ -4019,10 +4198,9 @@ async def retain_batch_async( f"[BATCH_RETAIN] bank={bank_id} operation {operation_id} cancelled (bank deleted), " f"stopping after {i - 1}/{len(sub_batches)} sub-batches" ) - if return_usage: - return all_results, total_usage - return all_results + raise _RetainOperationCancelled("Tracked Retain was cancelled between logical document batches") + sub_batch = [item for _source_index, item in indexed_sub_batch] sub_batch_tokens = sum(count_tokens(item.get("content", "")) for item in sub_batch) logger.info( f"Processing sub-batch {i}/{len(sub_batches)}: {len(sub_batch)} items, {sub_batch_tokens:,} tokens" @@ -4032,8 +4210,8 @@ async def retain_batch_async( bank_id=bank_id, contents=sub_batch, request_context=request_context, - document_id=document_id, - is_first_batch=i == 1, # Only upsert on first batch + document_id=None, + is_first_batch=True, fact_type_override=fact_type_override, document_tags=document_tags, operation_id=operation_id, @@ -4043,7 +4221,10 @@ async def retain_batch_async( # webhook delivery row is committed atomically with the final retain data. outbox_callback=outbox_callback if i == len(sub_batches) else None, ) - all_results.extend(sub_results) + if len(sub_results) != len(indexed_sub_batch): + raise RuntimeError("Retain sub-batch result count does not match its submitted content count") + for (source_index, _item), item_result in zip(indexed_sub_batch, sub_results): + all_results[source_index] = item_result total_usage = total_usage + sub_usage if total_processed_content_tokens is None or sub_processed is None: total_processed_content_tokens = None @@ -4054,14 +4235,16 @@ async def retain_batch_async( logger.info( f"RETAIN_BATCH_ASYNC (chunked) COMPLETE: {len(all_results)} results from {len(contents)} contents in {total_time:.3f}s" ) - result = all_results + if any(item is None for item in all_results): + raise RuntimeError("Retain token batching did not produce a result for every submitted content item") + result = cast(list[list[str]], all_results) else: # Small batch - use internal method directly result, total_usage, total_processed_content_tokens = await self._retain_batch_async_internal( bank_id=bank_id, contents=contents, request_context=request_context, - document_id=document_id, + document_id=None, is_first_batch=True, fact_type_override=fact_type_override, document_tags=document_tags, @@ -4159,9 +4342,6 @@ async def _retain_batch_async_internal( See ``RetainResult.processed_content_tokens`` for the semantics of the third element. """ - # Use the new modular orchestrator - from .retain import orchestrator - backend = await self._get_backend() # Resolve bank-specific config for this operation @@ -4200,25 +4380,37 @@ async def _retain_batch_async_internal( # Create parent span for retain operation with create_operation_span("retain", bank_id): - return await orchestrator.retain_batch( + from .ingestion import ( + RetainExecutionContext, + RetainInvocation, + RetainPipelineService, + ) + + invocation = RetainInvocation( + bank_id=bank_id, + raw_contents=tuple(cast(RetainContentDict, dict(item)) for item in contents), + request_context=request_context, + batch_document_id=document_id, + is_first_batch=is_first_batch, + fact_type_override=fact_type_override, + document_tags=tuple(document_tags) if document_tags is not None else None, + operation_id=operation_id, + outbox_callback=outbox_callback, + strategy=strategy, + sanitize_log_identifiers=_retain_extraction_mode == "chunks", + ) + execution = RetainExecutionContext( pool=self._backend, embeddings_model=self.embeddings, llm_config=self._retain_llm_config.with_config(resolved_config), entity_resolver=self.entity_resolver, format_date_fn=self._format_readable_date, - bank_id=bank_id, - contents_dicts=contents, - document_id=document_id, - is_first_batch=is_first_batch, - fact_type_override=fact_type_override, - document_tags=document_tags, - config=resolved_config, - operation_id=operation_id, + resolved_config=resolved_config, schema=_current_schema.get(), - outbox_callback=outbox_callback, db_semaphore=self._put_semaphore, - sanitize_log_identifiers=_retain_extraction_mode == "chunks", ) + outcome = await RetainPipelineService().retain(invocation, execution) + return outcome.as_tuple() def recall( self, @@ -10629,18 +10821,18 @@ async def get_operation_status( SELECT operation_id, status, error_message, result_metadata FROM {fq_table("async_operations")} WHERE bank_id = $1 - AND result_metadata::jsonb @> $2::jsonb + AND (result_metadata->>'parent_operation_id')::uuid = $2 ORDER BY (result_metadata->>'sub_batch_index')::int """, bank_id, - json.dumps({"parent_operation_id": operation_id}), + op_uuid, ) # Build child operations list and check if parent status needs updating child_statuses = [] all_done = True any_failed = False - all_completed = True + any_cancelled = False for child_row in child_rows: raw_crm = child_row["result_metadata"] @@ -10657,16 +10849,16 @@ async def get_operation_status( } ) - if child_status not in ("completed", "failed"): + if child_status not in ("completed", "failed", "cancelled"): all_done = False if child_status == "failed": any_failed = True - if child_status != "completed": - all_completed = False + if child_status == "cancelled": + any_cancelled = True # Self-healing: if parent status is out of sync with children, update it - if all_done and api_status == "pending": - correct_status = "failed" if any_failed else "completed" + if child_rows and all_done and api_status in {"pending", "processing"}: + correct_status = "failed" if any_failed else "cancelled" if any_cancelled else "completed" logger.warning( f"Parent operation {operation_id} status out of sync (DB: pending, should be: {correct_status}). Fixing." ) @@ -10675,11 +10867,22 @@ async def get_operation_status( UPDATE {fq_table("async_operations")} SET status = $2, updated_at = NOW(), completed_at = NOW() WHERE operation_id = $1 + AND status IN ('pending', 'processing') """, op_uuid, correct_status, ) - api_status = correct_status + refreshed_parent = await conn.fetchrow( + f""" + SELECT status + FROM {fq_table("async_operations")} + WHERE operation_id = $1 AND bank_id = $2 + """, + op_uuid, + bank_id, + ) + if refreshed_parent is not None: + api_status = refreshed_parent["status"] return { "operation_id": operation_id, @@ -10742,11 +10945,47 @@ async def cancel_operation( async with acquire_with_retry(backend) as conn: async with conn.transaction(): - # Lock the operation so a worker claim cannot race the ledger - # transition and leave one side pending while the other is - # cancelled. + # Read the operation identity first. Batch parents lock their + # active children before locking the parent, matching the + # child->parent order used by completion aggregation and + # avoiding a cancellation/completion deadlock. + preflight = await conn.fetchrow( + f"""SELECT bank_id, status, operation_type, result_metadata + FROM {fq_table("async_operations")} + WHERE operation_id = $1 AND bank_id = $2""", + op_uuid, + bank_id, + ) + if not preflight: + raise ValueError(f"Operation {operation_id} not found for bank {bank_id}") + + raw_preflight_metadata = preflight["result_metadata"] + preflight_metadata = conn.parse_json(raw_preflight_metadata) or {} + is_batch_parent = ( + preflight["operation_type"] == "batch_retain" + and isinstance(preflight_metadata, dict) + and preflight_metadata.get("is_parent") is True + ) + if is_batch_parent: + await conn.fetch( + f""" + SELECT operation_id + FROM {fq_table("async_operations")} + WHERE bank_id = $1 + AND (result_metadata->>'parent_operation_id')::uuid = $2 + AND status IN ('pending', 'processing') + ORDER BY operation_id + FOR UPDATE + """, + bank_id, + op_uuid, + ) + + # Re-read under lock after any concurrent child completion. + # The status check below is therefore the cancellation CAS + # precondition, not a stale preflight observation. result = await conn.fetchrow( - f"""SELECT bank_id, status, operation_type, task_payload + f"""SELECT bank_id, status, operation_type, task_payload, result_metadata FROM {fq_table("async_operations")} WHERE operation_id = $1 AND bank_id = $2 FOR UPDATE""", @@ -10765,6 +11004,11 @@ async def cancel_operation( 409, ) + raw_result_metadata = result["result_metadata"] + result_metadata = conn.parse_json(raw_result_metadata) or {} + parent_operation_id = ( + result_metadata.get("parent_operation_id") if isinstance(result_metadata, dict) else None + ) task_payload = conn.parse_json(result["task_payload"]) if ( result["operation_type"] == "file_convert_retain" @@ -10798,6 +11042,23 @@ async def cancel_operation( elif command.status not in {"failed", "cancelled"}: raise LedgerConflictError("multimodal document command cannot be cancelled") + if is_batch_parent: + # Cancelling every still-active child in the same + # transaction prevents new claims, retries, and late + # completion writes from reviving the batch. Children that + # already completed remain immutable. + await conn.execute( + f""" + UPDATE {fq_table("async_operations")} + SET status = 'cancelled', updated_at = NOW() + WHERE bank_id = $1 + AND (result_metadata->>'parent_operation_id')::uuid = $2 + AND status IN ('pending', 'processing') + """, + bank_id, + op_uuid, + ) + await conn.execute( f"""UPDATE {fq_table("async_operations")} SET status = 'cancelled', updated_at = NOW() @@ -10805,6 +11066,11 @@ async def cancel_operation( op_uuid, bank_id, ) + if parent_operation_id: + # A directly cancelled child is a terminal sibling. Resolve + # its parent now so the aggregate cannot wait forever when + # this was the final outstanding child. + await self._maybe_update_parent_operation(operation_id, conn) return { "success": True, @@ -10835,6 +11101,42 @@ async def retry_operation( async with acquire_with_retry(backend) as conn: async with conn.transaction(): + preflight = await conn.fetchrow( + f"""SELECT bank_id, status, operation_type, result_metadata + FROM {fq_table("async_operations")} + WHERE operation_id = $1 AND bank_id = $2""", + op_uuid, + bank_id, + ) + if not preflight: + raise ValueError(f"Operation {operation_id} not found for bank {bank_id}") + + raw_preflight_metadata = preflight["result_metadata"] + preflight_metadata = conn.parse_json(raw_preflight_metadata) or {} + is_batch_parent = ( + preflight["operation_type"] == "batch_retain" + and isinstance(preflight_metadata, dict) + and preflight_metadata.get("is_parent") is True + ) + locked_children = [] + if is_batch_parent: + # Keep the same child->parent lock order as cancellation + # and aggregation. A retried parent must reactivate its + # failed/cancelled executable children, not merely reopen + # the payload-less aggregate row. + locked_children = await conn.fetch( + f""" + SELECT operation_id, status + FROM {fq_table("async_operations")} + WHERE bank_id = $1 + AND (result_metadata->>'parent_operation_id')::uuid = $2 + ORDER BY operation_id + FOR UPDATE + """, + bank_id, + op_uuid, + ) + row = await conn.fetchrow( f"""SELECT bank_id, status, operation_type, result_metadata FROM {fq_table("async_operations")} @@ -10849,6 +11151,9 @@ async def retry_operation( raw_metadata = row["result_metadata"] result_metadata = conn.parse_json(raw_metadata) or {} + parent_operation_id = ( + result_metadata.get("parent_operation_id") if isinstance(result_metadata, dict) else None + ) multimodal_metadata = result_metadata.get("multimodal") if isinstance(result_metadata, dict) else None retry_completed_multimodal_parent = ( row["status"] == "completed" @@ -10865,6 +11170,75 @@ async def retry_operation( 409, ) + if is_batch_parent: + retryable_children = [ + child for child in locked_children if child["status"] in {"failed", "cancelled"} + ] + if not retryable_children: + raise OperationValidationError( + f"Operation {operation_id} has no failed or cancelled child operations to retry", + 409, + ) + await conn.execute( + f""" + UPDATE {fq_table("async_operations")} + SET status = 'pending', + error_message = NULL, + completed_at = NULL, + next_retry_at = NULL, + worker_id = NULL, + claimed_at = NULL, + retry_count = 0, + updated_at = NOW() + WHERE bank_id = $1 + AND (result_metadata->>'parent_operation_id')::uuid = $2 + AND status IN ('failed', 'cancelled') + """, + bank_id, + op_uuid, + ) + + if parent_operation_id: + parent_row = await conn.fetchrow( + f""" + SELECT status + FROM {fq_table("async_operations")} + WHERE operation_id = $1 AND bank_id = $2 + FOR UPDATE + """, + uuid.UUID(parent_operation_id), + bank_id, + ) + if parent_row is None: + raise OperationValidationError( + f"Operation {operation_id} cannot be retried because its parent no longer exists", + 409, + ) + if parent_row["status"] == "completed": + raise OperationValidationError( + f"Operation {operation_id} cannot be retried because its parent is completed", + 409, + ) + if parent_row["status"] in {"failed", "cancelled"}: + await conn.execute( + f""" + UPDATE {fq_table("async_operations")} + SET status = 'pending', + error_message = NULL, + completed_at = NULL, + updated_at = NOW() + WHERE operation_id = $1 AND bank_id = $2 + AND status IN ('failed', 'cancelled') + """, + uuid.UUID(parent_operation_id), + bank_id, + ) + + retryable_statuses = ( + "('failed', 'cancelled', 'completed')" + if retry_completed_multimodal_parent + else "('failed', 'cancelled')" + ) await conn.execute( f""" UPDATE {fq_table("async_operations")} @@ -10876,9 +11250,11 @@ async def retry_operation( claimed_at = NULL, retry_count = 0, updated_at = NOW() - WHERE operation_id = $1 + WHERE operation_id = $1 AND bank_id = $2 + AND status IN {retryable_statuses} """, op_uuid, + bank_id, ) return { @@ -11225,8 +11601,8 @@ async def submit_async_retain( ) -> dict[str, Any]: """Submit a batch retain operation to run asynchronously. - For large batches (exceeding retain_batch_chars threshold), automatically splits - into smaller sub-batches and creates a parent operation that tracks all children. + Large requests are split only between complete logical documents and + create a parent operation that tracks all child batches. """ await self._authenticate_tenant(request_context) @@ -11243,44 +11619,32 @@ async def submit_async_retain( if result and result.contents is not None: contents = result.contents - # Validate no duplicate document_ids in the batch - # Having duplicate document_ids causes race conditions in document upserts during parallel processing - doc_ids = [item.get("document_id") for item in contents if item.get("document_id")] - if len(doc_ids) != len(set(doc_ids)): - from collections import Counter + for item in contents: + if item.get("update_mode") == "append" and not _is_retain_document_id(item.get("document_id")): + raise ValueError("update_mode='append' requires a document_id") - duplicates = [doc_id for doc_id, count in Counter(doc_ids).items() if count > 1] - raise ValueError( - f"Batch contains duplicate document_ids: {duplicates}. " - f"Each content item in a batch must have a unique document_id to avoid race conditions." - ) + prepared_contents, document_groups = _prepare_retain_document_groups( + cast(list[RetainContentDict], [dict(item) for item in contents]), + operation_id=None, + ) + contents = cast(list[dict[str, Any]], prepared_contents) # Calculate total token count and determine if we need to split total_tokens = sum(count_tokens(item.get("content", "")) for item in contents) config = get_config() tokens_per_batch = config.retain_batch_tokens - # Split into sub-batches based on token count - sub_batches = [] - current_batch = [] - current_batch_tokens = 0 - - for item in contents: - item_tokens = count_tokens(item.get("content", "")) - - # If adding this item would exceed the limit, start a new batch - # (unless current batch is empty - then we must include it even if it's large) - if current_batch and current_batch_tokens + item_tokens > tokens_per_batch: - sub_batches.append(current_batch) - current_batch = [item] - current_batch_tokens = item_tokens - else: - current_batch.append(item) - current_batch_tokens += item_tokens - - # Add the last batch - if current_batch: - sub_batches.append(current_batch) + # Split only between complete logical documents. The worker's Retain + # service handles oversized single documents using bounded write + # windows, so submission never changes document ownership semantics. + indexed_sub_batches = _split_retain_document_groups( + document_groups, + tokens_per_batch=tokens_per_batch, + ) + sub_batches = [ + [cast(dict[str, Any], item) for _source_index, item in indexed_sub_batch] + for indexed_sub_batch in indexed_sub_batches + ] # Log splitting info if we actually split if len(sub_batches) > 1: diff --git a/core/dataplane/hms_api/engine/retain/chunk_storage.py b/core/dataplane/hms_api/engine/retain/chunk_storage.py index 0c8eb5a..bef8492 100644 --- a/core/dataplane/hms_api/engine/retain/chunk_storage.py +++ b/core/dataplane/hms_api/engine/retain/chunk_storage.py @@ -13,6 +13,8 @@ logger = logging.getLogger(__name__) +_ORACLE_IN_CHUNK_SIZE = 900 + def compute_chunk_hash(chunk_text: str) -> str: """Compute SHA256 hash of chunk text for delta comparison.""" @@ -63,10 +65,15 @@ async def delete_chunks_by_ids(conn, chunk_ids: list[str]) -> None: """ if not chunk_ids: return - await conn.execute( - f"DELETE FROM {fq_table('chunks')} WHERE chunk_id = ANY($1::text[])", - chunk_ids, + backend_type = getattr(conn, "backend_type", "postgresql") + batch_size = ( + _ORACLE_IN_CHUNK_SIZE if isinstance(backend_type, str) and backend_type.lower() == "oracle" else len(chunk_ids) ) + for start in range(0, len(chunk_ids), batch_size): + await conn.execute( + f"DELETE FROM {fq_table('chunks')} WHERE chunk_id = ANY($1::text[])", + chunk_ids[start : start + batch_size], + ) async def store_chunks_batch( diff --git a/core/dataplane/hms_api/engine/retain/embedding_utils.py b/core/dataplane/hms_api/engine/retain/embedding_utils.py index 6978eef..22056a3 100644 --- a/core/dataplane/hms_api/engine/retain/embedding_utils.py +++ b/core/dataplane/hms_api/engine/retain/embedding_utils.py @@ -23,7 +23,7 @@ def generate_embedding(embeddings_backend, text: str) -> list[float]: embeddings = embeddings_backend.encode([text]) return embeddings[0] except Exception as e: - raise Exception(f"Failed to generate embedding: {str(e)}") + raise RuntimeError(f"Failed to generate embedding ({type(e).__name__}): {e}") from e async def generate_embeddings_batch(embeddings_backend, texts: list[str]) -> list[list[float]]: @@ -48,7 +48,7 @@ async def generate_embeddings_batch(embeddings_backend, texts: list[str]) -> lis texts, ) except Exception as e: - raise Exception(f"Failed to generate batch embeddings: {str(e)}") + raise RuntimeError(f"Failed to generate batch embeddings ({type(e).__name__}): {e}") from e # Guarantee 1:1 alignment with input texts. A silent length mismatch here # propagates downstream as zip() drops items, eventually surfacing as an diff --git a/core/dataplane/hms_api/engine/retain/entity_labels.py b/core/dataplane/hms_api/engine/retain/entity_labels.py index f104a3f..6e38f69 100644 --- a/core/dataplane/hms_api/engine/retain/entity_labels.py +++ b/core/dataplane/hms_api/engine/retain/entity_labels.py @@ -54,10 +54,11 @@ def parse_entity_labels(raw: dict | list | None) -> EntityLabelsConfig | None: Accepts: - None → returns None - - list → list of attribute dicts (each may use legacy free_values/multi_value or new type field) + - list → list of attribute dicts (each may use deprecated + free_values/multi_value fields or the current type field) - dict → {attributes: [...]} - Legacy migration (backward-compat): + Deprecated-field normalization: - free_values=True → type="text" - multi_value=True → type="multi-values" - neither / free_values=False → type="value" @@ -88,7 +89,7 @@ def parse_entity_labels(raw: dict | list | None) -> EntityLabelsConfig | None: def _migrate_label_group(raw: dict) -> dict: - """Migrate legacy free_values/multi_value fields to the new type field.""" + """Normalize deprecated free_values/multi_value fields to ``type``.""" if not isinstance(raw, dict) or "type" in raw: return raw patched = dict(raw) @@ -98,7 +99,7 @@ def _migrate_label_group(raw: dict) -> dict: patched["type"] = "multi-values" else: patched["type"] = "value" - # Remove legacy keys so Pydantic doesn't error on unknown fields + # Remove deprecated keys so Pydantic does not reject unknown fields. patched.pop("free_values", None) patched.pop("multi_value", None) return patched diff --git a/core/dataplane/hms_api/engine/retain/entity_processing.py b/core/dataplane/hms_api/engine/retain/entity_processing.py index a486c18..3c6e292 100644 --- a/core/dataplane/hms_api/engine/retain/entity_processing.py +++ b/core/dataplane/hms_api/engine/retain/entity_processing.py @@ -6,6 +6,7 @@ import logging +from ..entity_resolution_contracts import EntityResolutionReadPlan from . import link_utils from .types import EntityLink, ProcessedFact @@ -101,6 +102,59 @@ async def resolve_entities( ) +async def plan_entities( + entity_resolver, + conn, + bank_id: str, + fact_keys: list[str], + facts: list[ProcessedFact], + log_buffer: list[str] | None = None, + user_entities_per_content: dict[int, list[dict]] | None = None, + entity_labels: list | None = None, +) -> EntityResolutionReadPlan: + """Build an entity plan without performing a Phase-1 write.""" + + if not fact_keys or not facts: + return EntityResolutionReadPlan(bank_id=bank_id, occurrences=()) + if len(fact_keys) != len(facts): + raise ValueError(f"Mismatch between fact_keys ({len(fact_keys)}) and facts ({len(facts)})") + + fact_texts, fact_dates, entities_per_fact = _prepare_facts_for_entity_processing( + facts, + user_entities_per_content, + ) + flattened, _all_entities, entity_to_unit = link_utils._prepare_entities_for_resolution( + fact_keys, + fact_texts, + fact_dates, + entities_per_fact, + log_buffer, + ) + if not flattened: + return EntityResolutionReadPlan(bank_id=bank_id, occurrences=()) + + occurrence_keys_by_unit: dict[str, list[str]] = {} + for entity, (fact_key, local_index, _fact_date) in zip(flattened, entity_to_unit, strict=True): + occurrence_key = f"{fact_key}:entity:{local_index}" + entity["occurrence_key"] = occurrence_key + entity["unit_key"] = fact_key + entity["local_index"] = local_index + occurrence_keys_by_unit.setdefault(fact_key, []).append(occurrence_key) + for entity in flattened: + entity["nearby_occurrence_keys"] = [ + key for key in occurrence_keys_by_unit[entity["unit_key"]] if key != entity["occurrence_key"] + ] + + return await entity_resolver.plan_entities_batch( + bank_id=bank_id, + entities_data=flattened, + context="", + unit_event_date=None, + conn=conn, + entity_labels=entity_labels, + ) + + async def build_entity_links( entity_resolver, conn, diff --git a/core/dataplane/hms_api/engine/retain/fact_extraction.py b/core/dataplane/hms_api/engine/retain/fact_extraction.py index 3c6541e..2a42eb5 100644 --- a/core/dataplane/hms_api/engine/retain/fact_extraction.py +++ b/core/dataplane/hms_api/engine/retain/fact_extraction.py @@ -9,7 +9,7 @@ import json import logging import re -from datetime import datetime, timedelta +from datetime import UTC, datetime, timedelta from typing import Literal, cast from pydantic import BaseModel, ConfigDict, Field, create_model, field_validator @@ -27,6 +27,17 @@ ) +def parse_datetime_flexible(value: object) -> datetime: + """Parse a datetime object or ISO string into a timezone-aware value.""" + + if isinstance(value, datetime): + return value.replace(tzinfo=UTC) if value.tzinfo is None else value + if isinstance(value, str): + parsed = datetime.fromisoformat(value.replace("Z", "+00:00")) + return parsed.replace(tzinfo=UTC) if parsed.tzinfo is None else parsed + raise TypeError(f"Expected datetime or string, got {type(value).__name__}") + + def _extract_map_entities( entity_obj: dict, fields: dict[str, MapField], @@ -113,6 +124,44 @@ def _sanitize_text(text: str | None) -> str | None: return sanitize_llm_output(text) +def _sanitize_model_response_strings(response: object, *, scope: str) -> object: + """Recursively sanitize string values at the model-response boundary. + + Dictionary keys are intentionally left untouched. Retain consumes several + nested response fields (including entities and structured labels), so doing + this once before field parsing keeps every downstream projection consistent. + """ + changed_values = 0 + removed_characters = 0 + + def sanitize_value(value: object) -> object: + nonlocal changed_values, removed_characters + + if isinstance(value, str): + sanitized = cast(str, sanitize_llm_output(value)) + if sanitized != value: + changed_values += 1 + removed_characters += len(value) - len(sanitized) + return sanitized + if isinstance(value, dict): + return {key: sanitize_value(item) for key, item in value.items()} + if isinstance(value, list): + return [sanitize_value(item) for item in value] + if isinstance(value, tuple): + return tuple(sanitize_value(item) for item in value) + return value + + sanitized_response = sanitize_value(response) + if changed_values: + logging.getLogger(__name__).warning( + "Sanitized model-response string values: scope=%s changed_values=%d removed_characters=%d", + scope, + changed_values, + removed_characters, + ) + return sanitized_response + + class Entity(BaseModel): """An entity extracted from text.""" @@ -156,6 +205,67 @@ class CausalRelation(BaseModel): ) +def _remap_causal_relations( + raw_relations: object, + *, + source_raw_index: int, + raw_to_retained_index: dict[int, int], +) -> list[CausalRelation]: + """Validate relations and translate raw response indices to retained positions. + + LLM responses can contain malformed facts that are skipped by the lenient + parser. A relation is safe to retain only when its target is an earlier raw + fact *and* that fact was successfully validated and retained. Never infer a + target from the compacted list position: doing so can silently connect the + relation to a different fact. + """ + if not isinstance(raw_relations, list): + return [] + + relation_logger = logging.getLogger(__name__) + validated_relations: list[CausalRelation] = [] + for relation in raw_relations: + if not isinstance(relation, dict): + continue + + target_raw_index = relation.get("target_index") + relation_type = relation.get("relation_type") + if ( + not isinstance(target_raw_index, int) + or isinstance(target_raw_index, bool) + or target_raw_index < 0 + or target_raw_index >= source_raw_index + or relation_type is None + ): + relation_logger.debug( + "Invalid target_index %r for raw fact %s; skipping causal relation", + target_raw_index, + source_raw_index, + ) + continue + + target_retained_index = raw_to_retained_index.get(target_raw_index) + if target_retained_index is None: + relation_logger.debug( + "Causal target raw fact %s was not retained for raw fact %s; skipping relation", + target_raw_index, + source_raw_index, + ) + continue + + try: + validated_relations.append( + CausalRelation( + target_fact_index=target_retained_index, + relation_type=relation_type, + ) + ) + except Exception as exc: + relation_logger.debug("Invalid causal relation %r: %s", relation, exc) + + return validated_relations + + class FactCausalRelation(BaseModel): """ Causal relationship from this fact to a PREVIOUS fact (embedded in each fact). @@ -647,7 +757,7 @@ def _chunk_conversation(turns: list[dict], max_chars: int) -> list[str]: ) -# Verbose extraction prompt - detailed, comprehensive facts (legacy mode) +# Verbose extraction prompt for detailed, comprehensive facts. VERBOSE_FACT_EXTRACTION_PROMPT = """Extract facts from text into structured format with FIVE required dimensions - BE EXTREMELY DETAILED. LANGUAGE: MANDATORY — Detect the language of the input text and produce ALL output in that EXACT same language. You are STRICTLY FORBIDDEN from translating or switching to any other language. Every single word of your output must be in the same language as the input. Do NOT output in a different language under any circumstance. @@ -978,7 +1088,10 @@ def _build_extraction_prompt_and_schema(config) -> tuple[str, type]: ), **dynamic_fields, ) - DynamicResponse = create_model("LabelsResponse", facts=(list[DynamicFact], ...)) # type: ignore[valid-type] + DynamicResponse = create_model( + "LabelsResponse", + facts=(list[DynamicFact], ...), # type: ignore[valid-type] # ty: ignore[invalid-type-form] + ) response_schema = DynamicResponse return prompt, response_schema @@ -994,8 +1107,6 @@ def _build_user_message( agent_name: str | None = None, ) -> str: """Build user message for fact extraction.""" - from .orchestrator import parse_datetime_flexible - sanitized_chunk = _sanitize_text(chunk) sanitized_context = _sanitize_text(context) if context else "none" @@ -1061,7 +1172,7 @@ async def _extract_facts_from_chunk( config, agent_name: str = None, metadata: dict[str, str] | None = None, -) -> tuple[list[dict[str, str]], TokenUsage]: +) -> tuple[list[Fact], TokenUsage]: """ Extract facts from a single chunk (internal helper for parallel processing). @@ -1116,9 +1227,14 @@ async def _extract_facts_from_chunk( return_usage=True, ) usage = usage + call_usage # Aggregate usage across retries + extraction_response_json = _sanitize_model_response_strings( + extraction_response_json, + scope="retain_extract_facts_sync", + ) # Lenient parsing of facts from raw JSON chunk_facts = [] + raw_to_retained_index: dict[int, int] = {} has_malformed_facts = False # Handle malformed LLM responses @@ -1237,9 +1353,14 @@ def get_value(field_name): # Validate and normalize each entity for ent in entities: if isinstance(ent, str): + if not ent.strip(): + continue # Normalize string to Entity object validated_entities.append(Entity(text=ent)) elif isinstance(ent, dict) and "text" in ent: + entity_text = ent.get("text") + if isinstance(entity_text, str) and not entity_text.strip(): + continue try: validated_entities.append(Entity.model_validate(ent)) except Exception as e: @@ -1302,35 +1423,11 @@ def get_value(field_name): # Add per-fact causal relations (only if enabled in config) if extract_causal_links: - validated_relations = [] - causal_relations_raw = get_value("causal_relations") - if causal_relations_raw: - for rel in causal_relations_raw: - if not isinstance(rel, dict): - continue - # New schema uses target_index - target_idx = rel.get("target_index") - relation_type = rel.get("relation_type") - - if target_idx is None or relation_type is None: - continue - - # Validate: target_index must be < current fact index - if target_idx < 0 or target_idx >= i: - logger.debug( - f"Invalid target_index {target_idx} for fact {i} (must be 0 to {i - 1}). Skipping." - ) - continue - - try: - validated_relations.append( - CausalRelation( - target_fact_index=target_idx, - relation_type=relation_type, - ) - ) - except Exception as e: - logger.debug(f"Invalid causal relation {rel}: {e}") + validated_relations = _remap_causal_relations( + get_value("causal_relations"), + source_raw_index=i, + raw_to_retained_index=raw_to_retained_index, + ) if validated_relations: fact_data["causal_relations"] = validated_relations @@ -1342,6 +1439,7 @@ def get_value(field_name): # Build Fact model instance try: fact = Fact(fact=combined_text, fact_type=fact_type, **fact_data) + raw_to_retained_index[i] = len(chunk_facts) chunk_facts.append(fact) except Exception as e: logger.error(f"Failed to create Fact model for fact {i}: {e}") @@ -1403,7 +1501,7 @@ async def _extract_facts_with_auto_split( config, agent_name: str = None, metadata: dict[str, str] | None = None, -) -> tuple[list[dict[str, str]], TokenUsage]: +) -> tuple[list[Fact], TokenUsage]: """ Extract facts from a chunk with automatic splitting if output exceeds token limits. @@ -1506,6 +1604,10 @@ async def _extract_facts_with_auto_split( all_facts = [] total_usage = TokenUsage() for sub_facts, sub_usage in sub_results: + subchunk_fact_start_idx = len(all_facts) + for fact in sub_facts: + for relation in fact.causal_relations or []: + relation.target_fact_index += subchunk_fact_start_idx all_facts.extend(sub_facts) total_usage = total_usage + sub_usage @@ -1587,7 +1689,7 @@ async def extract_facts_from_text( total_usage = TokenUsage() failed_chunks = [] for i, (chunk, result) in enumerate(zip(chunks, chunk_results)): - if isinstance(result, Exception): + if isinstance(result, BaseException): failed_chunks.append((i, result)) continue chunk_facts, chunk_usage = result @@ -1599,11 +1701,19 @@ async def extract_facts_from_text( # Fail the entire retain — partial extraction is not acceptable. # All successfully extracted facts are discarded because the transaction # hasn't committed yet. The worker poller will retry the entire task. - failed_summary = ", ".join(f"chunk {idx}: {type(err).__name__}" for idx, err in failed_chunks[:5]) + def summarize_failure(idx: int, err: BaseException) -> str: + detail = " ".join(str(err).split()) + if len(detail) > 300: + detail = f"{detail[:300]}...TRUNCATED" + suffix = f": {detail}" if detail else "" + return f"chunk {idx}: {type(err).__name__}{suffix}" + + failed_summary = ", ".join(summarize_failure(idx, err) for idx, err in failed_chunks[:5]) + first_error = failed_chunks[0][1] raise RuntimeError( f"Fact extraction failed: {len(failed_chunks)}/{len(chunks)} chunks failed. " f"First failures: {failed_summary}" - ) + ) from first_error return all_facts, chunk_metadata, total_usage @@ -1850,10 +1960,15 @@ async def extract_facts_from_contents_batch_api( ) ) continue + extraction_response_json = _sanitize_model_response_strings( + extraction_response_json, + scope=f"retain_extract_facts_batch:{custom_id}", + ) # Parse facts (reuse existing logic from _extract_facts_from_chunk) raw_facts = extraction_response_json.get("facts", []) chunk_facts = [] + raw_to_retained_index: dict[int, int] = {} for i, llm_fact in enumerate(raw_facts): if not isinstance(llm_fact, dict): @@ -1923,8 +2038,13 @@ def get_value(field_name): if entities: for ent in entities: if isinstance(ent, str): + if not ent.strip(): + continue validated_entities.append(Entity(text=ent)) elif isinstance(ent, dict) and "text" in ent: + entity_text = ent.get("text") + if isinstance(entity_text, str) and not entity_text.strip(): + continue try: validated_entities.append(Entity.model_validate(ent)) except Exception: @@ -1984,26 +2104,11 @@ def get_value(field_name): # Causal relations if extract_causal_links: - validated_relations = [] - causal_relations_raw = get_value("causal_relations") - if causal_relations_raw: - for rel in causal_relations_raw: - if not isinstance(rel, dict): - continue - target_idx = rel.get("target_index") - relation_type = rel.get("relation_type") - - if target_idx is None or relation_type is None: - continue - if target_idx < 0 or target_idx >= i: - continue - - try: - validated_relations.append( - CausalRelation(target_fact_index=target_idx, relation_type=relation_type) - ) - except Exception: - pass + validated_relations = _remap_causal_relations( + get_value("causal_relations"), + source_raw_index=i, + raw_to_retained_index=raw_to_retained_index, + ) if validated_relations: fact_data["causal_relations"] = validated_relations @@ -2014,6 +2119,7 @@ def get_value(field_name): try: fact = Fact(fact=combined_text, fact_type=fact_type, **fact_data) + raw_to_retained_index[i] = len(chunk_facts) chunk_facts.append(fact) except Exception as e: logger.error(f"Failed to create Fact model for fact {i}: {e}") @@ -2054,6 +2160,7 @@ def get_value(field_name): for chunk_meta, chunk_facts in facts_by_chunk: content = contents[chunk_meta.content_index] + chunk_fact_start_idx = global_fact_idx for fact_from_llm in chunk_facts: extracted_fact = ExtractedFactType( @@ -2062,7 +2169,10 @@ def get_value(field_name): entities=[e.text for e in (fact_from_llm.entities or [])], occurred_start=_parse_datetime(fact_from_llm.occurred_start) if fact_from_llm.occurred_start else None, occurred_end=_parse_datetime(fact_from_llm.occurred_end) if fact_from_llm.occurred_end else None, - causal_relations=_convert_causal_relations(fact_from_llm.causal_relations or [], global_fact_idx), + causal_relations=_convert_causal_relations( + fact_from_llm.causal_relations or [], + chunk_fact_start_idx, + ), content_index=chunk_meta.content_index, chunk_index=chunk_meta.chunk_index, context=content.context, @@ -2194,8 +2304,10 @@ async def extract_facts_from_contents( ) fact_extraction_tasks.append(task) - # Step 2: Wait for all fact extractions to complete. - # Use return_exceptions=True so one content item failure doesn't discard the rest. + # Step 2: Wait for all fact extractions to complete. Collect every result so + # sibling tasks are allowed to finish, then fail the whole extraction below + # if any content failed. A failed content has no trustworthy chunk metadata + # and therefore must never be represented as a successful zero-fact result. all_fact_results = await asyncio.gather(*fact_extraction_tasks, return_exceptions=True) # Step 3: Flatten and convert to typed objects @@ -2206,14 +2318,17 @@ async def extract_facts_from_contents( global_chunk_idx = 0 global_fact_idx = 0 - # Filter out failed content items + # A genuine zero-fact extraction still returns one ``(chunk, 0)`` entry for + # every input chunk. Historically this aggregation boundary replaced a + # per-content exception with ``([], [], TokenUsage())``; downstream code + # could not distinguish that silent data loss from a valid zero-fact + # document. Propagate the original exception, in deterministic input + # order, before projecting any partial result. valid_results = [] - for content, result in zip(contents, all_fact_results): - if isinstance(result, Exception): - logger.warning(f"Content extraction failed (skipping): {type(result).__name__}: {result}") - valid_results.append((content, ([], [], TokenUsage()))) - else: - valid_results.append((content, result)) + for content, result in zip(contents, all_fact_results, strict=True): + if isinstance(result, BaseException): + raise result + valid_results.append((content, result)) for content_index, (content, (facts_from_llm, chunks_from_llm, content_usage)) in enumerate(valid_results): total_usage = total_usage + content_usage @@ -2234,6 +2349,7 @@ async def extract_facts_from_contents( fact_idx_in_content = 0 for chunk_idx_in_content, (chunk_text, chunk_fact_count) in enumerate(chunks_from_llm): chunk_global_idx = chunk_start_idx + chunk_idx_in_content + chunk_fact_start_idx = global_fact_idx for _ in range(chunk_fact_count): if fact_idx_in_content < len(facts_from_llm): @@ -2253,7 +2369,8 @@ async def extract_facts_from_contents( if fact_from_llm.occurred_end else None, causal_relations=_convert_causal_relations( - fact_from_llm.causal_relations or [], global_fact_idx + fact_from_llm.causal_relations or [], + chunk_fact_start_idx, ), content_index=content_index, chunk_index=chunk_global_idx, @@ -2319,17 +2436,17 @@ def _parse_datetime(date_str: str): return None -def _convert_causal_relations(relations_from_llm, fact_start_idx: int) -> list[CausalRelationType]: +def _convert_causal_relations(relations_from_llm, chunk_fact_start_idx: int) -> list[CausalRelationType]: """ Convert causal relations from LLM format to ExtractedFact format. - Adjusts target_fact_index from content-relative to global indices. + Adjusts target_fact_index from chunk-relative to global indices. """ causal_relations = [] for rel in relations_from_llm: causal_relation = CausalRelationType( relation_type=rel.relation_type, - target_fact_index=fact_start_idx + rel.target_fact_index, + target_fact_index=chunk_fact_start_idx + rel.target_fact_index, ) causal_relations.append(causal_relation) return causal_relations @@ -2347,8 +2464,6 @@ def _add_temporal_offsets(facts: list[ExtractedFactType], contents: list[RetainC Modifies facts in place. """ - from .orchestrator import parse_datetime_flexible - for i, fact in enumerate(facts): # Use absolute position across all facts to ensure uniqueness across different contents offset = timedelta(seconds=i * SECONDS_PER_FACT) diff --git a/core/dataplane/hms_api/engine/retain/fact_storage.py b/core/dataplane/hms_api/engine/retain/fact_storage.py index 48b7495..5412e05 100644 --- a/core/dataplane/hms_api/engine/retain/fact_storage.py +++ b/core/dataplane/hms_api/engine/retain/fact_storage.py @@ -4,10 +4,13 @@ Handles insertion of facts into the database. """ +from __future__ import annotations + import json import logging import uuid from datetime import datetime +from typing import Any, Protocol from ...config import get_config from ..memory_engine import fq_table @@ -18,6 +21,12 @@ logger = logging.getLogger(__name__) +class _IdentifierLogSanitizer(Protocol): + """Minimal request-local identifier renderer used by trusted Retain callers.""" + + def identifier(self, value: Any) -> str: ... + + async def get_document_content( conn, bank_id: str, @@ -169,8 +178,9 @@ async def ensure_bank_exists(conn, bank_id: str, ops=None) -> None: async def delete_stale_observations_for_memories( conn, bank_id: str, - fact_ids: "list[str | uuid.UUID]", + fact_ids: list[str | uuid.UUID], ops=None, + log_sanitizer: _IdentifierLogSanitizer | None = None, ) -> int: """Delete observations whose source memories are about to be removed. @@ -258,9 +268,10 @@ async def delete_stale_observations_for_memories( remaining_source_ids, ) + log_bank_id = log_sanitizer.identifier(bank_id) if log_sanitizer is not None else bank_id logger.info( f"[OBSERVATIONS] Deleted {len(obs_ids)} observations, reset {len(remaining_source_ids)} " - f"source memories for re-consolidation in bank {bank_id}" + f"source memories for re-consolidation in bank {log_bank_id}" ) return len(obs_ids) @@ -274,6 +285,7 @@ async def handle_document_tracking( retain_params: dict | None = None, document_tags: list[str] | None = None, ops=None, + log_sanitizer: _IdentifierLogSanitizer | None = None, ) -> None: """ Handle document tracking in the database (full-replace mode). @@ -333,10 +345,17 @@ async def handle_document_tracking( ) existing_unit_ids = [row["id"] for row in existing_unit_rows] if existing_unit_ids: - invalidated = await delete_stale_observations_for_memories(conn, bank_id, existing_unit_ids, ops=ops) + invalidated = await delete_stale_observations_for_memories( + conn, + bank_id, + existing_unit_ids, + ops=ops, + log_sanitizer=log_sanitizer, + ) if invalidated: + log_document_id = log_sanitizer.identifier(document_id) if log_sanitizer is not None else document_id logger.info( - f"[RETAIN] Document {document_id} re-ingested: invalidated " + f"[RETAIN] Document {log_document_id} re-ingested: invalidated " f"{invalidated} observation(s) derived from {len(existing_unit_ids)} outgoing memory_units" ) # Explicitly delete memory_units by document_id BEFORE deleting the diff --git a/core/dataplane/hms_api/engine/retain/link_utils.py b/core/dataplane/hms_api/engine/retain/link_utils.py index 5f67de9..4eaa151 100644 --- a/core/dataplane/hms_api/engine/retain/link_utils.py +++ b/core/dataplane/hms_api/engine/retain/link_utils.py @@ -5,11 +5,15 @@ import logging import time from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING from uuid import UUID from ..memory_engine import fq_table from .types import EntityLink +if TYPE_CHECKING: + from ..db.ops import DataAccessOps + logger = logging.getLogger(__name__) # Sentinel UUID used in the unique index to represent NULL entity_id @@ -57,7 +61,7 @@ async def _bulk_insert_links( bank_id: str = "", chunk_size: int = 5000, skip_exists_check: bool = False, - ops=None, + ops: "DataAccessOps | None" = None, ) -> None: """Bulk-insert links using sorted INSERT FROM unnest(). @@ -77,6 +81,8 @@ async def _bulk_insert_links( """ if not links: return + if ops is None: + raise ValueError("bulk link insertion requires backend database operations") # Sort by (from_unit_id, to_unit_id) to guarantee consistent lock ordering # across concurrent transactions — prevents deadlocks. @@ -376,7 +382,7 @@ async def build_entity_links_from_resolved( unit_to_entity_ids: dict[str, list[str]], log_buffer: list[str] = None, skip_unit_entities_insert: bool = False, - ops=None, + ops: "DataAccessOps | None" = None, ) -> list["EntityLink"]: """ Build entity links between units that share entities. @@ -429,6 +435,8 @@ async def build_entity_links_from_resolved( entity_to_units = {} if all_entity_ids: + if ops is None: + raise ValueError("entity link construction requires backend database operations") query_start = time.time() import uuid @@ -502,7 +510,7 @@ async def create_temporal_links_batch_per_fact( unit_ids: list[str], time_window_hours: int = 24, log_buffer: list[str] = None, - ops=None, + ops: "DataAccessOps | None" = None, ) -> int: """ Create temporal links for multiple units, each with their own event_date. @@ -522,6 +530,8 @@ async def create_temporal_links_batch_per_fact( """ if not unit_ids: return 0 + if ops is None: + raise ValueError("temporal link construction requires backend database operations") try: import time as time_mod @@ -645,13 +655,15 @@ async def compute_semantic_links_ann( top_k: int = 50, threshold: float = 0.7, log_buffer: list[str] = None, + read_only: bool = False, ) -> list[tuple]: """ Phase 1: ANN search for semantic neighbors among existing units. - Runs on a separate connection OUTSIDE the write transaction to avoid - holding locks during expensive HNSW index probes. Uses a temp table + - LATERAL join to batch all probes in a single query. + Runs before the core write transaction to avoid holding write locks during + expensive HNSW index probes. The default path uses a temporary table; + ``read_only=True`` supplies seeds through a typed-array CTE so PostgreSQL + can enforce a read-only planning transaction. Queries are split by fact_type so PostgreSQL uses the per-bank partial HNSW indexes (idx_mu_emb_worl_*, idx_mu_emb_expr_*). Without the @@ -667,6 +679,8 @@ async def compute_semantic_links_ann( top_k: Max neighbors per unit threshold: Minimum cosine similarity log_buffer: Optional logging buffer + read_only: Use a DDL-free query suitable for an existing read-only + transaction. The caller owns that transaction. Returns: List of (from_id, to_id, "semantic", similarity, None) tuples @@ -703,34 +717,31 @@ async def compute_semantic_links_ann( # manually drop the temp table or reset hnsw.ef_search — the transaction # end handles both. rows: list = [] - async with conn.transaction(): - # Transaction-local ef_search bounds ANN work during semantic link - # creation. SET LOCAL auto-reverts at commit, so the setting does not - # leak to later recall queries through the connection pool. + if read_only: + # Phase 1 runs inside a database-enforced READ ONLY transaction. A + # temporary table is still a CREATE and PostgreSQL + # rejects it there, so pass seeds through typed arrays instead. The + # expensive operation remains the same indexed LATERAL HNSW probe. await conn.execute("SET LOCAL hnsw.ef_search = 60") - - t_setup = time_mod.time() - await conn.execute("CREATE TEMP TABLE _ann_seeds (unit_id text, emb_text text, fact_type text) ON COMMIT DROP") - - records = [ - (uid, emb if isinstance(emb, str) else str(emb), ft) - for uid, emb, ft in zip(unit_ids, embeddings, fact_types) - ] - await conn.copy_records_to_table("_ann_seeds", records=records, columns=["unit_id", "emb_text", "fact_type"]) - logger.debug(f"[ANN] Temp table setup: {time_mod.time() - t_setup:.3f}s ({len(records)} seeds)") - - # Run one ANN query per fact_type so each uses the right HNSW index. active_types = set(fact_types) for fact_type in active_types: + typed_seeds = [ + (uid, emb if isinstance(emb, str) else str(emb)) + for uid, emb, seed_type in zip(unit_ids, embeddings, fact_types) + if seed_type == fact_type + ] + seed_ids = [seed[0] for seed in typed_seeds] + seed_embeddings = [seed[1] for seed in typed_seeds] t_query = time_mod.time() - seed_count = sum(1 for ft in fact_types if ft == fact_type) - logger.debug(f"[ANN] Querying fact_type={fact_type}: {seed_count} seeds") ft_rows = await conn.fetch( f""" + WITH seeds(unit_id, emb_text) AS ( + SELECT * FROM unnest($4::text[], $5::text[]) + ) SELECT s.unit_id AS from_id, n.id::text AS to_id, n.similarity - FROM _ann_seeds s + FROM seeds s CROSS JOIN LATERAL ( SELECT mu.id, 1 - (mu.embedding <=> s.emb_text::vector) AS similarity @@ -741,14 +752,75 @@ async def compute_semantic_links_ann( ORDER BY mu.embedding <=> s.emb_text::vector LIMIT $3 ) n - WHERE s.fact_type = $2 """, bank_id, fact_type, top_k, + seed_ids, + seed_embeddings, + ) + logger.debug( + "[ANN] read-only fact_type=%s: %d rows in %.3fs", + fact_type, + len(ft_rows), + time_mod.time() - t_query, ) - logger.debug(f"[ANN] fact_type={fact_type}: {len(ft_rows)} rows in {time_mod.time() - t_query:.3f}s") rows.extend(ft_rows) + else: + async with conn.transaction(): + # Transaction-local ef_search. Default 400 is tuned for recall + # precision but at 164k units each HNSW probe takes 94ms. + # ef_search=60 gives 2.7ms per probe (35x faster) with sufficient + # accuracy for top-50 semantic link creation. SET LOCAL + # auto-reverts at commit, so pooled recall queries are unaffected. + await conn.execute("SET LOCAL hnsw.ef_search = 60") + + t_setup = time_mod.time() + await conn.execute( + "CREATE TEMP TABLE _ann_seeds (unit_id text, emb_text text, fact_type text) ON COMMIT DROP" + ) + + records = [ + (uid, emb if isinstance(emb, str) else str(emb), ft) + for uid, emb, ft in zip(unit_ids, embeddings, fact_types) + ] + await conn.copy_records_to_table( + "_ann_seeds", + records=records, + columns=["unit_id", "emb_text", "fact_type"], + ) + logger.debug(f"[ANN] Temp table setup: {time_mod.time() - t_setup:.3f}s ({len(records)} seeds)") + + # Run one ANN query per fact_type so each uses the right HNSW index. + active_types = set(fact_types) + for fact_type in active_types: + t_query = time_mod.time() + seed_count = sum(1 for ft in fact_types if ft == fact_type) + logger.debug(f"[ANN] Querying fact_type={fact_type}: {seed_count} seeds") + ft_rows = await conn.fetch( + f""" + SELECT s.unit_id AS from_id, + n.id::text AS to_id, + n.similarity + FROM _ann_seeds s + CROSS JOIN LATERAL ( + SELECT mu.id, + 1 - (mu.embedding <=> s.emb_text::vector) AS similarity + FROM {fq_table("memory_units")} mu + WHERE mu.bank_id = $1 + AND mu.fact_type = $2 + AND mu.embedding IS NOT NULL + ORDER BY mu.embedding <=> s.emb_text::vector + LIMIT $3 + ) n + WHERE s.fact_type = $2 + """, + bank_id, + fact_type, + top_k, + ) + logger.debug(f"[ANN] fact_type={fact_type}: {len(ft_rows)} rows in {time_mod.time() - t_query:.3f}s") + rows.extend(ft_rows) # Transaction commits here. _ann_seeds is dropped (ON COMMIT DROP). # hnsw.ef_search reverts (SET LOCAL). diff --git a/core/dataplane/hms_api/engine/retain/orchestrator.py b/core/dataplane/hms_api/engine/retain/orchestrator.py deleted file mode 100644 index 5a2f073..0000000 --- a/core/dataplane/hms_api/engine/retain/orchestrator.py +++ /dev/null @@ -1,2209 +0,0 @@ -""" -Main orchestrator for the retain pipeline. - -Coordinates all retain pipeline modules to store memories efficiently. -""" - -import asyncio -import hashlib -import json -import logging -import time -import uuid -from collections.abc import Awaitable, Callable -from datetime import UTC, datetime -from typing import Any - -from ...worker.stage import set_stage -from ..db.base import DatabaseBackend -from ..db_utils import acquire_with_retry -from ..embedding_fingerprint import embedding_model_version, ensure_bank_embedding_fingerprint -from ..memory_engine import count_tokens, fq_table -from . import bank_utils - - -def utcnow(): - """Get current UTC time.""" - return datetime.now(UTC) - - -def _merge_processed_content_tokens(a: int | None, b: int | None) -> int | None: - """Combine the processed-content-tokens signal across sub-results. - - Semantics (see RetainResult.processed_content_tokens): - * None means "this part of the retain did not go through chunk-level - dedup" — i.e. the entire submitted payload was processed. If any - sub-result is None, the aggregate is None so callers conservatively - bill the full content. - * Otherwise, accumulate the int values. - """ - if a is None or b is None: - return None - return a + b - - -def _count_delta_content_tokens(delta_contents: list["RetainContent"]) -> int: - """Sum content + context tokens across the chunk items that were - actually fed into the extraction pipeline on a partial-delta retain. - """ - total = 0 - for c in delta_contents: - total += count_tokens(c.content or "") - total += count_tokens(c.context or "") - return total - - -def parse_datetime_flexible(value: Any) -> datetime: - """ - Parse a datetime value that could be either a datetime object or an ISO string. - - This handles datetime values from both direct Python calls and deserialized JSON - (where datetime objects are serialized as ISO strings). - - Args: - value: Either a datetime object or an ISO format string - - Returns: - datetime object (timezone-aware) - - Raises: - TypeError: If value is neither datetime nor string - ValueError: If string is not a valid ISO datetime - """ - if isinstance(value, datetime): - # Ensure timezone-aware - if value.tzinfo is None: - return value.replace(tzinfo=UTC) - return value - elif isinstance(value, str): - # Parse ISO format string (handles both 'Z' and '+00:00' timezone formats) - dt = datetime.fromisoformat(value.replace("Z", "+00:00")) - # Ensure timezone-aware - if dt.tzinfo is None: - return dt.replace(tzinfo=UTC) - return dt - else: - raise TypeError(f"Expected datetime or string, got {type(value).__name__}") - - -import asyncpg - -from ..response_models import TokenUsage -from . import ( - chunk_storage, - embedding_processing, - entity_processing, - fact_extraction, - fact_storage, - link_creation, -) -from .types import ( - ChunkMetadata, - EntityResolutionResult, - Phase1Result, - Phase3Context, - ProcessedFact, - RetainContent, - RetainContentDict, -) - -logger = logging.getLogger(__name__) - - -def _embedding_model_version(embeddings_model) -> str: - return embedding_model_version(embeddings_model) - - -def _build_retain_params(contents_dicts, document_tags=None, doc_contents=None): - """Build retain_params and merged_tags from content dicts.""" - if doc_contents is not None: - # Per-document mode: doc_contents is list of (idx, content_dict) - items = [item for _, item in doc_contents] - else: - items = contents_dicts - - all_tags = set(document_tags or []) - for item in items: - item_tags = item.get("tags", []) or [] - all_tags.update(item_tags) - merged_tags = list(all_tags) - - retain_params = {} - if items: - first_item = items[0] - if first_item.get("context"): - retain_params["context"] = first_item["context"] - if first_item.get("event_date"): - retain_params["event_date"] = ( - first_item["event_date"].isoformat() - if hasattr(first_item["event_date"], "isoformat") - else str(first_item["event_date"]) - ) - if first_item.get("metadata"): - retain_params["metadata"] = first_item["metadata"] - - return retain_params, merged_tags - - -async def _pre_resolve_phase1( - pool: Any, - entity_resolver, - bank_id: str, - contents: list[RetainContent], - processed_facts: list[ProcessedFact], - config, - log_buffer: list[str], - skip_semantic_ann: bool = False, -) -> Phase1Result: - """ - Phase 1: Run expensive read-heavy operations on a separate connection - OUTSIDE the write transaction. - - - Entity resolution: trigram GIN scan + co-occurrence fetch + scoring - - Semantic ANN: HNSW index probes to find similar existing units - - Running these outside the transaction avoids holding row locks during - slow reads, eliminating TimeoutErrors under concurrent load. - """ - set_stage("retain.phase1.resolve") - from .link_utils import compute_semantic_links_ann - - user_entities_per_content = {idx: content.entities for idx, content in enumerate(contents) if content.entities} - - # Use placeholder unit_ids for grouping during resolution. The actual - # unit_ids are created later by insert_facts_batch inside the transaction, - # but entity resolution and ANN search only need them as grouping keys. - placeholder_unit_ids = [str(i) for i in range(len(processed_facts))] - embeddings = [fact.embedding for fact in processed_facts] - - async with acquire_with_retry(pool) as resolve_conn: - resolved_entity_ids, entity_to_unit, unit_to_entity_ids = await entity_processing.resolve_entities( - entity_resolver, - resolve_conn, - bank_id, - placeholder_unit_ids, - processed_facts, - log_buffer, - user_entities_per_content=user_entities_per_content, - entity_labels=getattr(config, "entity_labels", None), - ) - - # Semantic ANN search on the same connection (autocommit, no transaction). - # Skipped in streaming mode — deferred to Phase 3 to avoid O(bank_size) - # scaling bottleneck that makes later streaming batches progressively slower. - semantic_ann_links = [] - if not skip_semantic_ann and all(embedding is not None for embedding in embeddings): - fact_types = [fact.fact_type for fact in processed_facts] - semantic_ann_links = await compute_semantic_links_ann( - resolve_conn, bank_id, placeholder_unit_ids, embeddings, fact_types=fact_types, log_buffer=log_buffer - ) - elif not skip_semantic_ann: - log_buffer.append(" Semantic ANN precompute: skipped (missing embeddings)") - - return Phase1Result( - entities=EntityResolutionResult( - resolved_entity_ids=resolved_entity_ids, - entity_to_unit=entity_to_unit, - unit_to_entity_ids=unit_to_entity_ids, - ), - semantic_ann_links=semantic_ann_links, - ) - - -def _remap_phase1_results( - resolved_entity_ids: list[str], - entity_to_unit: list[tuple], - unit_to_entity_ids: dict[str, list[str]], - semantic_ann_links: list[tuple], - actual_unit_ids: list[str], -) -> tuple[list[tuple], dict[str, list[str]], list[tuple]]: - """ - Remap Phase 1 results from placeholder unit IDs to actual unit IDs. - - During Phase 1 we use str(fact_index) as placeholder unit IDs. - After insert_facts_batch creates real UUIDs, this function replaces the - placeholders so that all rows reference the correct memory_units. - """ - # Build placeholder -> actual mapping - placeholder_to_actual = {str(i): actual_id for i, actual_id in enumerate(actual_unit_ids)} - - # Remap entity_to_unit tuples - remapped_entity_to_unit = [ - (placeholder_to_actual.get(unit_id, unit_id), local_idx, fact_date) - for unit_id, local_idx, fact_date in entity_to_unit - ] - - # Remap unit_to_entity_ids keys - remapped_unit_to_entity_ids: dict[str, list[str]] = {} - for placeholder_id, entity_ids in unit_to_entity_ids.items(): - actual_id = placeholder_to_actual.get(placeholder_id, placeholder_id) - remapped_unit_to_entity_ids[actual_id] = entity_ids - - # Remap semantic ANN links (from_id uses placeholder) - remapped_semantic = [ - (placeholder_to_actual.get(lnk[0], lnk[0]), lnk[1], lnk[2], lnk[3], lnk[4]) for lnk in semantic_ann_links - ] - - return remapped_entity_to_unit, remapped_unit_to_entity_ids, remapped_semantic - - -async def _insert_facts_and_links( - conn, - embeddings_model, - entity_resolver, - bank_id: str, - contents: list[RetainContent], - extracted_facts: list, - processed_facts: list[ProcessedFact], - config, - log_buffer: list[str], - resolved_entity_ids: list[str], - entity_to_unit: list[tuple], - unit_to_entity_ids: dict[str, list[str]], - semantic_ann_links: list[tuple], - skip_semantic_links: bool = False, - outbox_callback=None, - ops=None, -) -> tuple[list[list[str]], Phase3Context]: - """ - Phase 2 of the retain pipeline: insert facts and retrieval-critical links. - - Runs inside a single database transaction to ensure atomicity of the data - that retrieval depends on (facts, unit_entities, temporal/semantic/causal links). - - Entity link generation and insertion for UI visualization are NOT done here — - only the unit_entities INSERT (FK to memory_units) stays in the transaction. - Entity link building is deferred to Phase 3 (post-transaction, best-effort). - """ - set_stage("retain.phase2.insert_facts") - # This runs inside the caller's write transaction. The bank row lock and - # conditional NULL initialisation prevent concurrent writers from mixing - # embeddings from two semantic spaces. - await ensure_bank_embedding_fingerprint( - conn, - bank_id, - embeddings_model, - policy=getattr(config, "embedding_fingerprint_policy", "strict"), - for_write=True, - legacy_attestation=getattr(config, "embedding_fingerprint_legacy_attestation", None), - ) - unit_ids = await fact_storage.insert_facts_batch(conn, bank_id, processed_facts, ops=ops) - step_start = time.time() - log_buffer.append(f" Insert facts: {len(unit_ids)} units in {time.time() - step_start:.3f}s") - - # Context for Phase 3 entity link building (after transaction commits) - phase3_context = Phase3Context() - - if unit_ids: - # Entity resolution was done in Phase 1 (separate connection). - # Remap placeholder IDs to actual unit IDs. - step_start = time.time() - remapped_entity_to_unit, remapped_unit_to_entity_ids, remapped_semantic = _remap_phase1_results( - resolved_entity_ids, entity_to_unit, unit_to_entity_ids, semantic_ann_links or [], unit_ids - ) - # Update semantic_ann_links with remapped IDs for Phase 2 - semantic_ann_links = remapped_semantic - # INSERT unit_entities (FK to memory_units, must be in transaction). - # Pass fact_date alongside so entity_cooccurrences.last_cooccurred - # tracks the event timeline, not the ingest moment. - unit_entity_pairs = [ - (unit_id, resolved_entity_ids[idx], fact_date) - for idx, (unit_id, _local_idx, fact_date) in enumerate(remapped_entity_to_unit) - ] - await entity_resolver.link_units_to_entities_batch(unit_entity_pairs, conn=conn) - log_buffer.append(f" Insert unit_entities: {len(unit_entity_pairs)} pairs in {time.time() - step_start:.3f}s") - # Save context for Phase 3 entity link building (after commit) - phase3_context = Phase3Context( - unit_ids=unit_ids, - resolved_entity_ids=resolved_entity_ids, - entity_to_unit=remapped_entity_to_unit, - unit_to_entity_ids=remapped_unit_to_entity_ids, - ) - - # Create temporal links - step_start = time.time() - temporal_link_count = await link_creation.create_temporal_links_batch( - conn, - bank_id, - unit_ids, - ops=ops, - write_temporal_links=getattr(config, "write_temporal_links", True), - ) - log_buffer.append(f" Temporal links: {temporal_link_count} links in {time.time() - step_start:.3f}s") - - # Create semantic links (within-batch + pre-computed ANN from Phase 1) - if not getattr(config, "write_semantic_links", True): - log_buffer.append(" Semantic links: skipped (mode=ann)") - semantic_link_count = 0 - elif skip_semantic_links: - log_buffer.append(" Semantic links: skipped (deferred to final ANN pass)") - semantic_link_count = 0 - else: - step_start = time.time() - embeddings_for_links = [fact.embedding for fact in processed_facts] - if not all(embedding is not None for embedding in embeddings_for_links): - semantic_link_count = 0 - log_buffer.append(" Semantic links: skipped (missing embeddings)") - else: - semantic_link_count = await link_creation.create_semantic_links_batch( - conn, - bank_id, - unit_ids, - embeddings_for_links, - pre_computed_ann_links=semantic_ann_links, - ops=ops, - write_semantic_links=getattr(config, "write_semantic_links", True), - ) - log_buffer.append(f" Semantic links: {semantic_link_count} links in {time.time() - step_start:.3f}s") - - # NOTE: Entity links are NOT inserted here. They are deferred to - # Phase 3 (post-transaction, best-effort) since retrieval uses the - # unit_entities self-join instead. Entity links only serve UI visualization. - - # Create causal links - step_start = time.time() - causal_link_count = await link_creation.create_causal_links_batch( - conn, bank_id, unit_ids, processed_facts, ops=ops - ) - log_buffer.append(f" Causal links: {causal_link_count} links in {time.time() - step_start:.3f}s") - - # Map results back to original content items. Use processed_facts (not - # extracted_facts) because unit_ids has 1:1 alignment with processed_facts — - # any upstream drop between extraction and processing would otherwise cause - # an IndexError (see issue #1037). - result_unit_ids = _map_results_to_contents(contents, processed_facts, unit_ids if unit_ids else []) - - if outbox_callback: - await outbox_callback(conn) - - return result_unit_ids, phase3_context - - -async def _build_and_insert_entity_links_phase3( - pool: Any, - entity_resolver, - bank_id: str, - phase3_ctx: Phase3Context, - config, - log_buffer: list[str], -) -> None: - """ - Phase 3 helper: build entity links from resolved data and insert them. - - Runs on a fresh connection after the main transaction has committed. - Entity links are for UI graph visualization only — retrieval uses - the unit_entities self-join instead. - """ - set_stage("retain.phase3.entity_links") - if not getattr(config, "write_entity_links", True): - log_buffer.append(" Entity links (viz): skipped (write_entity_links=false)") - return - - p3_unit_ids = phase3_ctx.unit_ids - p3_resolved = phase3_ctx.resolved_entity_ids - p3_entity_to_unit = phase3_ctx.entity_to_unit - p3_unit_to_entity_ids = phase3_ctx.unit_to_entity_ids - - if not p3_unit_ids or not p3_resolved: - return - - async with acquire_with_retry(pool) as conn: - step_start = time.time() - entity_links = await entity_processing.build_entity_links( - entity_resolver, - conn, - bank_id, - p3_unit_ids, - p3_resolved, - p3_entity_to_unit, - p3_unit_to_entity_ids, - log_buffer, - skip_unit_entities_insert=True, # Already inserted in Phase 2 - ops=pool.ops, - ) - if entity_links: - await entity_processing.insert_entity_links_batch(conn, entity_links, bank_id, ops=pool.ops) - log_buffer.append(f" Entity links (viz): {len(entity_links)} links in {time.time() - step_start:.3f}s") - - -async def _extract_and_embed( - contents: list[RetainContent], - llm_config, - agent_name: str, - config, - embeddings_model, - format_date_fn, - fact_type_override: str | None, - log_buffer: list[str], - pool: Any = None, - operation_id: str | None = None, - schema: str | None = None, -) -> tuple[list, list[ProcessedFact], list[ChunkMetadata], TokenUsage]: - """ - Shared pipeline: extract facts from contents and generate embeddings. - - Returns: - Tuple of (extracted_facts, processed_facts, chunks_metadata, usage) - """ - set_stage("retain.extract_and_embed") - step_start = time.time() - extracted_facts, chunks, usage = await fact_extraction.extract_facts_from_contents( - contents, llm_config, agent_name, config, pool, operation_id, schema - ) - log_buffer.append( - f" Extract facts: {len(extracted_facts)} facts, {len(chunks)} chunks " - f"from {len(contents)} contents in {time.time() - step_start:.3f}s" - ) - - if not extracted_facts: - return extracted_facts, [], chunks, usage - - if fact_type_override: - for fact in extracted_facts: - fact.fact_type = fact_type_override - - step_start = time.time() - augmented_texts = embedding_processing.augment_texts_with_dates(extracted_facts, format_date_fn) - try: - embeddings = await embedding_processing.generate_embeddings_batch(embeddings_model, augmented_texts) - log_buffer.append(f" Generate embeddings: {len(embeddings)} embeddings in {time.time() - step_start:.3f}s") - except Exception as exc: - logger.warning("Embedding generation failed; retaining facts with projection.embedding.ok=false: %s", exc) - embeddings = [None] * len(extracted_facts) - log_buffer.append(f" Generate embeddings: failed, storing {len(embeddings)} facts without embeddings") - - embedding_version = _embedding_model_version(embeddings_model) - extraction_version = getattr(config, "extraction_prompt_version", "5w-v1") - processed_facts = [ - ProcessedFact.from_extracted_fact( - ef, - emb, - extraction_prompt_version=extraction_version, - embedding_model_version=embedding_version, - ) - for ef, emb in zip(extracted_facts, embeddings) - ] - - return extracted_facts, processed_facts, chunks, usage - - -class _RetainLogBuffer(list[str]): - """Redact private identifiers only for trusted operation-scoped retains.""" - - def __init__(self, *, sanitized: bool, secrets: tuple[str | None, ...] = ()) -> None: - super().__init__() - self._sanitized = sanitized - self._secrets = {secret for secret in secrets if secret} - - def add_secret(self, secret: str | None) -> None: - if secret: - self._secrets.add(secret) - - def redact(self, message: str) -> str: - if not self._sanitized: - return message - for secret in sorted(self._secrets, key=len, reverse=True): - message = message.replace(secret, "") - return message - - def append(self, message: str) -> None: - super().append(self.redact(message)) - - -def _log_identifier(value: str | None, *, sanitized: bool) -> str: - return "" if sanitized and value else str(value) - - -class RetainPublicationAborted(RuntimeError): - """The retain no longer owns a document state that it could publish. - - Returning a normal result from this condition is unsafe for callers that - use ``outbox_callback`` as a transactional publication fence: the async - child would be marked completed even though that callback never committed. - Keep the message identifier-free because worker failures are persisted. - """ - - -async def _consume_streaming_batches( - chunk_queue: asyncio.Queue, - *, - chunk_batch_size: int, - process_batch: Callable[[list[tuple], int, bool], Awaitable[None]], - producer_error: list[BaseException], - pipeline_aborted: list[bool], -) -> None: - """Drain enriched chunks while preserving a real final-batch boundary. - - A full batch cannot be submitted as non-final until either another item is - observed or the producer terminates. Holding one full batch as lookahead - makes an exact multiple of ``chunk_batch_size`` indistinguishable from a - partial last batch in the *right* way: both invoke ``process_batch`` once - with ``is_last=True``. That is where the transactional publication callback - is attached. - - Producer failure and document takeover deliberately suppress the final - batch callback. They must propagate as failures instead of allowing the - enclosing async operation to report a false successful publication. - """ - - if chunk_batch_size <= 0: - raise ValueError("chunk_batch_size must be positive") - - batch: list[tuple] = [] - consumer_batch_idx = 0 - - while True: - item = await chunk_queue.get() - if item is None: - # The producer records extraction failures before publishing the - # sentinel. Do not commit a remaining partial document, and let the - # caller re-raise the original provider/extraction exception. - if producer_error: - return - if pipeline_aborted[0]: - raise RetainPublicationAborted("Retain publication ownership was lost") - if batch: - await process_batch(batch, consumer_batch_idx, True) - if pipeline_aborted[0]: - raise RetainPublicationAborted("Retain publication ownership was lost") - return - - if pipeline_aborted[0]: - # Keep draining so a producer blocked on the bounded queue can - # finish and publish its sentinel. No further work may be written. - continue - - # A real successor proves that the pending full batch is not the last. - if len(batch) >= chunk_batch_size: - await process_batch(batch, consumer_batch_idx, False) - consumer_batch_idx += 1 - batch = [] - if pipeline_aborted[0]: - continue - - batch.append(item) - - -async def retain_batch( - pool: Any, - embeddings_model, - llm_config, - entity_resolver, - format_date_fn, - bank_id: str, - contents_dicts: list[RetainContentDict], - config, - document_id: str | None = None, - is_first_batch: bool = True, - fact_type_override: str | None = None, - document_tags: list[str] | None = None, - operation_id: str | None = None, - schema: str | None = None, - outbox_callback: Callable[["asyncpg.Connection"], Awaitable[None]] | None = None, - db_semaphore: "asyncio.Semaphore | None" = None, - sanitize_log_identifiers: bool = False, -) -> tuple[list[list[str]], TokenUsage, int | None]: - """ - Process a batch of content through the retain pipeline. - - Supports delta retain: when upserting a document that already has chunks, - only re-processes chunks whose content has changed. Unchanged chunks keep - their existing facts, entities, and links. - - Returns a three-tuple of: - * per-content-item unit ID lists - * aggregate LLM token usage - * processed_content_tokens — content+context tokens that actually went - through extraction after chunk-level dedup, or ``None`` if this path - didn't dedup (caller should treat as "bill full submitted content"). - See ``RetainResult.processed_content_tokens`` for details. - """ - start_time = time.time() - total_chars = sum(len(item.get("content", "")) for item in contents_dicts) - - log_buffer = _RetainLogBuffer(sanitized=sanitize_log_identifiers, secrets=(bank_id,)) - log_buffer.append(f"{'=' * 60}") - log_buffer.append(f"RETAIN_BATCH START: {bank_id}") - log_buffer.append(f"Batch size: {len(contents_dicts)} content items, {total_chars:,} chars") - log_buffer.append(f"{'=' * 60}") - - # Get bank profile - profile = await bank_utils.get_bank_profile(pool, bank_id) - agent_name = profile["name"] - - # Fail before extraction/embedding work when an existing non-empty bank is - # known to be incompatible. The write path repeats this check under a row - # lock because this preflight alone cannot close concurrent-writer races. - async with acquire_with_retry(pool) as fingerprint_conn: - await ensure_bank_embedding_fingerprint( - fingerprint_conn, - bank_id, - embeddings_model, - policy=getattr(config, "embedding_fingerprint_policy", "strict"), - legacy_attestation=getattr(config, "embedding_fingerprint_legacy_attestation", None), - ) - - # Convert dicts to RetainContent objects - contents = _build_contents(contents_dicts, document_tags) - - # When contents have multiple distinct per-content document_ids and no - # batch-level document_id, group by doc_id and process each group - # independently so each document is tracked separately. - if not document_id: - per_content_doc_ids = [item.get("document_id") for item in contents_dicts] - unique_doc_ids = {d for d in per_content_doc_ids if d} - if len(unique_doc_ids) > 1: - # Group contents by document_id, preserving original order - groups: dict[str, tuple[list[RetainContentDict], list[RetainContent]]] = {} - original_indices: dict[str, list[int]] = {} - for idx, (cd, c) in enumerate(zip(contents_dicts, contents)): - doc_key = cd.get("document_id") or str(uuid.uuid4()) - if doc_key not in groups: - groups[doc_key] = ([], []) - original_indices[doc_key] = [] - groups[doc_key][0].append(cd) - groups[doc_key][1].append(c) - original_indices[doc_key].append(idx) - - # Process each group and merge results back in original order - result_unit_ids: list[list[str]] = [[] for _ in contents_dicts] - total_usage = TokenUsage() - total_processed_tokens: int | None = 0 - for doc_key, (group_dicts, group_contents) in groups.items(): - group_ids, group_usage, group_processed = await retain_batch( - pool=pool, - embeddings_model=embeddings_model, - llm_config=llm_config, - entity_resolver=entity_resolver, - format_date_fn=format_date_fn, - bank_id=bank_id, - contents_dicts=group_dicts, - config=config, - document_id=doc_key, - is_first_batch=is_first_batch, - fact_type_override=fact_type_override, - document_tags=document_tags, - operation_id=operation_id, - schema=schema, - outbox_callback=outbox_callback, - db_semaphore=db_semaphore, - sanitize_log_identifiers=sanitize_log_identifiers, - ) - for group_idx, orig_idx in enumerate(original_indices[doc_key]): - if group_idx < len(group_ids): - result_unit_ids[orig_idx] = group_ids[group_idx] - total_usage = total_usage + group_usage - total_processed_tokens = _merge_processed_content_tokens(total_processed_tokens, group_processed) - return result_unit_ids, total_usage, total_processed_tokens - - # Resolve effective document_id early so both delta and streaming paths - # can find existing chunks from a prior attempt. On retry, a generated - # document_id is recovered from operation result_metadata.document_ids[0]. - effective_doc_id = document_id - if not effective_doc_id: - doc_ids = {item.get("document_id") for item in contents_dicts if item.get("document_id")} - if len(doc_ids) == 1: - effective_doc_id = doc_ids.pop() - if not effective_doc_id and operation_id: - try: - async with acquire_with_retry(pool) as conn: - row = await conn.fetchrow( - f"SELECT result_metadata FROM {fq_table('async_operations')} WHERE operation_id = $1", - uuid.UUID(operation_id), - ) - if row and row["result_metadata"]: - meta = ( - row["result_metadata"] - if isinstance(row["result_metadata"], dict) - else json.loads(row["result_metadata"]) - ) - recovered = meta.get("document_ids") or [] - if recovered: - effective_doc_id = recovered[0] - except Exception: - pass - if not effective_doc_id: - effective_doc_id = str(uuid.uuid4()) - log_buffer.add_secret(effective_doc_id) - - # Record effective_doc_id on the operation (idempotent set-append). Captures - # both user-provided and generated ids so the operation shows every document - # it touched, and lets retries reuse the same generated id. - if operation_id: - try: - async with acquire_with_retry(pool) as conn: - await conn.execute( - f""" - UPDATE {fq_table("async_operations")} - SET result_metadata = jsonb_set( - COALESCE(result_metadata, '{{}}'::jsonb), - '{{document_ids}}', - CASE - WHEN COALESCE(result_metadata->'document_ids', '[]'::jsonb) @> $1::jsonb - THEN result_metadata->'document_ids' - ELSE COALESCE(result_metadata->'document_ids', '[]'::jsonb) || $1::jsonb - END, - true - ), - updated_at = now() - WHERE operation_id = $2 - """, - json.dumps([effective_doc_id]), - uuid.UUID(operation_id), - ) - except Exception: - logger.warning("Failed to persist document_id", exc_info=True) - - # --- Append mode: prepend existing document content to new content --- - # When update_mode="append", fetch the existing document text and prepend it - # so the full document is reprocessed (delta retain will skip unchanged chunks). - update_mode = None - for item in contents_dicts: - item_mode = item.get("update_mode") - if item_mode: - update_mode = item_mode - break - - if update_mode == "append" and effective_doc_id and is_first_batch: - async with acquire_with_retry(pool) as conn: - existing_text = await fact_storage.get_document_content(conn, bank_id, effective_doc_id) - if existing_text: - # Prepend existing text as a new content item at the beginning - existing_content: RetainContentDict = {"content": existing_text} - # Copy context/tags from first item for consistency - first = contents_dicts[0] - if first.get("context"): - existing_content["context"] = first["context"] - if first.get("tags"): - existing_content["tags"] = first["tags"] - contents_dicts = [existing_content, *contents_dicts] - # Rebuild contents list to match - contents = _build_contents(contents_dicts, document_tags) - log_buffer.append( - f"[append] Prepended {len(existing_text):,} chars from existing document {effective_doc_id}" - ) - - # --- Stale-request check (best-effort, before LLM extraction) --- - # If the document was already updated by a more recent retain (updated_at > our - # start_time), skip this request entirely to avoid overwriting newer content - # (e.g. a longer conversation) with older data. This is an optimization — the - # real correctness guarantee comes from the FOR UPDATE + content_hash check - # inside each batch TXN (see _run_mini_batch_db_work). - async with acquire_with_retry(pool) as conn: - doc_row = await conn.fetchrow( - f"SELECT updated_at FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2", - effective_doc_id, - bank_id, - ) - if doc_row and doc_row["updated_at"]: - doc_updated = doc_row["updated_at"].timestamp() - if doc_updated > start_time: - log_buffer.append( - f"[stale] Skipping retain: document {effective_doc_id} was updated at " - f"{doc_row['updated_at'].isoformat()} (after this request started at " - f"{datetime.fromtimestamp(start_time, tz=UTC).isoformat()})" - ) - logger.info("\n" + "\n".join(log_buffer) + "\n") - if outbox_callback is not None: - # A publication callback is not a notification after the fact; - # multimodal retain uses it as the ledger CAS that proves this - # document command actually published. A stale no-op cannot - # satisfy that contract and must not look like child success. - raise RetainPublicationAborted("Retain was superseded before publication") - # No new content was processed — report 0 so callers can skip - # billing cleanly instead of falling back to full-content billing. - return [[] for _ in contents], TokenUsage(), 0 - - # --- Delta retain: check if we can skip unchanged chunks --- - if is_first_batch: - delta_result = await _try_delta_retain( - pool, - embeddings_model, - llm_config, - entity_resolver, - format_date_fn, - bank_id, - contents_dicts, - contents, - config, - effective_doc_id, - fact_type_override, - document_tags, - agent_name, - log_buffer, - start_time, - operation_id, - schema, - outbox_callback, - db_semaphore, - sanitize_log_identifiers, - ) - if delta_result is not None: - return delta_result - - # --- Always use the streaming pipeline (producer-consumer batching) --- - # Even small documents go through the same path — they just end up as a - # single batch. This eliminates the maintenance burden of two separate - # retain code paths. - chunk_batch_size = getattr(config, "retain_chunk_batch_size", 100) - chunk_size = getattr(config, "retain_chunk_size", 3000) - all_pre_chunks: list[str] = [] - chunk_to_content: list[int] = [] # maps chunk index -> index into contents - for content_idx, content in enumerate(contents): - content_chunks = fact_extraction.chunk_text(content.content, chunk_size) - all_pre_chunks.extend(content_chunks) - chunk_to_content.extend([content_idx] * len(content_chunks)) - - # Memory: after chunking, the original content bodies in RetainContent are - # no longer needed (all_pre_chunks holds the working set). Clear them so - # Python can reclaim the (potentially multi-MB) strings. - # Note: contents_dicts["content"] is still needed briefly for hash computation - # inside _streaming_retain_batch, but gets cleared there after use. - for content in contents: - content.content = "" - - total_pre_chunks = len(all_pre_chunks) - num_batches = (total_pre_chunks + chunk_batch_size - 1) // chunk_batch_size if total_pre_chunks > 0 else 1 - log_buffer.append( - f"[streaming] {total_pre_chunks} chunks, batch_size {chunk_batch_size} — " - f"{num_batches} batch{'es' if num_batches != 1 else ''}" - ) - - return await _streaming_retain_batch( - pool=pool, - embeddings_model=embeddings_model, - llm_config=llm_config, - entity_resolver=entity_resolver, - format_date_fn=format_date_fn, - bank_id=bank_id, - contents_dicts=contents_dicts, - contents=contents, - config=config, - document_id=effective_doc_id, - is_first_batch=is_first_batch, - fact_type_override=fact_type_override, - document_tags=document_tags, - agent_name=agent_name, - log_buffer=log_buffer, - start_time=start_time, - all_pre_chunks=all_pre_chunks, - chunk_to_content=chunk_to_content, - chunk_batch_size=chunk_batch_size, - operation_id=operation_id, - schema=schema, - outbox_callback=outbox_callback, - db_semaphore=db_semaphore, - sanitize_log_identifiers=sanitize_log_identifiers, - ) - - -# --------------------------------------------------------------------------- -# Final semantic ANN pass (post-commit) -# --------------------------------------------------------------------------- - -_ANN_CHUNK_SIZE = 1000 # Max seeds per ANN query — smaller chunks avoid timeouts -_ANN_PARALLELISM = 4 # Max concurrent ANN chunks to avoid pool saturation - - -async def _run_final_semantic_ann( - pool: Any, - bank_id: str, - unit_ids: list[str], - config, - log_buffer: list[str], -) -> None: - """ - Create semantic links for all committed units in a single pass. - - Called after all streaming batches have committed. Loads embeddings and - fact_types from the database, then runs ANN in chunks of _ANN_CHUNK_SIZE - seeds. This replaces per-batch within-batch + fire-and-forget ANN with - one efficient pass that sees the full bank. - """ - from .link_utils import _bulk_insert_links, compute_semantic_links_ann - - if not getattr(config, "write_semantic_links", True): - log_buffer.append("[streaming] Final ANN: semantic links skipped (mode=ann)") - return - - if not unit_ids: - return - - # Load embeddings and fact_types for all committed units - load_start = time.time() - async with acquire_with_retry(pool) as conn: - rows = await conn.fetch( - f""" - SELECT id::text, embedding::text, fact_type - FROM {fq_table("memory_units")} - WHERE bank_id = $1 AND id = ANY($2::uuid[]) - ORDER BY id - """, - bank_id, - unit_ids, - ) - - if not rows: - log_buffer.append("[streaming] Final ANN: no units found in DB (unexpected)") - return - - # Build lookup: unit_id -> (embedding_text, fact_type) - unit_map: dict[str, tuple[str, str]] = {} - for row in rows: - unit_map[row["id"]] = (row["embedding"], row["fact_type"]) - - # Filter to units that have embeddings - ann_unit_ids = [] - ann_embeddings = [] - ann_fact_types = [] - for uid in unit_ids: - if uid in unit_map and unit_map[uid][0] is not None: - ann_unit_ids.append(uid) - ann_embeddings.append(unit_map[uid][0]) # embedding as text (for temp table) - ann_fact_types.append(unit_map[uid][1]) - - log_buffer.append( - f"[streaming] Final ANN: loaded {len(ann_unit_ids)} units with embeddings in {time.time() - load_start:.3f}s" - ) - - if not ann_unit_ids: - return - - # Process in parallel chunks — each chunk runs ANN query + INSERT on its own connection. - # Parallelism bounded by _ANN_PARALLELISM to avoid saturating the connection pool. - num_chunks = (len(ann_unit_ids) + _ANN_CHUNK_SIZE - 1) // _ANN_CHUNK_SIZE - ann_semaphore = asyncio.Semaphore(_ANN_PARALLELISM) - chunk_link_counts: list[int] = [0] * num_chunks - - async def _process_ann_chunk(chunk_idx: int) -> None: - chunk_start = chunk_idx * _ANN_CHUNK_SIZE - chunk_end = min(chunk_start + _ANN_CHUNK_SIZE, len(ann_unit_ids)) - chunk_ids = ann_unit_ids[chunk_start:chunk_end] - chunk_embs = ann_embeddings[chunk_start:chunk_end] - chunk_ftypes = ann_fact_types[chunk_start:chunk_end] - - async with ann_semaphore: - t0 = time.time() - async with acquire_with_retry(pool) as conn: - ann_links = await compute_semantic_links_ann( - conn, - bank_id, - chunk_ids, - chunk_embs, - fact_types=chunk_ftypes, - top_k=20, # Recall uses at most 20 neighbors - log_buffer=log_buffer, - ) - if ann_links: - await _bulk_insert_links(conn, ann_links, bank_id=bank_id, ops=pool.ops) - chunk_link_counts[chunk_idx] = len(ann_links) - logger.info( - f"[streaming] Final ANN chunk {chunk_idx + 1}/{num_chunks}: " - f"{len(ann_links)} links in {time.time() - t0:.3f}s" - ) - - await asyncio.gather(*[_process_ann_chunk(i) for i in range(num_chunks)]) - total_links = sum(chunk_link_counts) - log_buffer.append(f"[streaming] Final ANN: {total_links} total semantic links") - - -# --------------------------------------------------------------------------- -# Streaming chunk batching -# --------------------------------------------------------------------------- - - -async def _streaming_retain_batch( - pool: Any, - embeddings_model, - llm_config, - entity_resolver, - format_date_fn, - bank_id: str, - contents_dicts: list[RetainContentDict], - contents: list[RetainContent], - config, - document_id: str | None, - is_first_batch: bool, - fact_type_override: str | None, - document_tags: list[str] | None, - agent_name: str, - log_buffer: list[str], - start_time: float, - all_pre_chunks: list[str], - chunk_to_content: list[int], - chunk_batch_size: int, - operation_id: str | None = None, - schema: str | None = None, - outbox_callback: Callable[["asyncpg.Connection"], Awaitable[None]] | None = None, - db_semaphore: "asyncio.Semaphore | None" = None, - sanitize_log_identifiers: bool = False, -) -> tuple[list[list[str]], TokenUsage, int | None]: - """ - Process a large document in streaming mini-batches to bound memory usage. - - Instead of extracting facts from ALL chunks at once (which can OOM for 17k+ - chunk documents), this splits the pre-chunked content into batches of - ``chunk_batch_size`` chunks. Each mini-batch goes through the full - extract -> embed -> Phase 1/2/3 pipeline and commits to the DB before the - next batch starts, so memory is released between batches. - - All mini-batches share the same ``document_id`` so that: - - Delta retain can detect already-committed chunks on retry - - The document row tracks the full content - - Chunks are associated with the correct document - """ - total_chunks = len(all_pre_chunks) - total_usage = TokenUsage() - all_unit_ids: list[str] = [] - - # document_id is already resolved by retain_batch (includes recovery from - # operation result_metadata on retry). - effective_doc_id = document_id - - # Default template for metadata (context, event_date, etc.) when content list is empty. - _default_content = RetainContent(content="") - - # --------------------------------------------------------------------------- - # Recovery detection (read-only, before LLM extraction) - # --------------------------------------------------------------------------- - # Check if this is a retry of the same content (crash recovery). If the - # document exists with a matching content_hash and has committed chunks, - # the producer can skip already-extracted chunks to avoid duplicate work. - existing_chunk_hashes: set[str] = set() - combined_content = "\n".join([c.get("content", "") for c in contents_dicts]) - # Memory: contents_dicts content strings are now captured in combined_content. - # Clear them from the dicts to release the per-item copies (can be multi-MB each). - for d in contents_dicts: - d.pop("content", None) - # Sanitize before hashing to match what handle_document_tracking stores - sanitized_content = fact_extraction._sanitize_text(combined_content) or "" - new_content_hash = hashlib.sha256(sanitized_content.encode()).hexdigest() - # Memory: sanitized_content is only needed for the hash; free it immediately. - sanitized_content = "" - is_recovery = False - - async def _finalize_existing_publication() -> None: - """Publish an already-written/recovered document under a row lock.""" - - if outbox_callback is None: - return - async with acquire_with_retry(pool) as conn: - async with conn.transaction(): - existing_hash = await conn.fetchval( - f"SELECT content_hash FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2 FOR UPDATE", - effective_doc_id, - bank_id, - ) - if existing_hash != new_content_hash: - raise RetainPublicationAborted("Retain publication ownership was lost") - await outbox_callback(conn) - - try: - async with acquire_with_retry(pool) as conn: - doc_row = await conn.fetchrow( - f"SELECT content_hash FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2", - effective_doc_id, - bank_id, - ) - if doc_row and doc_row["content_hash"] == new_content_hash: - existing_rows = await chunk_storage.load_existing_chunks(conn, bank_id, effective_doc_id) - existing_chunk_hashes = {c.content_hash for c in existing_rows if c.content_hash} - if existing_chunk_hashes: - is_recovery = True - log_buffer.append( - f"[streaming] RECOVERY: found {len(existing_chunk_hashes)} already-committed chunks — " - f"will skip matching and preserve existing data" - ) - except Exception: - pass # If we can't load, just process all chunks - - # --------------------------------------------------------------------------- - # Document tracking is DEFERRED to the first consumer batch TXN. - # --------------------------------------------------------------------------- - # Previously, document tracking (cascade-delete old data + insert doc row) - # ran in a separate transaction BEFORE LLM extraction. This left a gap - # between the cascade-delete and the first chunk write, allowing concurrent - # requests to interleave and produce duplicates. - # - # Now, document tracking runs atomically inside the first batch's write TXN, - # using SELECT ... FOR UPDATE on the document row for serialization across - # workers. Each batch TXN also verifies document ownership via content_hash - # to detect when a concurrent request has taken over the document. - # See _run_mini_batch_db_work() for the implementation. - retain_params, merged_tags = _build_retain_params(contents_dicts, document_tags) - # Track whether document tracking has been done (by the first batch) - doc_tracking_done = [False] - - # --------------------------------------------------------------------------- - # Producer-consumer pipeline: LLM extraction runs concurrently with DB writes - # --------------------------------------------------------------------------- - num_batches = (total_chunks + chunk_batch_size - 1) // chunk_batch_size - - # Queue for enriched chunks (extracted facts + embeddings). - # Buffer up to 2x batch_size items so the producer can stay ahead of the consumer. - chunk_queue: asyncio.Queue = asyncio.Queue(maxsize=chunk_batch_size * 2) - - # Shared mutable state for the producer to report skipped chunks and usage - producer_error: list[BaseException] = [] - # Set to True by _run_mini_batch_db_work when a concurrent request takes - # over the document (content_hash mismatch). The consumer checks this and - # stops processing further batches. - pipeline_aborted: list[bool] = [False] - - # ---- LLM Producer ---- - # Fires all chunk extractions as concurrent tasks (bounded by the LLM - # semaphore inside fact_extraction to 32 concurrent). As each completes - # it pushes the enriched result into the queue for the DB consumer. - async def _llm_producer() -> None: - async def _extract_one(global_idx: int, chunk_text: str) -> None: - source = contents[chunk_to_content[global_idx]] if contents else _default_content - content = RetainContent( - content=chunk_text, - context=source.context, - event_date=source.event_date, - metadata=source.metadata, - entities=source.entities, - tags=source.tags, - observation_scopes=source.observation_scopes, - ) - extracted, processed, chunk_meta, usage = await _extract_and_embed( - [content], - llm_config, - agent_name, - config, - embeddings_model, - format_date_fn, - fact_type_override, - log_buffer, - pool, - operation_id, - schema, - ) - await chunk_queue.put((global_idx, content, extracted, processed, chunk_meta, usage)) - # Memory: release the chunk text from the shared list now that it's - # been extracted and queued. The queued RetainContent holds its own copy. - all_pre_chunks[global_idx] = "" - - tasks: list[asyncio.Task] = [] - skipped_total = 0 - for i, chunk_text in enumerate(all_pre_chunks): - chunk_hash = chunk_storage.compute_chunk_hash(chunk_text) - if chunk_hash in existing_chunk_hashes: - # Memory: skipped chunks aren't needed either. - all_pre_chunks[i] = "" - skipped_total += 1 - continue - tasks.append(asyncio.create_task(_extract_one(i, chunk_text))) - - if skipped_total > 0: - log_buffer.append(f"[streaming] Producer: skipped {skipped_total}/{total_chunks} already-committed chunks") - - # Wait for all extractions; collect exceptions - results = await asyncio.gather(*tasks, return_exceptions=True) - for r in results: - if isinstance(r, BaseException): - producer_error.append(r) - - # Signal the consumer that production is done - await chunk_queue.put(None) - - # ---- DB Consumer ---- - # Drains enriched chunks from the queue in batches and runs - # Phase 1 (entity resolution) -> Phase 2 (write txn) -> Phase 3 (ANN fire-and-forget). - async def _db_consumer() -> None: - await _consume_streaming_batches( - chunk_queue, - chunk_batch_size=chunk_batch_size, - process_batch=_process_db_batch, - producer_error=producer_error, - pipeline_aborted=pipeline_aborted, - ) - - async def _process_db_batch( - batch: list[tuple], - consumer_batch_idx: int, - is_last: bool, - ) -> None: - """Run Phase 1 + Phase 2 + Phase 3 for a batch of pre-extracted chunks.""" - # Allow clearing combined_content after the no-facts skip path runs - # doc tracking — see the assignment further below. - nonlocal combined_content - # Combine results from individual chunk extractions - batch_contents: list[RetainContent] = [] - batch_extracted: list = [] - batch_processed: list[ProcessedFact] = [] - batch_chunk_meta: list[ChunkMetadata] = [] - batch_usage = TokenUsage() - - for global_idx, content, extracted, processed, chunk_meta, usage in batch: - content_idx_in_batch = len(batch_contents) - # Adjust chunk indices to use the original global position (global_idx) - # so that chunk_id = {bank}_{doc}_{chunk_index} is deterministic regardless - # of task completion order. content_index is batch-relative for result grouping. - for fact in extracted: - fact.content_index = content_idx_in_batch - if fact.chunk_index is not None: - fact.chunk_index = global_idx - for pf in processed: - pf.content_index = content_idx_in_batch - for cm in chunk_meta: - cm.chunk_index = global_idx - - batch_contents.append(content) - batch_extracted.extend(extracted) - batch_processed.extend(processed) - batch_chunk_meta.extend(chunk_meta) - batch_usage = batch_usage + usage - - nonlocal total_usage - total_usage = total_usage + batch_usage - - if not batch_extracted: - # Even with 0 facts, the first batch must still run document tracking - # (cascade-delete + insert doc row) to establish ownership and prevent - # concurrent requests from interleaving. Later batches can safely skip. - if not doc_tracking_done[0]: - async with acquire_with_retry(pool) as conn: - async with conn.transaction(): - await ensure_bank_embedding_fingerprint( - conn, - bank_id, - embeddings_model, - policy=getattr(config, "embedding_fingerprint_policy", "strict"), - for_write=True, - legacy_attestation=getattr(config, "embedding_fingerprint_legacy_attestation", None), - ) - await conn.execute( - f"INSERT INTO {fq_table('documents')} (id, bank_id, original_text, content_hash) " - f"VALUES ($1, $2, '', '__pending__') " - f"ON CONFLICT (id, bank_id) DO NOTHING", - effective_doc_id, - bank_id, - ) - await conn.fetchval( - f"SELECT content_hash FROM {fq_table('documents')} " - f"WHERE id = $1 AND bank_id = $2 FOR UPDATE", - effective_doc_id, - bank_id, - ) - if is_recovery: - await fact_storage.upsert_document_metadata( - conn, - bank_id, - effective_doc_id, - combined_content, - retain_params, - merged_tags, - ) - else: - await fact_storage.handle_document_tracking( - conn, - bank_id, - effective_doc_id, - combined_content, - is_first_batch, - retain_params, - merged_tags, - ops=pool.ops, - ) - if is_last and outbox_callback is not None: - await outbox_callback(conn) - doc_tracking_done[0] = True - # Memory: combined_content has been persisted; release - # it now so the rest of the consumer loop doesn't pin - # a multi-MB string. Nothing reads it after tracking. - combined_content = "" - log_buffer.append(f"[streaming] Document {effective_doc_id} tracked (0 facts in first batch)") - elif is_last: - await _finalize_existing_publication() - log_buffer.append( - f"[streaming] Consumer batch {consumer_batch_idx + 1}: " - f"0 facts extracted from {len(batch)} chunks, skipping" - ) - return - - log_buffer.append( - f"[streaming] Consumer batch {consumer_batch_idx + 1}: " - f"processing {len(batch_extracted)} facts from {len(batch)} chunks" - ) - - async def _run_mini_batch_db_work() -> None: - # Allow clearing combined_content after the doc-tracking call so - # subsequent batches don't carry the per-document text in memory. - nonlocal combined_content - entity_resolver.discard_pending_stats() - mb_start = time.time() - - # Phase 1 — Entity Resolution only (no ANN — deferred to Phase 3) - p1_start = time.time() - phase1 = await _pre_resolve_phase1( - pool, - entity_resolver, - bank_id, - batch_contents, - batch_processed, - config, - log_buffer, - skip_semantic_ann=True, - ) - - logger.info(f"[streaming] Phase 1 (entity resolution): {time.time() - p1_start:.3f}s") - - # Phase 2 — Write transaction - # ----------------------------------------------------------------- - # Concurrent-safety via row-level locking: - # - # The streaming pipeline splits work across multiple batch TXNs. - # Without protection, two concurrent retains for the same document - # can interleave: Request A writes batch1, Request B cascade-deletes - # A's doc and writes its own batch1, then A's batch2 adds stale data - # on top of B's → duplicates. - # - # To prevent this, every batch TXN: - # 1. SELECT ... FOR UPDATE on the document row — serializes all - # writers for this document at the DB level (works across workers). - # 2. Check content_hash — if it doesn't match ours, another request - # took over the document → abort remaining batches. - # 3. First batch only: run handle_document_tracking (cascade-delete - # old data + insert doc row) atomically with the first chunk write. - # This eliminates the gap between "delete old" and "insert new" - # that previously allowed interleaving. - # ----------------------------------------------------------------- - - p2_start = time.time() - batch_result_ids = None - phase3_ctx = None - async with acquire_with_retry(pool) as conn: - async with conn.transaction(): - # --- Document ownership gate --- - # Lock the document row to serialize all concurrent writers. - # SELECT ... FOR UPDATE doesn't lock non-existent rows, so we - # first ensure the row exists with a lightweight upsert, THEN lock it. - # The content_hash='__pending__' placeholder is immediately overwritten - # by handle_document_tracking or upsert_document_metadata below. - await conn.execute( - f"INSERT INTO {fq_table('documents')} (id, bank_id, original_text, content_hash) " - f"VALUES ($1, $2, '', '__pending__') " - f"ON CONFLICT (id, bank_id) DO NOTHING", - effective_doc_id, - bank_id, - ) - existing_hash = await conn.fetchval( - f"SELECT content_hash FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2 FOR UPDATE", - effective_doc_id, - bank_id, - ) - - if not doc_tracking_done[0]: - # --- First batch: document tracking (atomic with chunk write) --- - if is_recovery: - await fact_storage.upsert_document_metadata( - conn, - bank_id, - effective_doc_id, - combined_content, - retain_params, - merged_tags, - ) - log_buffer.append( - f"[streaming] Document {effective_doc_id} updated " - f"(recovery, preserving existing chunks)" - ) - else: - await fact_storage.handle_document_tracking( - conn, - bank_id, - effective_doc_id, - combined_content, - is_first_batch, - retain_params, - merged_tags, - ops=pool.ops, - ) - log_buffer.append(f"[streaming] Document {effective_doc_id} tracked (full content)") - doc_tracking_done[0] = True - # Memory: combined_content is no longer needed after - # this first-batch tracking call. Release it so the - # remaining consumer batches don't pin the string. - combined_content = "" - else: - # --- Later batches: verify we still own the document --- - # If another request took over (cascade-deleted our doc and - # inserted its own), the content_hash won't match ours. - if existing_hash is not None and existing_hash != new_content_hash: - log_buffer.append( - f"[streaming] Document {effective_doc_id} taken over by " - f"concurrent request (hash mismatch) — aborting remaining batches" - ) - logger.info("\n" + "\n".join(log_buffer) + "\n") - # Signal the consumer to stop processing further batches - pipeline_aborted[0] = True - return - - # Store chunks with correct global indices - step_start = time.time() - chunk_id_map = {} - if batch_chunk_meta: - chunk_id_map = await chunk_storage.store_chunks_batch( - conn, bank_id, effective_doc_id, batch_chunk_meta, ops=pool.ops - ) - log_buffer.append( - f" Store chunks: {len(batch_chunk_meta)} chunks in {time.time() - step_start:.3f}s" - ) - - # Map document_id and chunk_id to processed facts - for fact, processed_fact in zip(batch_extracted, batch_processed): - processed_fact.document_id = effective_doc_id - if batch_chunk_meta and fact.chunk_index is not None: - chunk_id = chunk_id_map.get(fact.chunk_index) - if chunk_id: - processed_fact.chunk_id = chunk_id - - # Insert facts and links — skip semantic links entirely in streaming - # mode; they are created in a single final ANN pass after all batches. - batch_result_ids, phase3_ctx = await _insert_facts_and_links( - conn, - embeddings_model, - entity_resolver, - bank_id, - batch_contents, - batch_extracted, - batch_processed, - config, - log_buffer, - resolved_entity_ids=phase1.entities.resolved_entity_ids, - entity_to_unit=phase1.entities.entity_to_unit, - unit_to_entity_ids=phase1.entities.unit_to_entity_ids, - semantic_ann_links=[], - skip_semantic_links=True, - outbox_callback=outbox_callback if is_last else None, - ops=pool.ops, - ) - - logger.info(f"[streaming] Phase 2 (write txn): {time.time() - p2_start:.3f}s") - - # Best-effort: entity viz + stats (fast, not semantic ANN) - if phase3_ctx is not None: - try: - await entity_resolver.flush_pending_stats() - await _build_and_insert_entity_links_phase3( - pool, entity_resolver, bank_id, phase3_ctx, config, log_buffer - ) - except Exception: - logger.warning(f"Phase 3 stats (consumer batch {consumer_batch_idx + 1}) failed", exc_info=True) - - logger.info( - f"[streaming] Consumer batch {consumer_batch_idx + 1} total " - f"(excluding fire-and-forget): {time.time() - mb_start:.3f}s" - ) - - # Collect unit_ids from this batch - if batch_result_ids: - for content_ids in batch_result_ids: - all_unit_ids.extend(content_ids) - - if db_semaphore is not None: - async with db_semaphore: - await _run_mini_batch_db_work() - else: - await _run_mini_batch_db_work() - - # Memory: after DB write, clear the batch-local lists that hold extracted - # facts and embedding vectors. These can be large (384 floats per fact × - # thousands of facts) and are no longer needed after commit. - batch_contents.clear() - batch_extracted.clear() - batch_processed.clear() - batch_chunk_meta.clear() - - # --------------------------------------------------------------------------- - # Check if facts are already committed (recovery from previous crash). - # If so, skip extraction+writes and jump straight to final ANN pass. - # --------------------------------------------------------------------------- - facts_already_committed = False - if operation_id: - try: - async with acquire_with_retry(pool) as conn: - row = await conn.fetchrow( - f"SELECT result_metadata FROM {fq_table('async_operations')} WHERE operation_id = $1", - uuid.UUID(operation_id), - ) - if row and row["result_metadata"]: - meta = ( - row["result_metadata"] - if isinstance(row["result_metadata"], dict) - else json.loads(row["result_metadata"]) - ) - committed_doc_ids = meta.get("facts_committed_document_ids") or [] - document_ids = meta.get("document_ids") or [] - # Legacy path: operations created before per-document checkpoint - # tracking only wrote facts_committed=true without document IDs. - # Treat those as committed only for single-doc operations. - legacy_single_doc_checkpoint = ( - meta.get("facts_committed") - and not committed_doc_ids - and (len(document_ids) <= 1 or document_ids == [effective_doc_id]) - ) - if effective_doc_id in committed_doc_ids or legacy_single_doc_checkpoint: - facts_already_committed = True - log_buffer.append( - f"[streaming] Recovery: facts already committed ({meta.get('unit_ids_count', '?')} units), " - f"skipping to final ANN pass" - ) - except Exception: - logger.warning("Failed to check operation recovery state", exc_info=True) - - if not facts_already_committed: - # Run producer and consumer concurrently - await asyncio.gather(_llm_producer(), _db_consumer()) - - # Propagate producer errors (e.g. LLM failures) - if producer_error: - raise producer_error[0] - - # If no batch was processed (e.g. zero facts extracted from gibberish - # content, or all chunks skipped in recovery), the document row was - # never created by the first batch TXN. Create it now so the document - # is tracked regardless of extraction results. - if not doc_tracking_done[0] and not pipeline_aborted[0]: - async with acquire_with_retry(pool) as conn: - async with conn.transaction(): - await ensure_bank_embedding_fingerprint( - conn, - bank_id, - embeddings_model, - policy=getattr(config, "embedding_fingerprint_policy", "strict"), - for_write=True, - legacy_attestation=getattr(config, "embedding_fingerprint_legacy_attestation", None), - ) - await conn.execute( - f"INSERT INTO {fq_table('documents')} (id, bank_id, original_text, content_hash) " - f"VALUES ($1, $2, '', '__pending__') " - f"ON CONFLICT (id, bank_id) DO NOTHING", - effective_doc_id, - bank_id, - ) - await conn.fetchval( - f"SELECT content_hash FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2 FOR UPDATE", - effective_doc_id, - bank_id, - ) - if is_recovery: - await fact_storage.upsert_document_metadata( - conn, - bank_id, - effective_doc_id, - combined_content, - retain_params, - merged_tags, - ) - else: - await fact_storage.handle_document_tracking( - conn, - bank_id, - effective_doc_id, - combined_content, - is_first_batch, - retain_params, - merged_tags, - ops=pool.ops, - ) - if outbox_callback is not None: - await outbox_callback(conn) - doc_tracking_done[0] = True - # Memory: combined_content has been persisted and won't be - # read again — release the per-document text now. - combined_content = "" - log_buffer.append(f"[streaming] Document {effective_doc_id} tracked (no facts extracted)") - - # Mark facts as committed in operation metadata (crash recovery checkpoint) - if operation_id and all_unit_ids: - try: - async with acquire_with_retry(pool) as conn: - # Append effective_doc_id to the committed document set if not - # already present, so multi-doc batches track each document - # independently for crash recovery. - await conn.execute( - f""" - UPDATE {fq_table("async_operations")} - SET result_metadata = jsonb_set( - result_metadata || $1::jsonb, - '{{facts_committed_document_ids}}', - CASE - WHEN COALESCE(result_metadata->'facts_committed_document_ids', '[]'::jsonb) @> $2::jsonb - THEN result_metadata->'facts_committed_document_ids' - ELSE COALESCE(result_metadata->'facts_committed_document_ids', '[]'::jsonb) || $2::jsonb - END, - true - ), - updated_at = now() - WHERE operation_id = $3 - """, - json.dumps({"facts_committed": True, "unit_ids_count": len(all_unit_ids)}), - json.dumps([effective_doc_id]), - uuid.UUID(operation_id), - ) - log_buffer.append(f"[streaming] Checkpoint: {len(all_unit_ids)} facts committed, ANN pass next") - except Exception: - logger.warning("Failed to save facts_committed checkpoint", exc_info=True) - else: - # Recovery path: load committed unit IDs from DB - async with acquire_with_retry(pool) as conn: - rows = await conn.fetch( - f""" - SELECT id::text FROM {fq_table("memory_units")} - WHERE bank_id = $1 AND document_id = $2 - ORDER BY created_at - """, - bank_id, - effective_doc_id, - ) - all_unit_ids = [row["id"] for row in rows] - log_buffer.append(f"[streaming] Recovery: loaded {len(all_unit_ids)} unit IDs from DB") - await _finalize_existing_publication() - - # --------------------------------------------------------------------------- - # Final ANN pass: create semantic links for ALL committed units at once. - # This replaces per-batch within-batch + fire-and-forget ANN with a single - # efficient pass after all facts are in the database. - # --------------------------------------------------------------------------- - if all_unit_ids and not pipeline_aborted[0]: - ann_start = time.time() - try: - await _run_final_semantic_ann(pool, bank_id, all_unit_ids, config, log_buffer) - except Exception: - # ANN pass is best-effort. FK violations can occur if a concurrent - # retain cascade-deleted our units between the batch commit and here. - log_document_id = _log_identifier(effective_doc_id, sanitized=sanitize_log_identifiers) - logger.warning( - f"[streaming] Final ANN pass failed for document {log_document_id} " - f"(units may have been superseded by concurrent retain)", - exc_info=True, - ) - log_buffer.append(f"[streaming] Final ANN pass: {time.time() - ann_start:.3f}s for {len(all_unit_ids)} units") - - total_time = time.time() - start_time - log_buffer.append(f"{'=' * 60}") - if pipeline_aborted[0]: - log_buffer.append( - f"STREAMING RETAIN ABORTED: document {effective_doc_id} was taken over by " - f"a concurrent request after {total_time:.3f}s — data from this request was discarded" - ) - else: - log_buffer.append( - f"STREAMING RETAIN COMPLETE: {len(all_unit_ids)} units across {num_batches} batches in {total_time:.3f}s" - ) - log_buffer.append(f"Document: {effective_doc_id}") - log_buffer.append(f"{'=' * 60}") - logger.info("\n" + "\n".join(log_buffer) + "\n") - - if pipeline_aborted[0]: - raise RetainPublicationAborted("Retain publication ownership was lost") - - # Map all unit_ids back to the original content items. - # For streaming mode with a single document, all units belong to content 0. - result_unit_ids = [all_unit_ids] + [[] for _ in contents[1:]] - # The streaming path doesn't compute per-chunk content-hash dedup in - # a way that lets us report a partial-processed tokens count — signal - # ``None`` so callers bill against the full submitted payload. - return result_unit_ids, total_usage, None - - -# --------------------------------------------------------------------------- -# Delta retain -# --------------------------------------------------------------------------- - - -async def _try_delta_retain( - pool: Any, - embeddings_model, - llm_config, - entity_resolver, - format_date_fn, - bank_id, - contents_dicts, - contents, - config, - document_id, - fact_type_override, - document_tags, - agent_name, - log_buffer, - start_time, - operation_id, - schema, - outbox_callback, - db_semaphore: "asyncio.Semaphore | None" = None, - sanitize_log_identifiers: bool = False, -) -> tuple[list[list[str]], TokenUsage, int | None] | None: - """ - Attempt delta retain for a document upsert. Returns result tuple if delta - was performed, or None to fall back to full retain. - - When a result tuple is returned, the third element is the content+context - token count for the chunks that actually went through extraction - (``0`` if the submission matched prior content exactly and nothing was - re-extracted). - """ - # Need a single document_id - effective_doc_id = document_id - if not effective_doc_id: - doc_ids = {item.get("document_id") for item in contents_dicts if item.get("document_id")} - if len(doc_ids) != 1: - return None - effective_doc_id = doc_ids.pop() - - # Load existing chunks and snapshot the document's content_hash. This is - # outside the write TXN, so a concurrent retain could modify the document - # between this read and the write. The write TXN verifies the hash hasn't - # changed; if it has, we fall back to streaming (which has full protection). - async with acquire_with_retry(pool) as conn: - existing_chunks = await chunk_storage.load_existing_chunks(conn, bank_id, effective_doc_id) - doc_hash_at_load = await conn.fetchval( - f"SELECT content_hash FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2", - effective_doc_id, - bank_id, - ) - - if not existing_chunks: - return None - - if any(c.content_hash is None for c in existing_chunks): - log_document_id = _log_identifier(effective_doc_id, sanitized=sanitize_log_identifiers) - logger.info(f"Delta retain skipped for {log_document_id}: existing chunks lack content_hash (pre-migration)") - return None - - # Chunk new content and classify changes - step_start = time.time() - new_chunks_with_contents = _chunk_contents_for_delta(contents, config) - log_buffer.append( - f"[delta] Chunked new content: {len(new_chunks_with_contents)} chunks in {time.time() - step_start:.3f}s" - ) - - existing_by_index = {c.chunk_index: c for c in existing_chunks} - new_hashes = {idx: chunk_storage.compute_chunk_hash(text) for idx, text in new_chunks_with_contents.items()} - - unchanged_indices, changed_indices, new_indices, removed_indices = [], [], [], [] - for idx, new_hash in new_hashes.items(): - existing = existing_by_index.get(idx) - if existing and existing.content_hash == new_hash: - unchanged_indices.append(idx) - elif existing: - changed_indices.append(idx) - else: - new_indices.append(idx) - for idx in existing_by_index: - if idx not in new_hashes: - removed_indices.append(idx) - - log_buffer.append( - f"[delta] Chunk diff: {len(unchanged_indices)} unchanged, " - f"{len(changed_indices)} changed, {len(new_indices)} new, " - f"{len(removed_indices)} removed" - ) - - if not unchanged_indices: - log_document_id = _log_identifier(effective_doc_id, sanitized=sanitize_log_identifiers) - logger.info(f"Delta retain: no unchanged chunks for {log_document_id}, falling back to full retain") - return None - - chunks_to_process = changed_indices + new_indices - - if not chunks_to_process and not removed_indices: - # Nothing changed — just update document metadata/tags - log_buffer.append("[delta] No chunk changes detected — updating document metadata only") - return await _delta_metadata_only( - pool, - bank_id, - contents_dicts, - contents, - effective_doc_id, - document_tags, - log_buffer, - start_time, - outbox_callback, - ) - - # Build content items for only the changed/new chunks - delta_contents, delta_chunk_map = _build_delta_contents(contents, new_chunks_with_contents, chunks_to_process) - - if not delta_contents: - return await _delta_metadata_only( - pool, - bank_id, - contents_dicts, - contents, - effective_doc_id, - document_tags, - log_buffer, - start_time, - outbox_callback, - ) - - # Extract facts and generate embeddings (shared pipeline) - extracted_facts, processed_facts, new_chunk_metadata, usage = await _extract_and_embed( - delta_contents, - llm_config, - agent_name, - config, - embeddings_model, - format_date_fn, - fact_type_override, - log_buffer, - pool, - operation_id, - schema, - ) - - # Database transaction - result_unit_ids: list[list[str]] = [] - log_buffer_pre_db = len(log_buffer) - - async def _run_delta_db_work() -> bool: - nonlocal result_unit_ids - del log_buffer[log_buffer_pre_db:] - for pf in processed_facts: - pf.document_id = None - pf.chunk_id = None - entity_resolver.discard_pending_stats() - - # PHASE 1 — Entity Resolution + Semantic ANN (separate connection, read-heavy) - phase1 = await _pre_resolve_phase1( - pool, - entity_resolver, - bank_id, - delta_contents, - processed_facts, - config, - log_buffer, - skip_semantic_ann=not getattr(config, "write_semantic_links", True), - ) - - # PHASE 2 — Core Write Transaction (atomic) - # Lock the document row and verify ownership. Delta loaded existing - # chunks OUTSIDE this TXN, so a concurrent retain may have cascade-deleted - # and replaced the document since then. If the content_hash changed, - # the chunk state we based our delta diff on is stale — abort. - async with acquire_with_retry(pool) as conn: - async with conn.transaction(): - current_hash = await conn.fetchval( - f"SELECT content_hash FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2 FOR UPDATE", - effective_doc_id, - bank_id, - ) - # Verify the document hasn't been replaced since we loaded chunks. - # Compare the current hash against what we snapshotted at load time. - if current_hash is not None and doc_hash_at_load is not None and current_hash != doc_hash_at_load: - log_buffer.append( - f"[delta] Document {effective_doc_id} was modified by concurrent request " - f"since chunks were loaded — aborting delta, falling back to full retain" - ) - logger.info("\n" + "\n".join(log_buffer) + "\n") - # Tell the caller to fall back to streaming (which has full - # FOR UPDATE protection). This result must be propagated: - # treating the empty ``result_unit_ids`` as a successful - # delta would skip the publication callback while allowing - # an async child operation to be marked completed. - return False - - # Update document metadata (no delete) - step_start = time.time() - combined_content = "\n".join([c.get("content", "") for c in contents_dicts]) - retain_params, merged_tags = _build_retain_params(contents_dicts, document_tags) - await fact_storage.upsert_document_metadata( - conn, - bank_id, - effective_doc_id, - combined_content, - retain_params, - merged_tags, - ) - log_buffer.append(f" Document metadata update in {time.time() - step_start:.3f}s") - - # Delete changed and removed chunks (cascades to memory_units and links) - step_start = time.time() - chunks_to_delete = [ - existing_by_index[idx].chunk_id - for idx in changed_indices + removed_indices - if idx in existing_by_index - ] - await chunk_storage.delete_chunks_by_ids(conn, chunks_to_delete) - log_buffer.append( - f" Deleted {len(chunks_to_delete)} chunks " - f"({len(changed_indices)} changed + {len(removed_indices)} removed) " - f"in {time.time() - step_start:.3f}s" - ) - - # Update tags on unchanged chunks' memory units - step_start = time.time() - updated_count = await fact_storage.update_memory_units_tags( - conn, bank_id, effective_doc_id, merged_tags - ) - log_buffer.append( - f" Updated tags on {updated_count} existing memory units in {time.time() - step_start:.3f}s" - ) - - # Store new/changed chunks - step_start = time.time() - chunk_id_map_by_doc = {} - if new_chunk_metadata: - remapped_chunks = [ - ChunkMetadata( - chunk_text=cm.chunk_text, - fact_count=cm.fact_count, - content_index=cm.content_index, - chunk_index=delta_chunk_map.get(cm.chunk_index, cm.chunk_index), - ) - for cm in new_chunk_metadata - ] - chunk_id_map = await chunk_storage.store_chunks_batch( - conn, bank_id, effective_doc_id, remapped_chunks, ops=pool.ops - ) - for chunk_idx, chunk_id in chunk_id_map.items(): - chunk_id_map_by_doc[(effective_doc_id, chunk_idx)] = chunk_id - log_buffer.append( - f" Stored {len(remapped_chunks)} new/changed chunks in {time.time() - step_start:.3f}s" - ) - - # Map chunk_ids and document_ids to processed facts - for ef, pf in zip(extracted_facts, processed_facts): - pf.document_id = effective_doc_id - if ef.chunk_index is not None: - original_idx = delta_chunk_map.get(ef.chunk_index, ef.chunk_index) - chunk_id = chunk_id_map_by_doc.get((effective_doc_id, original_idx)) - if chunk_id: - pf.chunk_id = chunk_id - - # Insert facts and retrieval-critical links. - # Use delta_contents (the changed/new chunks) as the content list, - # since extracted_facts have content_index relative to delta_contents. - result_unit_ids, phase3_ctx = await _insert_facts_and_links( - conn, - embeddings_model, - entity_resolver, - bank_id, - delta_contents, - extracted_facts, - processed_facts, - config, - log_buffer, - resolved_entity_ids=phase1.entities.resolved_entity_ids, - entity_to_unit=phase1.entities.entity_to_unit, - unit_to_entity_ids=phase1.entities.unit_to_entity_ids, - semantic_ann_links=phase1.semantic_ann_links, - outbox_callback=outbox_callback, - ops=pool.ops, - ) - - # PHASE 3 — Best-Effort Display Data (post-transaction) - try: - await entity_resolver.flush_pending_stats() - await _build_and_insert_entity_links_phase3( - pool, - entity_resolver, - bank_id, - phase3_ctx, - config, - log_buffer, - ) - except Exception: - logger.warning("Phase 3 (best-effort display data) failed — retrieval unaffected", exc_info=True) - - total_time = time.time() - start_time - log_buffer.append(f"{'=' * 60}") - log_buffer.append( - f"DELTA RETAIN COMPLETE: {len(processed_facts)} new units, " - f"{len(unchanged_indices)} chunks unchanged in {total_time:.3f}s" - ) - log_buffer.append(f"Document: {effective_doc_id}") - log_buffer.append(f"{'=' * 60}") - logger.info("\n" + "\n".join(log_buffer) + "\n") - - return True - - if db_semaphore is not None: - async with db_semaphore: - delta_committed = await _run_delta_db_work() - else: - delta_committed = await _run_delta_db_work() - if not delta_committed: - entity_resolver.discard_pending_stats() - if outbox_callback is not None: - # A publication-fenced request must not re-enter streaming after it - # has proved that another writer changed the document. Streaming's - # first batch establishes new ownership before the final callback; - # doing that here could let an obsolete command overwrite already - # published data before its final CAS rejects it. Fail this attempt - # and let the operation retry/reconcile from a fresh snapshot. - raise RetainPublicationAborted("Document changed during retain publication") - return None - # Count content + context tokens that actually went through extraction. - # ``delta_contents`` holds the per-chunk RetainContent items for the - # changed/new chunks (see ``_build_delta_contents``) — i.e. exactly what - # the LLM pipeline saw this call. Unchanged chunks contribute zero. - processed_tokens = _count_delta_content_tokens(delta_contents) - return result_unit_ids, usage, processed_tokens - - -async def _delta_metadata_only( - pool: Any, - bank_id, - contents_dicts, - contents, - document_id, - document_tags, - log_buffer, - start_time, - outbox_callback, -): - """Handle the case where no chunks changed — just update document metadata and tags.""" - async with acquire_with_retry(pool) as conn: - async with conn.transaction(): - # Lock the document row to serialize with concurrent retains - await conn.fetchval( - f"SELECT content_hash FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2 FOR UPDATE", - document_id, - bank_id, - ) - combined_content = "\n".join([c.get("content", "") for c in contents_dicts]) - retain_params, merged_tags = _build_retain_params(contents_dicts, document_tags) - await fact_storage.upsert_document_metadata( - conn, - bank_id, - document_id, - combined_content, - retain_params, - merged_tags, - ) - await fact_storage.update_memory_units_tags(conn, bank_id, document_id, merged_tags) - if outbox_callback: - await outbox_callback(conn) - - total_time = time.time() - start_time - log_buffer.append(f"DELTA RETAIN (no changes): metadata updated in {total_time:.3f}s") - logger.info("\n" + "\n".join(log_buffer) + "\n") - # Nothing went through the extraction pipeline — report 0 processed - # content tokens so callers can bill accordingly (a caller that's been - # told ``0`` knows the retain was a pure metadata update and should - # charge nothing for content). - return [[] for _ in contents], TokenUsage(), 0 - - -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- - - -def _build_contents(contents_dicts: list[RetainContentDict], document_tags: list[str] | None) -> list[RetainContent]: - """Convert content dicts to RetainContent objects.""" - contents = [] - for item in contents_dicts: - item_tags = item.get("tags", []) or [] - merged_tags = list(set(item_tags + (document_tags or []))) - - if "event_date" in item and item["event_date"] is None: - event_date_value = None - elif item.get("event_date"): - event_date_value = parse_datetime_flexible(item["event_date"]) - else: - event_date_value = utcnow() - - content = RetainContent( - content=item["content"], - context=item.get("context", ""), - event_date=event_date_value, - metadata=item.get("metadata", {}), - entities=item.get("entities", []), - tags=merged_tags, - observation_scopes=item.get("observation_scopes"), - ) - contents.append(content) - return contents - - -def _chunk_contents_for_delta(contents: list[RetainContent], config) -> dict[int, str]: - """ - Chunk contents the same way the streaming path does, returning a map of - global_chunk_index -> chunk_text. - - Must use the same chunk_size as the streaming path (default 3000) so that - chunk boundaries match and delta can detect unchanged chunks. - Previously defaulted to 120000, causing all chunks to appear changed on retry. - """ - result = {} - global_chunk_idx = 0 - for content in contents: - chunk_size = getattr(config, "retain_chunk_size", 3000) - chunks = fact_extraction.chunk_text(content.content, chunk_size) - for chunk_text in chunks: - result[global_chunk_idx] = chunk_text - global_chunk_idx += 1 - return result - - -def _build_delta_contents( - original_contents: list[RetainContent], - new_chunks_with_contents: dict[int, str], - chunks_to_process: list[int], -) -> tuple[list[RetainContent], dict[int, int]]: - """ - Build RetainContent items containing only the chunks that need processing. - - Returns: - - List of RetainContent items (one per chunk to process) - - Map of delta_chunk_index -> original_chunk_index - """ - if not chunks_to_process or not original_contents: - return [], {} - - template_content = original_contents[0] - delta_contents = [] - delta_chunk_map = {} - - for original_chunk_idx in sorted(chunks_to_process): - chunk_text = new_chunks_with_contents.get(original_chunk_idx) - if not chunk_text: - continue - delta_content = RetainContent( - content=chunk_text, - context=template_content.context, - event_date=template_content.event_date, - metadata=template_content.metadata, - entities=template_content.entities, - tags=template_content.tags, - observation_scopes=template_content.observation_scopes, - ) - delta_contents.append(delta_content) - delta_chunk_map[len(delta_contents) - 1] = original_chunk_idx - - return delta_contents, delta_chunk_map - - -def _map_results_to_contents( - contents: list[RetainContent], - processed_facts: list[ProcessedFact], - unit_ids: list[str], -) -> list[list[str]]: - """Map created unit IDs back to original content items. - - `processed_facts` and `unit_ids` must have the same length: each unit_id - corresponds to the processed_fact at the same index. - """ - if len(processed_facts) != len(unit_ids): - raise ValueError(f"processed_facts ({len(processed_facts)}) and unit_ids ({len(unit_ids)}) length mismatch") - - facts_by_content: dict[int, list[int]] = {i: [] for i in range(len(contents))} - for i, fact in enumerate(processed_facts): - # Normalize content_index: some LLM providers return 1-indexed values. - # Clamp to valid range to prevent KeyError. - idx = fact.content_index - if idx < 0 or idx >= len(contents): - idx = min(max(idx, 0), len(contents) - 1) if len(contents) > 0 else 0 - facts_by_content[idx].append(i) - - result_unit_ids = [] - for content_index in range(len(contents)): - content_unit_ids = [unit_ids[i] for i in facts_by_content[content_index]] - result_unit_ids.append(content_unit_ids) - - return result_unit_ids diff --git a/core/dataplane/hms_api/engine/retain/types.py b/core/dataplane/hms_api/engine/retain/types.py index 54427d2..8c4246c 100644 --- a/core/dataplane/hms_api/engine/retain/types.py +++ b/core/dataplane/hms_api/engine/retain/types.py @@ -10,6 +10,8 @@ from typing import Literal, TypedDict from uuid import UUID +from ..entity_resolution_contracts import EntityResolutionReadPlan + class RetainContentDict(TypedDict, total=False): """Type definition for content items in retain_batch_async. @@ -280,6 +282,14 @@ class Phase1Result: semantic_ann_links: list[tuple] +@dataclass +class EntityReadPlanPhase1Result: + """Phase-1 output with unresolved entity creation deferred to the UoW.""" + + entity_read_plan: EntityResolutionReadPlan + semantic_ann_links: list[tuple] + + @dataclass class EntityLink: """ diff --git a/core/dataplane/hms_api/worker/poller.py b/core/dataplane/hms_api/worker/poller.py index b63b880..9e80af8 100644 --- a/core/dataplane/hms_api/worker/poller.py +++ b/core/dataplane/hms_api/worker/poller.py @@ -463,6 +463,7 @@ async def _mark_completed(self, operation_id: str, schema: str | None): UPDATE {table} SET status = 'completed', completed_at = now(), updated_at = now() WHERE operation_id = $1 + AND status IN ('pending', 'processing') """, operation_id, ) @@ -480,6 +481,7 @@ async def _mark_failed(self, operation_id: str, error_message: str, schema: str UPDATE {table} SET status = 'failed', error_message = $2, completed_at = now(), updated_at = now() WHERE operation_id = $1 + AND status IN ('pending', 'processing') """, operation_id, error_message, @@ -519,12 +521,14 @@ async def _maybe_update_parent_operation(self, child_operation_id: str, schema: # Lock parent to prevent concurrent sibling updates parent_row = await conn.fetchrow( - f"SELECT operation_id FROM {table} WHERE operation_id = $1 AND bank_id = $2 FOR UPDATE", + f"SELECT operation_id, status FROM {table} WHERE operation_id = $1 AND bank_id = $2 FOR UPDATE", uuid.UUID(parent_operation_id), bank_id, ) if not parent_row: return + if parent_row["status"] not in {"pending", "processing"}: + return # Check whether all siblings are done. Pull error_message too so a # parent that fails can inherit a representative child reason -- @@ -535,38 +539,54 @@ async def _maybe_update_parent_operation(self, child_operation_id: str, schema: f""" SELECT status, error_message FROM {table} WHERE bank_id = $1 - AND result_metadata::jsonb @> $2::jsonb + AND (result_metadata->>'parent_operation_id')::uuid = $2 """, bank_id, - json.dumps({"parent_operation_id": parent_operation_id}), + uuid.UUID(parent_operation_id), ) - if not siblings or not all(s["status"] in ("completed", "failed") for s in siblings): + if not siblings or not all(s["status"] in ("completed", "failed", "cancelled") for s in siblings): return any_failed = any(s["status"] == "failed" for s in siblings) + any_cancelled = any(s["status"] == "cancelled" for s in siblings) if any_failed: await conn.execute( f""" UPDATE {table} SET status = 'failed', error_message = $2, updated_at = now() - WHERE operation_id = $1 + WHERE operation_id = $1 AND bank_id = $3 + AND status IN ('pending', 'processing') """, uuid.UUID(parent_operation_id), _summarise_child_error_messages(siblings), + bank_id, + ) + new_status = "failed" + elif any_cancelled: + await conn.execute( + f""" + UPDATE {table} + SET status = 'cancelled', updated_at = now() + WHERE operation_id = $1 AND bank_id = $2 + AND status IN ('pending', 'processing') + """, + uuid.UUID(parent_operation_id), + bank_id, ) + new_status = "cancelled" else: await conn.execute( f""" UPDATE {table} SET status = 'completed', updated_at = now(), completed_at = now() - WHERE operation_id = $1 + WHERE operation_id = $1 AND bank_id = $2 + AND status IN ('pending', 'processing') """, uuid.UUID(parent_operation_id), + bank_id, ) - logger.info( - f"Poller updated parent operation {parent_operation_id} to " - f"{'failed' if any_failed else 'completed'} (all siblings done)" - ) + new_status = "completed" + logger.info(f"Poller updated parent operation {parent_operation_id} to {new_status} (all siblings done)") except Exception as e: # Log but don't re-raise — the child has already been marked failed, # which is the critical state change. A stuck parent will be caught on @@ -583,7 +603,7 @@ async def _schedule_retry(self, operation_id: str, retry_at: "Any", error_messag UPDATE {table} SET status = 'pending', next_retry_at = $2, worker_id = NULL, claimed_at = NULL, retry_count = retry_count + 1, error_message = $3, updated_at = now() - WHERE operation_id = $1 + WHERE operation_id = $1 AND status = 'processing' """, operation_id, retry_at, @@ -604,7 +624,7 @@ async def _defer_operation(self, operation_id: str, exec_date: "Any", reason: st UPDATE {table} SET status = 'pending', next_retry_at = $2, worker_id = NULL, claimed_at = NULL, updated_at = now() - WHERE operation_id = $1 + WHERE operation_id = $1 AND status = 'processing' """, operation_id, exec_date, diff --git a/core/dataplane/pyproject.toml b/core/dataplane/pyproject.toml index 4d25a03..e241e8d 100644 --- a/core/dataplane/pyproject.toml +++ b/core/dataplane/pyproject.toml @@ -8,6 +8,8 @@ version = "0.6.1" description = "HMS: Agent Memory That Works Like Human Memory" readme = "README.md" requires-python = ">=3.11" +license = "MIT" +license-files = ["LICENSE", "THIRD_PARTY_NOTICES.md"] dependencies = [ "asyncpg>=0.29.0", "python-dotenv>=1.0.0", diff --git a/core/dataplane/tests/e2e/test_multimodal_security_e2e.py b/core/dataplane/tests/e2e/test_multimodal_security_e2e.py index 8af086d..cd88d19 100644 --- a/core/dataplane/tests/e2e/test_multimodal_security_e2e.py +++ b/core/dataplane/tests/e2e/test_multimodal_security_e2e.py @@ -192,12 +192,12 @@ def handler(request: httpx.Request) -> httpx.Response: and "multimodal" in record.getMessage().lower() ) assert multimodal_worker_logs - orchestrator_logs = "\n".join( - record.getMessage() for record in caplog.records if record.name == "hms_api.engine.retain.orchestrator" + ingestion_logs = "\n".join( + record.getMessage() + for record in caplog.records + if record.name.startswith("hms_api.engine.ingestion") ) - assert orchestrator_logs - assert "" in orchestrator_logs - multimodal_pipeline_logs = f"{multimodal_worker_logs}\n{orchestrator_logs}" + multimodal_pipeline_logs = f"{multimodal_worker_logs}\n{ingestion_logs}" for forbidden in (bank_id, document_id, normalized.asset.sha256, "security-surface.png"): assert forbidden not in multimodal_pipeline_logs diff --git a/core/dataplane/tests/test_async_batch_retain.py b/core/dataplane/tests/test_async_batch_retain.py index bc014b9..6499059 100644 --- a/core/dataplane/tests/test_async_batch_retain.py +++ b/core/dataplane/tests/test_async_batch_retain.py @@ -1,11 +1,13 @@ """Test async batch retain with smart batching and parent-child operations.""" import asyncio +import hashlib import json import uuid import pytest - +from hms_api.engine.cross_encoder import RRFPassthroughCrossEncoder +from hms_api.engine.embeddings import Embeddings from hms_api.extensions import RequestContext # These tests submit async operations and rely on the engine-owned worker to @@ -16,6 +18,42 @@ pytestmark = pytest.mark.xdist_group("worker_tests") +class _DeterministicEmbeddings(Embeddings): + """Provide stable vectors without loading a local model.""" + + model_name = "hms-async-batch-test-hash-v1" + + @property + def provider_name(self) -> str: + return "async-batch-test" + + @property + def dimension(self) -> int: + return 384 + + async def initialize(self) -> None: + return None + + def encode(self, texts: list[str]) -> list[list[float]]: + vectors: list[list[float]] = [] + for text in texts: + digest = hashlib.sha256(text.encode()).digest() + vectors.append([((digest[index % len(digest)] / 255.0) * 2.0) - 1.0 for index in range(self.dimension)]) + return vectors + + +@pytest.fixture(scope="session") +def embeddings() -> Embeddings: + """Use deterministic embeddings that need no optional ML dependencies.""" + return _DeterministicEmbeddings() + + +@pytest.fixture(scope="session") +def cross_encoder() -> RRFPassthroughCrossEncoder: + """Use the dependency-free reciprocal-rank fusion reranker.""" + return RRFPassthroughCrossEncoder() + + async def _ensure_bank(pool, bank_id: str) -> None: """Upsert a minimal bank row so FK on async_operations passes.""" await pool.execute( @@ -26,40 +64,47 @@ async def _ensure_bank(pool, bank_id: str) -> None: @pytest.mark.asyncio -async def test_duplicate_document_ids_rejected_async(memory, request_context): - """Test that async retain rejects batches with duplicate document_ids.""" - bank_id = "test_duplicate_async" +async def test_repeated_document_ids_stay_in_one_async_document_group(memory, request_context): + """Repeated IDs remain one logical document and are never split across children.""" + bank_id = f"test_repeated_async_{uuid.uuid4().hex[:8]}" contents = [ {"content": "First item", "document_id": "doc1"}, {"content": "Second item", "document_id": "doc2"}, - {"content": "Third item", "document_id": "doc1"}, # Duplicate! + {"content": "Third item", "document_id": "doc1"}, ] - # Should raise ValueError due to duplicate document_ids - with pytest.raises(ValueError, match="duplicate document_ids.*doc1"): - await memory.submit_async_retain( - bank_id=bank_id, - contents=contents, - request_context=request_context, - ) + result = await memory.submit_async_retain( + bank_id=bank_id, + contents=contents, + request_context=request_context, + ) + status = await memory.get_operation_status( + bank_id=bank_id, + operation_id=result["operation_id"], + request_context=request_context, + ) + + assert status["status"] == "completed" + assert status["result_metadata"]["num_sub_batches"] == 1 + assert status["child_operations"][0]["items_count"] == 3 @pytest.mark.asyncio -async def test_duplicate_document_ids_rejected_sync(memory, request_context): - """Test that sync retain also rejects batches with duplicate document_ids.""" - bank_id = "test_duplicate_sync" +async def test_repeated_document_ids_are_supported_sync(memory, request_context): + """The public sync API maps every repeated-ID input back to its own result slot.""" + bank_id = f"test_repeated_sync_{uuid.uuid4().hex[:8]}" contents = [ {"content": "First item", "document_id": "doc1"}, - {"content": "Second item", "document_id": "doc1"}, # Duplicate! + {"content": "Second item", "document_id": "doc1"}, ] - # Should raise ValueError due to duplicate document_ids - with pytest.raises(ValueError, match="duplicate document_ids.*doc1"): - await memory.retain_batch_async( - bank_id=bank_id, - contents=contents, - request_context=request_context, - ) + results = await memory.retain_batch_async( + bank_id=bank_id, + contents=contents, + request_context=request_context, + ) + + assert len(results) == 2 @pytest.mark.asyncio diff --git a/core/dataplane/tests/test_config_validation.py b/core/dataplane/tests/test_config_validation.py index 00167c2..f711c21 100644 --- a/core/dataplane/tests/test_config_validation.py +++ b/core/dataplane/tests/test_config_validation.py @@ -23,6 +23,7 @@ def setup_test_env(): "HMS_API_LLM_MODEL", "HMS_API_DATABASE_URL", "HMS_API_MIGRATION_DATABASE_URL", + "HMS_API_RETAIN_EMBEDDING_FAILURE_POLICY", ] # Save original values @@ -102,6 +103,40 @@ def test_valid_retain_config_succeeds(): assert config.retain_chunk_size == 3000 +def test_retain_embedding_failure_policy_defaults_to_nonfatal_storage(monkeypatch): + """Embedding outages remain non-fatal unless fail-closed is explicitly enabled.""" + + from hms_api.config import HMSConfig + + monkeypatch.delenv("HMS_API_RETAIN_EMBEDDING_FAILURE_POLICY", raising=False) + monkeypatch.setenv("HMS_API_LLM_PROVIDER", "mock") + + config = HMSConfig.from_env() + + assert config.retain_embedding_failure_policy == "store_without_embedding" + + +def test_retain_embedding_failure_policy_accepts_raise(monkeypatch): + from hms_api.config import HMSConfig + + monkeypatch.setenv("HMS_API_RETAIN_EMBEDDING_FAILURE_POLICY", "RAISE") + monkeypatch.setenv("HMS_API_LLM_PROVIDER", "mock") + + config = HMSConfig.from_env() + + assert config.retain_embedding_failure_policy == "raise" + + +def test_retain_embedding_failure_policy_rejects_unknown_value(monkeypatch): + from hms_api.config import HMSConfig + + monkeypatch.setenv("HMS_API_RETAIN_EMBEDDING_FAILURE_POLICY", "ignore") + monkeypatch.setenv("HMS_API_LLM_PROVIDER", "mock") + + with pytest.raises(ValueError, match="HMS_API_RETAIN_EMBEDDING_FAILURE_POLICY"): + HMSConfig.from_env() + + def test_log_config_masks_database_urls(caplog): """Config startup logs must not expose database credentials.""" from hms_api.config import HMSConfig diff --git a/core/dataplane/tests/test_db_abstraction.py b/core/dataplane/tests/test_db_abstraction.py index 68fa90c..6bd1a9d 100644 --- a/core/dataplane/tests/test_db_abstraction.py +++ b/core/dataplane/tests/test_db_abstraction.py @@ -5,6 +5,8 @@ """ import json +import uuid +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -469,6 +471,83 @@ def test_default_database_backend(self): assert DEFAULT_DATABASE_BACKEND == "postgresql" +# --------------------------------------------------------------------------- +# PostgreSQLOps fact identity tests (mock DatabaseConnection, no live DB) +# --------------------------------------------------------------------------- + + +class TestPostgreSQLOpsInsertFactsBatch: + """Fact IDs must be assigned before SQL, in exact input order.""" + + @staticmethod + def _make_batch(n: int = 3) -> dict: + return dict( + bank_id="bank-pg", + fact_texts=[f"fact-{index}" for index in range(n)], + embeddings=[None] * n, + event_dates=[None] * n, + occurred_starts=[None] * n, + occurred_ends=[None] * n, + mentioned_ats=[None] * n, + contexts=[f"context-{index}" for index in range(n)], + fact_types=["world"] * n, + metadata_jsons=["{}"] * n, + chunk_ids=[f"chunk-{index}" for index in range(n)], + document_ids=["doc-pg"] * n, + tags_list=["[]"] * n, + observation_scopes_list=[None] * n, + text_signals_list=[None] * n, + projection_jsons=["{}"] * n, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("text_search_extension", ["native", "vchord"]) + async def test_client_generated_ids_are_inserted_and_returned_in_input_order( + self, + text_search_extension, + ): + from hms_api.engine.db.ops_postgresql import PostgreSQLOps + + generated = [ + uuid.UUID("00000000-0000-0000-0000-000000000003"), + uuid.UUID("00000000-0000-0000-0000-000000000001"), + uuid.UUID("00000000-0000-0000-0000-000000000002"), + ] + config = SimpleNamespace( + database_backend="postgresql", + database_schema="public", + text_search_extension=text_search_extension, + ) + connection = AsyncMock(spec=DatabaseConnection) + batch = self._make_batch() + + with ( + patch("hms_api.engine.db.ops_postgresql.uuid4", side_effect=generated), + patch("hms_api.config.get_config", return_value=config), + patch("hms_api.engine.schema.get_config", return_value=config), + patch("hms_api.engine.memory_engine.get_config", return_value=config), + ): + result = await PostgreSQLOps().insert_facts_batch( + conn=connection, + text_search_extension=text_search_extension, + **batch, + ) + + connection.execute.assert_awaited_once() + connection.fetch.assert_not_awaited() + query, bank_id, inserted_ids, fact_texts, *remaining = connection.execute.await_args.args + assert bank_id == "bank-pg" + assert inserted_ids == generated + assert fact_texts == batch["fact_texts"] + assert remaining[-1] == batch["projection_jsons"] + assert result == [str(value) for value in generated] + assert "$2::uuid[]" in query + assert "AS t(id, text, embedding" in query + assert "INSERT INTO public.memory_units (id, bank_id" in query + assert "RETURNING id" not in query + assert ("bm25_catalog.bm25vector" in query) is (text_search_extension == "vchord") + + # --------------------------------------------------------------------------- # OracleOps unit tests (mock DatabaseConnection, no live DB) # --------------------------------------------------------------------------- diff --git a/core/dataplane/tests/test_delta_retain.py b/core/dataplane/tests/test_delta_retain.py index 3075c8d..a5100aa 100644 --- a/core/dataplane/tests/test_delta_retain.py +++ b/core/dataplane/tests/test_delta_retain.py @@ -888,21 +888,6 @@ async def test_delta_retain_recall_with_chunks(memory, request_context): # content+context tokens of the chunks that were actually processed). -def test_merge_processed_content_tokens_helper(): - """Unit check on the None-propagating aggregator used by the engine.""" - from hms_api.engine.retain.orchestrator import ( - _merge_processed_content_tokens, - ) - - assert _merge_processed_content_tokens(0, 0) == 0 - assert _merge_processed_content_tokens(5, 7) == 12 - # None "wins" in either slot — once any sub-result bypassed dedup, the - # aggregate is None so callers bill full content. - assert _merge_processed_content_tokens(None, 10) is None - assert _merge_processed_content_tokens(10, None) is None - assert _merge_processed_content_tokens(None, None) is None - - @pytest.mark.asyncio async def test_processed_content_tokens_first_retain_is_none(memory, request_context): """ diff --git a/core/dataplane/tests/test_fact_extraction_retry.py b/core/dataplane/tests/test_fact_extraction_retry.py index 8b5dfb8..9c19eae 100644 --- a/core/dataplane/tests/test_fact_extraction_retry.py +++ b/core/dataplane/tests/test_fact_extraction_retry.py @@ -7,6 +7,7 @@ """ from datetime import datetime, timezone +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -42,6 +43,36 @@ def _make_llm_config(mock_response): return llm +@pytest.mark.asyncio +async def test_chunk_failure_summary_preserves_original_exception(monkeypatch): + from hms_api.engine.retain import fact_extraction + + provider_error = TimeoutError("gateway connect timed out") + + monkeypatch.setattr(fact_extraction, "chunk_text", lambda *_args, **_kwargs: ["failed chunk", "ok chunk"]) + + async def extract_chunk(*, chunk_index, **_kwargs): + if chunk_index == 0: + raise provider_error + return [], MagicMock() + + monkeypatch.setattr(fact_extraction, "_extract_facts_with_auto_split", extract_chunk) + + with pytest.raises( + RuntimeError, + match=r"chunk 0: TimeoutError: gateway connect timed out", + ) as raised: + await fact_extraction.extract_facts_from_text( + text="document", + event_date=None, + llm_config=object(), + agent_name="test-agent", + config=SimpleNamespace(retain_chunk_size=1000), + ) + + assert raised.value.__cause__ is provider_error + + @pytest.mark.asyncio async def test_non_dict_json_all_retries_returns_empty(): """ diff --git a/core/dataplane/tests/test_file_retain.py b/core/dataplane/tests/test_file_retain.py index 68a3b83..9d16224 100644 --- a/core/dataplane/tests/test_file_retain.py +++ b/core/dataplane/tests/test_file_retain.py @@ -547,7 +547,7 @@ def name(self) -> str: async def test_file_retain_maps_timestamp_to_event_date(memory_no_llm_verify, sample_txt_content): """Regression (PR #1092): file retain must translate 'timestamp' -> 'event_date'. - The retain orchestrator only reads 'event_date' from each content dict. + The Retain pipeline only reads 'event_date' from each content dict. _handle_file_convert_retain previously forwarded 'timestamp' unchanged, so every file-retained memory silently defaulted to utcnow() and the 'unset' sentinel was a no-op. This test intercepts the inner batch_retain task the handler diff --git a/core/dataplane/tests/test_ingestion_oracle_contracts.py b/core/dataplane/tests/test_ingestion_oracle_contracts.py new file mode 100644 index 0000000..8ad8a8f --- /dev/null +++ b/core/dataplane/tests/test_ingestion_oracle_contracts.py @@ -0,0 +1,602 @@ +"""Offline Oracle contracts for the Retain ingestion pipeline. + +These tests exercise adapter selection and Oracle-specific SQL boundaries with +fakes. They do not replace validation against an Oracle 23ai instance. +""" + +from __future__ import annotations + +import copy +import json +import uuid +from array import array +from contextlib import asynccontextmanager +from types import SimpleNamespace + +import pytest +from hms_api.engine.db.ops_oracle import OracleOps +from hms_api.engine.db.oracle import _convert_arg, _rewrite_pg_to_oracle +from hms_api.engine.entity_resolution_contracts import ( + EntityOccurrenceBinding, + EntityResolutionReadPlan, + ExistingEntityBinding, +) +from hms_api.engine.entity_resolver import EntityResolver +from hms_api.engine.ingestion import runtime as runtime_module +from hms_api.engine.ingestion import service as service_module +from hms_api.engine.ingestion.adapters.oracle_semantic import ( + compute_oracle_semantic_links_ann, +) +from hms_api.engine.ingestion.adapters.postgres_fresh_ownership import ( + FreshDocumentOwnershipConflict, + FreshPostgresDocumentOwnership, +) +from hms_api.engine.ingestion.persistence.backend import retain_backend_adapters +from hms_api.engine.ingestion.persistence.operation_fence import OperationActivityFence +from hms_api.engine.ingestion.persistence.oracle import ( + FreshOracleDocumentOwnership, + OracleCheckpointStore, + OracleDocumentOwnership, + OraclePlanningRepository, +) +from hms_api.engine.ingestion.persistence.postgres import ( + PostgresCheckpointStore, + PostgresDocumentOwnership, + PostgresPlanningRepository, +) +from hms_api.engine.retain import chunk_storage + + +class _Transaction: + def __init__(self, events: list[str]) -> None: + self._events = events + + async def __aenter__(self): + self._events.append("begin") + return self + + async def __aexit__(self, exc_type, _exc, _traceback): + self._events.append("rollback" if exc_type is not None else "commit") + return False + + +class _SnapshotConnection: + def __init__(self) -> None: + self.events: list[str] = [] + + def transaction(self) -> _Transaction: + return _Transaction(self.events) + + async def execute(self, query: str, *_args): + self.events.append(" ".join(query.split())) + return "OK 0" + + +def test_backend_selector_exposes_matching_retain_adapters() -> None: + postgres = retain_backend_adapters("postgresql") + oracle = retain_backend_adapters("ORACLE") + + assert isinstance(postgres.planning_repository(object()), PostgresPlanningRepository) + assert isinstance(postgres.checkpoint_store(object()), PostgresCheckpointStore) + assert isinstance(postgres.document_ownership(), PostgresDocumentOwnership) + assert isinstance(postgres.document_ownership(fresh=True), FreshPostgresDocumentOwnership) + assert isinstance(postgres.operation_activity_fence(str(uuid.uuid4())), OperationActivityFence) + assert postgres.operation_activity_fence(None) is None + + assert isinstance(oracle.planning_repository(object()), OraclePlanningRepository) + assert isinstance(oracle.checkpoint_store(object()), OracleCheckpointStore) + assert isinstance(oracle.document_ownership(), OracleDocumentOwnership) + assert isinstance(oracle.document_ownership(fresh=True), FreshOracleDocumentOwnership) + assert isinstance(oracle.operation_activity_fence(str(uuid.uuid4())), OperationActivityFence) + assert oracle.operation_activity_fence(None) is None + + with pytest.raises(ValueError, match="Unsupported Retain database backend"): + retain_backend_adapters("sqlite") + + +@pytest.mark.asyncio +async def test_backend_snapshots_preserve_database_transaction_rules() -> None: + postgres_connection = _SnapshotConnection() + async with retain_backend_adapters("postgresql").planning_snapshot(postgres_connection): + postgres_connection.events.append("body") + assert postgres_connection.events == [ + "begin", + "SET TRANSACTION ISOLATION LEVEL REPEATABLE READ READ ONLY", + "body", + "commit", + ] + + oracle_connection = _SnapshotConnection() + async with retain_backend_adapters("oracle").planning_snapshot(oracle_connection): + oracle_connection.events.append("body") + assert oracle_connection.events == [ + "SET TRANSACTION READ ONLY", + "body", + ] + + +def test_service_route_accepts_oracle_and_rejects_unknown_backends() -> None: + oracle_execution = SimpleNamespace( + pool=SimpleNamespace(backend_type="oracle"), + resolved_config=SimpleNamespace( + database_backend="oracle", + retain_extraction_mode="chunks", + ), + ) + service_module._require_supported_route(oracle_execution) + + unknown_execution = SimpleNamespace( + pool=SimpleNamespace(backend_type="sqlite"), + resolved_config=SimpleNamespace( + database_backend="sqlite", + retain_extraction_mode="chunks", + ), + ) + with pytest.raises(service_module.RetainDatabaseUnsupportedError): + service_module._require_supported_route(unknown_execution) + + +def test_oracle_rewrites_json_uuid_text_comparisons_to_raw_uuid_comparisons() -> None: + parent_operation_id = uuid.uuid4() + + rewritten, ignore_duplicate, returning_columns = _rewrite_pg_to_oracle( + """ + SELECT operation_id + FROM async_operations + WHERE bank_id = $1 + AND (result_metadata->>'parent_operation_id')::uuid = $2 + """ + ) + + assert "HEXTORAW(REPLACE(JSON_VALUE(result_metadata, '$.parent_operation_id'), '-', '')) = :2" in rewritten + assert "->>" not in rewritten + assert "::uuid" not in rewritten + assert ignore_duplicate is False + assert returning_columns is None + assert _convert_arg(parent_operation_id) == parent_operation_id.bytes + assert _convert_arg(str(parent_operation_id)) == parent_operation_id.bytes + + +def test_oracle_preserves_numeric_ordering_for_json_text_casts() -> None: + rewritten, _, _ = _rewrite_pg_to_oracle( + "SELECT operation_id FROM async_operations ORDER BY (result_metadata->>'sub_batch_index')::int" + ) + + assert "ORDER BY TO_NUMBER(JSON_VALUE(result_metadata, '$.sub_batch_index'))" in rewritten + assert "->>" not in rewritten + assert "::int" not in rewritten + + +class _PlanningConnection: + def __init__(self, rows: list[dict]) -> None: + self.rows = rows + self.query = "" + self.args: tuple = () + + async def fetch(self, query: str, *args): + self.query = query + self.args = args + return self.rows + + +@pytest.mark.asyncio +async def test_oracle_planning_reorders_checkpoint_bindings_in_python() -> None: + first_id = uuid.uuid4() + second_id = uuid.uuid4() + connection = _PlanningConnection( + [ + {"unit_id": first_id, "chunk_index": 0}, + {"unit_id": second_id, "chunk_index": 1}, + ] + ) + repository = OraclePlanningRepository(connection, schema="tenant") + + bindings = await repository.load_document_unit_bindings( + "bank", + "document", + expected_unit_ids=(str(second_id), str(first_id)), + ) + + assert [binding.unit_id for binding in bindings] == [str(second_id), str(first_id)] + assert [binding.chunk_index for binding in bindings] == [1, 0] + assert connection.args == ("bank", "document") + assert "array_position" not in connection.query + assert "ANY(" not in connection.query + assert '"tenant".memory_units' in connection.query + + with pytest.raises(ValueError, match="checkpoint unit IDs are missing"): + await repository.load_document_unit_bindings( + "bank", + "document", + expected_unit_ids=(str(uuid.uuid4()),), + ) + + +class _OwnershipConnection: + def __init__(self, *, insert_status: str = "INSERT 0 1", locked_value: str | None = "document") -> None: + self.insert_status = insert_status + self.locked_value = locked_value + self.queries: list[str] = [] + + async def execute(self, query: str, *_args): + self.queries.append(query) + if query.lstrip().startswith("INSERT"): + return self.insert_status + return "UPDATE 1" + + async def fetchval(self, query: str, *_args): + self.queries.append(query) + return self.locked_value + + +class _UnhashedOwnershipConnection: + def __init__(self, row) -> None: + self.row = row + self.query = "" + self.args: tuple = () + + async def fetchrow(self, query: str, *args): + self.query = query + self.args = args + return self.row + + +@pytest.mark.asyncio +@pytest.mark.parametrize("ownership_type", [PostgresDocumentOwnership, OracleDocumentOwnership]) +@pytest.mark.parametrize( + "row,expected", + [ + ({"content_hash": None}, True), + ({"content_hash": ""}, True), + ({"content_hash": "new-hash"}, False), + (None, False), + ], +) +async def test_unhashed_ownership_requires_a_locked_existing_row(ownership_type, row, expected) -> None: + connection = _UnhashedOwnershipConnection(row) + + result = await ownership_type(schema="tenant").validate_unhashed_window( + connection, + bank_id="bank", + document_id="document", + ) + + assert result is expected + assert "FOR UPDATE" in connection.query + assert '"tenant".documents' in connection.query + assert connection.args == ("document", "bank") + + +@pytest.mark.asyncio +async def test_oracle_fresh_ownership_claim_and_transition_use_row_counts() -> None: + connection = _OwnershipConnection() + ownership = FreshOracleDocumentOwnership(schema="tenant") + + await ownership.prepare_first_window( + connection, + bank_id="bank", + document_id="document", + ) + transitioned = await ownership.transition_content_hash( + connection, + bank_id="bank", + document_id="document", + expected_content_hash="old", + new_content_hash="new", + ) + + assert transitioned is True + assert "FOR UPDATE" in connection.queries[1] + assert "RETURNING" not in connection.queries[2] + + conflict_connection = _OwnershipConnection(insert_status="INSERT 0 0") + with pytest.raises(FreshDocumentOwnershipConflict): + await ownership.prepare_first_window( + conflict_connection, + bank_id="bank", + document_id="document", + ) + + +class _CheckpointConnection: + def __init__(self, metadata: dict) -> None: + self.metadata = copy.deepcopy(metadata) + self.events: list[str] = [] + self.queries: list[str] = [] + + def transaction(self) -> _Transaction: + return _Transaction(self.events) + + async def fetchrow(self, query: str, *_args): + self.queries.append(query) + return {"result_metadata": json.dumps(self.metadata)} + + async def execute(self, query: str, *args): + self.queries.append(query) + self.metadata = json.loads(args[0]) + return "UPDATE 1" + + +@pytest.mark.asyncio +async def test_oracle_checkpoint_read_modify_write_preserves_resume_contract() -> None: + operation_id = str(uuid.uuid4()) + connection = _CheckpointConnection( + { + "keep": {"caller": "value"}, + "batch_id": "provider-job", + "batch_provider": "provider", + "chunk_count": 2, + } + ) + store = OracleCheckpointStore(connection, schema="tenant") + + await store.record_document_id(operation_id, "document-a") + await store.record_document_id(operation_id, "document-a") + await store.record_core_committed( + operation_id, + "document-a", + unit_ids=("unit-b", "unit-a"), + requires_final_ann=True, + ) + checkpoint = await store.recover(operation_id) + + assert checkpoint.document_ids == ("document-a",) + assert checkpoint.core_committed_document_ids == ("document-a",) + assert checkpoint.final_ann_pending_document_ids == ("document-a",) + assert checkpoint.unit_ids_for_document("document-a") == ("unit-b", "unit-a") + assert checkpoint.unscoped_facts_committed is True + assert connection.metadata["keep"] == {"caller": "value"} + + await store.record_core_committed( + operation_id, + "document-empty", + unit_ids=(), + requires_final_ann=True, + ) + checkpoint = await store.recover(operation_id) + assert checkpoint.unit_ids_for_document("document-empty") == () + assert checkpoint.final_ann_pending_document_ids == ("document-a",) + + await store.record_final_ann_completed(operation_id, "document-a") + await store.clear_provider_batch(operation_id) + checkpoint = await store.recover(operation_id) + + assert checkpoint.final_ann_pending_document_ids == () + assert connection.metadata["keep"] == {"caller": "value"} + assert "batch_id" not in connection.metadata + assert "batch_provider" not in connection.metadata + assert "chunk_count" not in connection.metadata + assert all("jsonb_" not in query.lower() for query in connection.queries) + assert any("FOR UPDATE" in query for query in connection.queries) + + +class _SemanticConnection: + def __init__(self) -> None: + self.calls: list[tuple[str, tuple]] = [] + + async def fetch(self, query: str, *args): + self.calls.append((query, args)) + return [ + {"to_id": uuid.uuid4(), "similarity": 0.82}, + {"to_id": uuid.uuid4(), "similarity": 0.64}, + ] + + +@pytest.mark.asyncio +async def test_oracle_semantic_ann_uses_vector_selects_without_postgres_ddl() -> None: + connection = _SemanticConnection() + + links = await compute_oracle_semantic_links_ann( + connection, + "bank", + ["seed"], + [array("f", [0.1, 0.2, 0.3])], + fact_types=["world"], + top_k=20, + threshold=0.7, + ) + + assert len(links) == 1 + assert links[0][0] == "seed" + assert isinstance(links[0][1], str) + query, args = connection.calls[0] + assert "VECTOR_DISTANCE" in query + assert "FETCH FIRST 20 ROWS ONLY" in query + assert args[:2] == ("bank", "world") + assert isinstance(args[2], array) + assert list(args[2]) == pytest.approx([0.1, 0.2, 0.3]) + assert all(token not in query for token in ("SET LOCAL", "unnest", "TEMP TABLE", "LATERAL")) + + +@pytest.mark.asyncio +async def test_oracle_entity_lookup_preserves_original_input_name() -> None: + entity_id = uuid.uuid4() + + class _EntityConnection: + async def fetchrow(self, _query: str, _bank_id: str, input_name: str): + return { + "id": entity_id, + "name_lower": input_name.lower(), + } + + rows = await OracleOps().fetch_missing_entity_ids( + _EntityConnection(), + "entities", + "bank", + ["Mixed Case"], + ) + + assert len(rows) == 1 + assert rows[0]["id"] == entity_id + assert rows[0]["name_lower"] == "mixed case" + assert rows[0]["input_name"] == "Mixed Case" + + +class _ExistingEntityValidationConnection: + def __init__(self, backend_type: str) -> None: + self.backend_type = backend_type + self.calls: list[tuple[str, tuple]] = [] + + async def fetch(self, query: str, *args): + self.calls.append((query, args)) + if self.backend_type == "oracle": + assert len(args[1]) <= 900 + return [{"id": entity_id} for entity_id in args[1]] + + +@pytest.mark.parametrize( + ("backend_type", "expected_batch_sizes"), + (("oracle", [900, 101]), ("postgresql", [1001])), +) +@pytest.mark.asyncio +async def test_existing_entity_validation_chunks_only_oracle_bind_lists( + backend_type: str, + expected_batch_sizes: list[int], +) -> None: + entity_ids = [str(uuid.uuid4()) for _ in range(1001)] + occurrences = tuple( + EntityOccurrenceBinding( + occurrence_key=f"occurrence-{index}", + unit_key=f"unit-{index}", + local_index=0, + event_date=None, + ) + for index in range(len(entity_ids)) + ) + plan = EntityResolutionReadPlan( + bank_id="bank", + occurrences=occurrences, + existing_bindings=tuple( + ExistingEntityBinding( + occurrence_key=occurrence.occurrence_key, + entity_id=entity_id, + ) + for occurrence, entity_id in zip(occurrences, entity_ids, strict=True) + ), + ) + connection = _ExistingEntityValidationConnection(backend_type) + resolver = EntityResolver(SimpleNamespace(ops=OracleOps())) + + finalized = await resolver.finalize_entity_read_plan( + connection, + "bank", + plan, + entities_table="entities", + ) + + assert [len(call_args[1]) for _query, call_args in connection.calls] == expected_batch_sizes + assert list(finalized.resolved_entity_ids) == entity_ids + + +class _ChunkDeletionConnection: + def __init__(self, backend_type: str) -> None: + self.backend_type = backend_type + self.calls: list[tuple[str, tuple]] = [] + + async def execute(self, query: str, *args): + self.calls.append((query, args)) + if self.backend_type == "oracle": + assert len(args[0]) <= 900 + return f"DELETE {len(args[0])}" + + +@pytest.mark.parametrize( + ("backend_type", "expected_batch_sizes"), + (("oracle", [900, 101]), ("postgresql", [1001])), +) +@pytest.mark.asyncio +async def test_delta_chunk_deletion_chunks_only_oracle_bind_lists( + backend_type: str, + expected_batch_sizes: list[int], +) -> None: + connection = _ChunkDeletionConnection(backend_type) + + await chunk_storage.delete_chunks_by_ids( + connection, + [f"chunk-{index}" for index in range(1001)], + ) + + assert [len(call_args[0]) for _query, call_args in connection.calls] == expected_batch_sizes + assert all("ANY($1::text[])" in query for query, _call_args in connection.calls) + + +@pytest.mark.asyncio +async def test_runtime_oracle_planning_uses_read_only_snapshot_and_oracle_ann(monkeypatch) -> None: + connection = _SnapshotConnection() + oracle_ann_calls: list[tuple] = [] + + @asynccontextmanager + async def acquire(_pool): + yield connection + + async def plan_entities(*_args, **_kwargs): + return SimpleNamespace(occurrences=()) + + async def oracle_ann(*args, **kwargs): + oracle_ann_calls.append((args, kwargs)) + return [("0", "neighbor", "semantic", 0.9, None)] + + async def postgres_ann(*_args, **_kwargs): + raise AssertionError("PostgreSQL ANN must not run for Oracle") + + monkeypatch.setattr(runtime_module, "acquire_with_retry", acquire) + monkeypatch.setattr(runtime_module.entity_processing, "plan_entities", plan_entities) + monkeypatch.setattr(runtime_module, "compute_oracle_semantic_links_ann", oracle_ann) + monkeypatch.setattr(runtime_module, "compute_semantic_links_ann", postgres_ann) + + result = await runtime_module.pre_resolve_entities( + SimpleNamespace(backend_type="oracle"), + object(), + "bank", + [SimpleNamespace(entities=[])], + ["fact-key"], + [SimpleNamespace(embedding=[0.1, 0.2], fact_type="world")], + SimpleNamespace(entity_labels=None), + [], + ) + + assert connection.events == ["SET TRANSACTION READ ONLY"] + assert len(oracle_ann_calls) == 1 + assert result.semantic_ann_links == [("0", "neighbor", "semantic", 0.9, None)] + + +@pytest.mark.asyncio +async def test_runtime_final_oracle_ann_maps_driver_uuid_rows_to_string_ids(monkeypatch) -> None: + unit_id = str(uuid.uuid4()) + inserted_links: list[tuple] = [] + + class _FinalAnnConnection: + async def fetch(self, query: str, *_args): + if "VECTOR_DISTANCE" in query: + return [{"to_id": uuid.uuid4(), "similarity": 0.91}] + return [ + { + "id": uuid.UUID(unit_id), + "embedding": array("f", [0.1, 0.2]), + "fact_type": "world", + } + ] + + connection = _FinalAnnConnection() + + @asynccontextmanager + async def acquire(_pool): + yield connection + + async def insert_links(_connection, links, **_kwargs): + inserted_links.extend(links) + + monkeypatch.setattr(runtime_module, "acquire_with_retry", acquire) + monkeypatch.setattr(runtime_module, "_bulk_insert_links", insert_links) + + await runtime_module.run_final_semantic_ann( + SimpleNamespace(backend_type="oracle", ops=object()), + "bank", + [unit_id], + SimpleNamespace(write_semantic_links=True), + [], + ) + + assert len(inserted_links) == 1 + assert inserted_links[0][0] == unit_id + assert isinstance(inserted_links[0][1], str) diff --git a/core/dataplane/tests/test_ingestion_oracle_live.py b/core/dataplane/tests/test_ingestion_oracle_live.py new file mode 100644 index 0000000..a4cac8d --- /dev/null +++ b/core/dataplane/tests/test_ingestion_oracle_live.py @@ -0,0 +1,208 @@ +"""Live Oracle 23ai smoke coverage for the Retain ingestion pipeline.""" + +from __future__ import annotations + +import hashlib +import importlib.util +import os +import uuid + +import pytest +from hms_api import MemoryEngine, RequestContext +from hms_api.config import clear_config_cache +from hms_api.engine.cross_encoder import RRFPassthroughCrossEncoder +from hms_api.engine.embeddings import Embeddings +from hms_api.engine.memory_engine import Budget +from hms_api.engine.query_analyzer import DateparserQueryAnalyzer +from hms_api.engine.task_backend import SyncTaskBackend, WorkerTaskBackend + +pytestmark = [ + pytest.mark.oracle, + pytest.mark.skipif( + importlib.util.find_spec("oracledb") is None, + reason="the oracle optional dependency is required", + ), + pytest.mark.skipif( + not os.getenv("ORACLE_TEST_DSN"), + reason="ORACLE_TEST_DSN is required for live Oracle tests", + ), +] + + +class _DeterministicEmbeddings(Embeddings): + """Small deterministic vectors with the production schema dimension.""" + + model_name = "hms-oracle-live-hash-v1" + + @property + def provider_name(self) -> str: + return "oracle-live-test" + + @property + def dimension(self) -> int: + return 384 + + async def initialize(self) -> None: + return None + + def encode(self, texts: list[str]) -> list[list[float]]: + vectors: list[list[float]] = [] + for text in texts: + digest = hashlib.sha256(text.encode()).digest() + vectors.append([((digest[index % len(digest)] / 255.0) * 2.0) - 1.0 for index in range(self.dimension)]) + return vectors + + +@pytest.mark.asyncio +async def test_oracle_retain_fresh_delta_and_recall( + oracle_db_url: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Exercise durable fresh, no-op, replacement, and retrieval paths.""" + + monkeypatch.setenv("HMS_API_DATABASE_BACKEND", "oracle") + clear_config_cache() + memory = MemoryEngine( + db_url=oracle_db_url, + memory_llm_provider="none", + memory_llm_model="none", + embeddings=_DeterministicEmbeddings(), + cross_encoder=RRFPassthroughCrossEncoder(), + query_analyzer=DateparserQueryAnalyzer(), + pool_min_size=1, + pool_max_size=3, + run_migrations=False, + task_backend=SyncTaskBackend(), + skip_llm_verification=True, + ) + await memory.initialize() + + bank_id = f"oracle-retain-live-{uuid.uuid4().hex[:12]}" + document_id = "profile" + request_context = RequestContext() + original = "Alice maintains the Atlas search service." + replacement = "Alice maintains the Atlas search service. Bob owns incident response." + + try: + first_ids = await memory.retain_async( + bank_id=bank_id, + content=original, + document_id=document_id, + request_context=request_context, + ) + assert first_ids + + unchanged_ids = await memory.retain_async( + bank_id=bank_id, + content=original, + document_id=document_id, + request_context=request_context, + ) + assert unchanged_ids == [] + + replacement_ids = await memory.retain_async( + bank_id=bank_id, + content=replacement, + document_id=document_id, + request_context=request_context, + ) + assert replacement_ids + + document = await memory.get_document( + document_id, + bank_id, + request_context=request_context, + ) + assert document is not None + assert document["original_text"] == replacement + + recalled = await memory.recall_async( + bank_id=bank_id, + query="Who owns incident response?", + budget=Budget.LOW, + max_tokens=512, + request_context=request_context, + ) + assert recalled.results + assert any("incident response" in result.text for result in recalled.results) + finally: + await memory.delete_bank(bank_id, request_context=request_context) + await memory.close() + clear_config_cache() + + +@pytest.mark.asyncio +async def test_oracle_async_batch_parent_lifecycle_uses_typed_uuid_predicates( + oracle_db_url: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Exercise parent lookup, cancellation, and retry against Oracle JSON.""" + + monkeypatch.setenv("HMS_API_DATABASE_BACKEND", "oracle") + clear_config_cache() + memory = MemoryEngine( + db_url=oracle_db_url, + memory_llm_provider="none", + memory_llm_model="none", + embeddings=_DeterministicEmbeddings(), + cross_encoder=RRFPassthroughCrossEncoder(), + query_analyzer=DateparserQueryAnalyzer(), + pool_min_size=1, + pool_max_size=3, + run_migrations=False, + task_backend=WorkerTaskBackend(), + skip_llm_verification=True, + ) + await memory.initialize() + + bank_id = f"oracle-parent-live-{uuid.uuid4().hex[:12]}" + request_context = RequestContext() + + try: + submitted = await memory.submit_async_retain( + bank_id=bank_id, + contents=[{"content": "Alice owns the Atlas service.", "document_id": "profile"}], + request_context=request_context, + ) + parent_operation_id = submitted["operation_id"] + + pending = await memory.get_operation_status( + bank_id=bank_id, + operation_id=parent_operation_id, + request_context=request_context, + ) + assert pending["status"] == "pending" + assert len(pending["child_operations"]) == 1 + assert pending["child_operations"][0]["status"] == "pending" + + await memory.cancel_operation( + bank_id=bank_id, + operation_id=parent_operation_id.upper(), + request_context=request_context, + ) + cancelled = await memory.get_operation_status( + bank_id=bank_id, + operation_id=parent_operation_id, + request_context=request_context, + ) + assert cancelled["status"] == "cancelled" + assert cancelled["child_operations"][0]["status"] == "cancelled" + + await memory.retry_operation( + bank_id=bank_id, + operation_id=parent_operation_id.upper(), + request_context=request_context, + ) + retried = await memory.get_operation_status( + bank_id=bank_id, + operation_id=parent_operation_id, + request_context=request_context, + ) + assert retried["status"] == "pending" + assert retried["child_operations"][0]["status"] == "pending" + finally: + try: + await memory.delete_bank(bank_id, request_context=request_context) + finally: + await memory.close() + clear_config_cache() diff --git a/core/dataplane/tests/test_ingestion_pipeline_contracts.py b/core/dataplane/tests/test_ingestion_pipeline_contracts.py new file mode 100644 index 0000000..80f06a5 --- /dev/null +++ b/core/dataplane/tests/test_ingestion_pipeline_contracts.py @@ -0,0 +1,1473 @@ +"""Focused offline contracts for the Retain ingestion application boundary.""" + +from __future__ import annotations + +import logging +import uuid +from contextlib import asynccontextmanager +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from hms_api.engine import embedding_fingerprint as fingerprint_module +from hms_api.engine import memory_engine as memory_engine_module +from hms_api.engine.embedding_fingerprint import ( + EmbeddingFingerprintMismatchError, +) +from hms_api.engine.embedding_fingerprint import ( + embedding_model_version as shared_embedding_model_version, +) +from hms_api.engine.ingestion import ( + RetainExecutionContext, + RetainInvocation, + RetainOperationInactiveError, + RetainPublicationAborted, +) +from hms_api.engine.ingestion import service as service_module +from hms_api.engine.ingestion.domain import DocumentChangeKind +from hms_api.engine.ingestion.persistence import writer as writer_module +from hms_api.engine.ingestion.persistence.operation_fence import OperationActivityFence +from hms_api.engine.ingestion.persistence.unit_of_work import ( + CoreGraphWrite, + FirstFullWriteWindow, + MetadataOnlyWriteRequest, + RetainUnitOfWork, + WriteWindowRequest, +) +from hms_api.engine.ingestion.persistence.writer import PersistenceWriter +from hms_api.engine.ingestion.redaction import IdentifierSanitizer +from hms_api.engine.ingestion.runtime import embedding_model_version +from hms_api.engine.memory_engine import MemoryEngine +from hms_api.engine.response_models import TokenUsage +from hms_api.engine.retain import fact_storage +from hms_api.engine.retain.types import RetainContent + + +def _embedding_model() -> SimpleNamespace: + return SimpleNamespace( + provider_name="local", + model="BAAI/bge-small-en-v1.5", + dimension=384, + normalization=True, + ) + + +def _config(**overrides) -> SimpleNamespace: + values = { + "database_backend": "postgresql", + "retain_extraction_mode": "chunks", + "retain_chunk_size": 3000, + "embedding_fingerprint_policy": "strict", + "embedding_fingerprint_legacy_attestation": None, + } + values.update(overrides) + return SimpleNamespace(**values) + + +def _execution(*, config=None) -> RetainExecutionContext: + return RetainExecutionContext( + pool=SimpleNamespace(backend_type="postgresql", ops=object()), + embeddings_model=_embedding_model(), + llm_config=object(), + entity_resolver=SimpleNamespace(discard_pending_stats=lambda: None), + format_date_fn=lambda *_args, **_kwargs: "", + resolved_config=config or _config(), + ) + + +class _Ownership: + def __init__(self, events: list[str], *, owns: bool = True) -> None: + self._events = events + self._owns = owns + + async def prepare_first_window(self, _connection, *, bank_id, document_id) -> None: + del bank_id, document_id + self._events.append("ownership") + + async def validate_later_window( + self, + _connection, + *, + bank_id, + document_id, + expected_content_hash, + ) -> bool: + del bank_id, document_id, expected_content_hash + self._events.append("ownership") + return self._owns + + async def validate_unhashed_window( + self, + _connection, + *, + bank_id, + document_id, + ) -> bool: + del bank_id, document_id + self._events.append("unhashed-ownership") + return self._owns + + async def transition_content_hash( + self, + _connection, + *, + bank_id, + document_id, + expected_content_hash, + new_content_hash, + ) -> bool: + del bank_id, document_id, expected_content_hash, new_content_hash + self._events.append("transition") + return self._owns + + +class _Connection: + def __init__(self, events: list[str]) -> None: + self._events = events + self.in_transaction = False + + def transaction(self): + connection = self + + class _Transaction: + async def __aenter__(self): + connection.in_transaction = True + connection._events.append("begin") + return connection + + async def __aexit__(self, exc_type, _exc, _traceback): + connection._events.append("rollback" if exc_type is not None else "commit") + connection.in_transaction = False + return False + + return _Transaction() + + +def _writer( + events: list[str], + *, + sanitize: bool = False, + owns: bool = True, + operation_activity=None, +) -> PersistenceWriter: + return PersistenceWriter( + pool=SimpleNamespace(ops=object()), + embeddings_model=_embedding_model(), + entity_resolver=object(), + config=_config(), + ownership=_Ownership(events, owns=owns), + operation_activity=operation_activity, + sanitize_log_identifiers=sanitize, + ) + + +def _unit_of_work(writer: PersistenceWriter, connection: _Connection) -> RetainUnitOfWork: + @asynccontextmanager + async def connection_scope(): + yield connection + + return RetainUnitOfWork(connection_scope=connection_scope, adapter=writer) + + +def test_runtime_uses_shared_embedding_compatibility_version() -> None: + model = _embedding_model() + + version = embedding_model_version(model) + + assert version == shared_embedding_model_version(model) + assert version.startswith("fp:") + assert len(version) == len("fp:") + 64 + + +@pytest.mark.asyncio +async def test_tracked_anonymous_batch_is_one_retry_stable_document(monkeypatch) -> None: + engine = object.__new__(MemoryEngine) + engine._operation_validator = None + engine._authenticate_tenant = AsyncMock() + engine._check_op_alive = AsyncMock(return_value=True) + engine._replace_vector_index_document = AsyncMock() + engine._replace_vector_index_fact_type = AsyncMock() + engine._sync_vector_index_units = AsyncMock() + engine._config_resolver = SimpleNamespace( + resolve_full_config=AsyncMock(return_value=SimpleNamespace(enable_observations=False)) + ) + monkeypatch.setattr( + memory_engine_module, + "get_config", + lambda: SimpleNamespace(retain_batch_tokens=1), + ) + + committed_units_by_document: dict[str, list[str]] = {} + internal_calls: list[tuple[tuple[str, ...], bool, bool]] = [] + outbox_calls: list[object] = [] + + async def outbox_callback(connection) -> None: + outbox_calls.append(connection) + + async def retain_internal(**kwargs): + document_ids = tuple(item["document_id"] for item in kwargs["contents"]) + callback = kwargs["outbox_callback"] + internal_calls.append( + ( + document_ids, + kwargs["is_first_batch"], + callback is not None, + ) + ) + assert len(set(document_ids)) == 1 + document_id = document_ids[0] + is_recovery = document_id in committed_units_by_document + if not is_recovery: + committed_units_by_document[document_id] = [f"unit-{index + 1}" for index in range(len(kwargs["contents"]))] + if callback is not None and not is_recovery: + await callback(object()) + return [[unit_id] for unit_id in committed_units_by_document[document_id]], TokenUsage(), 1 + + engine._retain_batch_async_internal = retain_internal + operation_id = str(uuid.uuid4()) + + request = { + "bank_id": "bank", + "contents": [ + {"content": "first anonymous item"}, + {"content": "second anonymous item"}, + {"content": "third anonymous item"}, + ], + "request_context": object(), + "operation_id": operation_id, + "outbox_callback": outbox_callback, + } + result = await engine.retain_batch_async(**request) + retry_result = await engine.retain_batch_async(**request) + + assert result == [["unit-1"], ["unit-2"], ["unit-3"]] + assert retry_result == result + assert len(committed_units_by_document) == 1 + assert len(internal_calls) == 2 + assert internal_calls[0][0] == internal_calls[1][0] + assert internal_calls[0][1:] == (True, True) + assert internal_calls[1][1:] == (True, True) + assert len(outbox_calls) == 1 + + +@pytest.mark.asyncio +async def test_tracked_token_retry_survives_threshold_change_without_checkpoint_alias(monkeypatch) -> None: + engine = object.__new__(MemoryEngine) + engine._operation_validator = None + engine._authenticate_tenant = AsyncMock() + engine._check_op_alive = AsyncMock(return_value=True) + engine._replace_vector_index_document = AsyncMock() + engine._replace_vector_index_fact_type = AsyncMock() + engine._sync_vector_index_units = AsyncMock() + engine._config_resolver = SimpleNamespace( + resolve_full_config=AsyncMock(return_value=SimpleNamespace(enable_observations=False)) + ) + config = SimpleNamespace(retain_batch_tokens=1) + monkeypatch.setattr(memory_engine_module, "get_config", lambda: config) + + operation_id = str(uuid.uuid4()) + committed_units_by_document: dict[str, str] = {} + anonymous_ids: list[str] = [] + outbox_calls: list[object] = [] + fail_second_call = True + internal_call_count = 0 + + async def outbox_callback(connection) -> None: + outbox_calls.append(connection) + + async def retain_internal(**kwargs): + nonlocal fail_second_call, internal_call_count + internal_call_count += 1 + for item in kwargs["contents"]: + if item["content"] == "anonymous payload": + anonymous_ids.append(item["document_id"]) + if fail_second_call and internal_call_count == 2: + fail_second_call = False + raise RuntimeError("synthetic crash after first document commit") + + results: list[list[str]] = [] + committed_new_document = False + for item in kwargs["contents"]: + document_id = item["document_id"] + if document_id not in committed_units_by_document: + committed_units_by_document[document_id] = f"unit-{len(committed_units_by_document) + 1}" + committed_new_document = True + results.append([committed_units_by_document[document_id]]) + if kwargs["outbox_callback"] is not None and committed_new_document: + await kwargs["outbox_callback"](object()) + return results, TokenUsage(), 1 + + engine._retain_batch_async_internal = retain_internal + request = { + "bank_id": "bank", + "contents": [ + {"content": "first explicit payload", "document_id": "document-a"}, + {"content": "anonymous payload"}, + {"content": "second explicit payload", "document_id": "document-b"}, + ], + "request_context": object(), + "operation_id": operation_id, + "outbox_callback": outbox_callback, + } + + with pytest.raises(RuntimeError, match="synthetic crash"): + await engine.retain_batch_async(**request) + + config.retain_batch_tokens = 10_000 + result = await engine.retain_batch_async(**request) + + assert len(result) == 3 + assert len(committed_units_by_document) == 3 + assert len(anonymous_ids) == 2 + assert len(set(anonymous_ids)) == 1 + assert len(outbox_calls) == 1 + + +@pytest.mark.asyncio +async def test_token_batching_keeps_repeated_document_group_and_restores_input_order(monkeypatch) -> None: + engine = object.__new__(MemoryEngine) + engine._operation_validator = None + engine._authenticate_tenant = AsyncMock() + engine._check_op_alive = AsyncMock(return_value=True) + engine._replace_vector_index_document = AsyncMock() + engine._replace_vector_index_fact_type = AsyncMock() + engine._sync_vector_index_units = AsyncMock() + engine._config_resolver = SimpleNamespace( + resolve_full_config=AsyncMock(return_value=SimpleNamespace(enable_observations=False)) + ) + monkeypatch.setattr( + memory_engine_module, + "get_config", + lambda: SimpleNamespace(retain_batch_tokens=1), + ) + submitted_groups: list[tuple[str, ...]] = [] + + async def retain_internal(**kwargs): + submitted_groups.append(tuple(item["slot"] for item in kwargs["contents"])) + return [[item["slot"]] for item in kwargs["contents"]], TokenUsage(), 1 + + engine._retain_batch_async_internal = retain_internal + result = await engine.retain_batch_async( + bank_id="bank", + contents=[ + {"content": "alpha one", "document_id": "document-a", "slot": "slot-0"}, + {"content": "beta", "document_id": "document-b", "slot": "slot-1"}, + {"content": "alpha two", "document_id": "document-a", "slot": "slot-2"}, + ], + request_context=object(), + operation_id=str(uuid.uuid4()), + ) + + assert submitted_groups == [("slot-0", "slot-2"), ("slot-1",)] + assert result == [["slot-0"], ["slot-1"], ["slot-2"]] + + +@pytest.mark.asyncio +async def test_tracked_token_cancellation_raises_without_partial_success_or_outbox(monkeypatch) -> None: + engine = object.__new__(MemoryEngine) + engine._operation_validator = None + engine._authenticate_tenant = AsyncMock() + engine._check_op_alive = AsyncMock(side_effect=[True, False]) + engine._replace_vector_index_document = AsyncMock() + engine._replace_vector_index_fact_type = AsyncMock() + engine._sync_vector_index_units = AsyncMock() + engine._config_resolver = SimpleNamespace( + resolve_full_config=AsyncMock(return_value=SimpleNamespace(enable_observations=False)) + ) + monkeypatch.setattr( + memory_engine_module, + "get_config", + lambda: SimpleNamespace(retain_batch_tokens=1), + ) + internal = AsyncMock(return_value=([["unit-a"]], TokenUsage(), 1)) + engine._retain_batch_async_internal = internal + outbox = AsyncMock() + + with pytest.raises( + memory_engine_module._RetainOperationCancelled, + match="cancelled between logical document batches", + ): + await engine.retain_batch_async( + bank_id="bank", + contents=[ + {"content": "first payload", "document_id": "document-a"}, + {"content": "second payload", "document_id": "document-b"}, + ], + request_context=object(), + operation_id=str(uuid.uuid4()), + outbox_callback=outbox, + ) + + assert internal.await_count == 1 + outbox.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_mark_completed_uses_terminal_safe_compare_and_set(monkeypatch) -> None: + engine = object.__new__(MemoryEngine) + engine._get_backend = AsyncMock(return_value=object()) + queries: list[str] = [] + + class CompletionConnection: + @asynccontextmanager + async def transaction(self): + yield self + + async def fetchrow(self, query, *_args): + queries.append(query) + return None + + @asynccontextmanager + async def connection_scope(*_args, **_kwargs): + yield CompletionConnection() + + monkeypatch.setattr(memory_engine_module, "acquire_with_retry", connection_scope) + + await engine._mark_operation_completed(str(uuid.uuid4())) + + assert len(queries) == 1 + assert "status IN ('pending', 'processing')" in queries[0] + + +@pytest.mark.asyncio +async def test_execute_task_treats_retain_cancellation_as_terminal() -> None: + engine = object.__new__(MemoryEngine) + engine._audit_logger = None + engine._handle_batch_retain = AsyncMock( + side_effect=memory_engine_module._RetainOperationCancelled("cancelled"), + ) + engine._mark_operation_completed = AsyncMock() + + await engine.execute_task( + { + "type": "batch_retain", + "bank_id": "bank", + "contents": [], + } + ) + + engine._mark_operation_completed.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_child_liveness_fails_closed_when_batch_parent_is_cancelled(monkeypatch) -> None: + engine = object.__new__(MemoryEngine) + engine._get_backend = AsyncMock(return_value=object()) + child_id = uuid.uuid4() + parent_id = uuid.uuid4() + queries: list[str] = [] + + class LivenessConnection: + def parse_json(self, value): + return value + + async def fetchrow(self, query, *_args): + queries.append(query) + if len(queries) == 1: + return { + "status": "processing", + "bank_id": "bank", + "result_metadata": {"parent_operation_id": str(parent_id)}, + } + return {"status": "cancelled"} + + @asynccontextmanager + async def connection_scope(*_args, **_kwargs): + yield LivenessConnection() + + monkeypatch.setattr(memory_engine_module, "acquire_with_retry", connection_scope) + + assert await engine._check_op_alive(str(child_id)) is False + assert len(queries) == 2 + + +@pytest.mark.asyncio +async def test_parent_aggregation_never_overwrites_terminal_cancellation() -> None: + engine = object.__new__(MemoryEngine) + parent_id = uuid.uuid4() + fetch_siblings = AsyncMock(side_effect=AssertionError("terminal parent must stop aggregation")) + execute = AsyncMock(side_effect=AssertionError("terminal parent must not be updated")) + + class AggregationConnection: + def __init__(self) -> None: + self.fetchrow_calls = 0 + self.fetch = fetch_siblings + self.execute = execute + + def parse_json(self, value): + return value + + async def fetchrow(self, _query, *_args): + self.fetchrow_calls += 1 + if self.fetchrow_calls == 1: + return { + "bank_id": "bank", + "result_metadata": {"parent_operation_id": str(parent_id)}, + } + return {"operation_id": parent_id, "status": "cancelled"} + + await engine._maybe_update_parent_operation(str(uuid.uuid4()), AggregationConnection()) + + fetch_siblings.assert_not_awaited() + execute.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_cancelled_child_is_terminal_and_cancels_active_parent() -> None: + engine = object.__new__(MemoryEngine) + parent_id = uuid.uuid4() + updates: list[tuple[str, tuple[object, ...]]] = [] + + class AggregationConnection: + def __init__(self) -> None: + self.fetchrow_calls = 0 + + def parse_json(self, value): + return value + + async def fetchrow(self, _query, *_args): + self.fetchrow_calls += 1 + if self.fetchrow_calls == 1: + return { + "bank_id": "bank", + "result_metadata": {"parent_operation_id": str(parent_id)}, + } + return {"operation_id": parent_id, "status": "pending"} + + async def fetch(self, _query, *_args): + return [ + {"status": "completed", "error_message": None}, + {"status": "cancelled", "error_message": None}, + ] + + async def execute(self, query, *args): + updates.append((query, args)) + + await engine._maybe_update_parent_operation(str(uuid.uuid4()), AggregationConnection()) + + assert len(updates) == 1 + query, args = updates[0] + assert args[1] == "cancelled" + assert "status IN ('pending', 'processing')" in query + + +@pytest.mark.asyncio +async def test_poller_parent_aggregation_treats_cancelled_as_terminal() -> None: + from hms_api.worker.poller import WorkerPoller + + poller = object.__new__(WorkerPoller) + parent_id = uuid.uuid4() + updates: list[tuple[str, tuple[object, ...]]] = [] + + class PollerAggregationConnection: + def __init__(self) -> None: + self.fetchrow_calls = 0 + + async def fetchrow(self, _query, *_args): + self.fetchrow_calls += 1 + if self.fetchrow_calls == 1: + return { + "bank_id": "bank", + "result_metadata": {"parent_operation_id": str(parent_id)}, + } + return {"operation_id": parent_id, "status": "pending"} + + async def fetch(self, _query, *_args): + return [ + {"status": "completed", "error_message": None}, + {"status": "cancelled", "error_message": None}, + ] + + async def execute(self, query, *args): + updates.append((query, args)) + + await poller._maybe_update_parent_operation( + str(uuid.uuid4()), + None, + PollerAggregationConnection(), + ) + + assert len(updates) == 1 + query, _args = updates[0] + assert "SET status = 'cancelled'" in query + assert "status IN ('pending', 'processing')" in query + + +@pytest.mark.asyncio +async def test_retry_parent_reopens_retryable_children_before_parent(monkeypatch) -> None: + engine = object.__new__(MemoryEngine) + engine._operation_validator = None + engine._authenticate_tenant = AsyncMock() + engine._get_backend = AsyncMock(return_value=object()) + parent_id = uuid.uuid4() + events: list[tuple[str, str, tuple[object, ...]]] = [] + + class RetryConnection: + def __init__(self) -> None: + self.fetchrow_calls = 0 + + def parse_json(self, value): + return value + + @asynccontextmanager + async def transaction(self): + yield self + + async def fetchrow(self, query, *args): + self.fetchrow_calls += 1 + events.append(("fetchrow", query, args)) + return { + "bank_id": "bank", + "status": "failed", + "operation_type": "batch_retain", + "result_metadata": {"is_parent": True}, + } + + async def fetch(self, query, *args): + events.append(("fetch", query, args)) + return [ + {"operation_id": uuid.uuid4(), "status": "completed"}, + {"operation_id": uuid.uuid4(), "status": "failed"}, + {"operation_id": uuid.uuid4(), "status": "cancelled"}, + ] + + async def execute(self, query, *args): + events.append(("execute", query, args)) + + @asynccontextmanager + async def connection_scope(*_args, **_kwargs): + yield RetryConnection() + + monkeypatch.setattr(memory_engine_module, "acquire_with_retry", connection_scope) + + await engine.retry_operation( + bank_id="bank", + operation_id=str(parent_id).upper(), + request_context=object(), + ) + + child_lock_index = next(index for index, event in enumerate(events) if event[0] == "fetch") + parent_lock_index = next( + index for index, event in enumerate(events) if event[0] == "fetchrow" and "FOR UPDATE" in event[1] + ) + assert child_lock_index < parent_lock_index + updates = [event for event in events if event[0] == "execute"] + assert len(updates) == 2 + assert "status IN ('failed', 'cancelled')" in updates[0][1] + assert "result_metadata->>'parent_operation_id'" in updates[0][1] + assert updates[0][2][1] == parent_id + assert "status IN ('failed', 'cancelled')" in updates[1][1] + + +@pytest.mark.asyncio +async def test_retry_child_reopens_only_failed_or_cancelled_parent(monkeypatch) -> None: + engine = object.__new__(MemoryEngine) + engine._operation_validator = None + engine._authenticate_tenant = AsyncMock() + engine._get_backend = AsyncMock(return_value=object()) + child_id = uuid.uuid4() + parent_id = uuid.uuid4() + events: list[tuple[str, str, tuple[object, ...]]] = [] + + class RetryConnection: + def __init__(self) -> None: + self.fetchrow_calls = 0 + + def parse_json(self, value): + return value + + @asynccontextmanager + async def transaction(self): + yield self + + async def fetchrow(self, query, *args): + self.fetchrow_calls += 1 + events.append(("fetchrow", query, args)) + if self.fetchrow_calls <= 2: + return { + "bank_id": "bank", + "status": "cancelled", + "operation_type": "retain", + "result_metadata": {"parent_operation_id": str(parent_id)}, + } + return {"status": "failed"} + + async def execute(self, query, *args): + events.append(("execute", query, args)) + + @asynccontextmanager + async def connection_scope(*_args, **_kwargs): + yield RetryConnection() + + monkeypatch.setattr(memory_engine_module, "acquire_with_retry", connection_scope) + + await engine.retry_operation( + bank_id="bank", + operation_id=str(child_id), + request_context=object(), + ) + + updates = [event for event in events if event[0] == "execute"] + assert len(updates) == 2 + assert str(updates[0][2][0]) == str(parent_id) + assert "status IN ('failed', 'cancelled')" in updates[0][1] + assert str(updates[1][2][0]) == str(child_id) + assert "status IN ('failed', 'cancelled')" in updates[1][1] + + +@pytest.mark.asyncio +async def test_service_preflights_fingerprint_before_extraction(monkeypatch) -> None: + events: list[str] = [] + invocation = RetainInvocation( + bank_id="bank", + raw_contents=({"content": "payload", "document_id": "doc"},), + request_context=object(), + ) + execution = _execution() + pipeline = service_module.RetainPipelineService() + plan = SimpleNamespace( + change=SimpleNamespace(kind=DocumentChangeKind.FULL), + recovered_unit_ids=None, + intent=SimpleNamespace(items=()), + ) + + monkeypatch.setattr( + service_module, + "normalize_contents", + lambda *_args, **_kwargs: [SimpleNamespace(document_id="doc")], + ) + monkeypatch.setattr(service_module, "plan_documents", lambda *_args, **_kwargs: (object(),)) + monkeypatch.setattr( + pipeline, + "_recover_checkpoint", + AsyncMock(return_value=SimpleNamespace(document_ids=(), core_committed_document_ids=())), + ) + monkeypatch.setattr(pipeline, "_preflight_documents", AsyncMock(return_value=(plan,))) + + async def get_bank_profile(*_args, **_kwargs): + events.append("bank") + return {"name": "agent"} + + async def ensure_fingerprint(*_args, **kwargs): + assert kwargs.get("for_write", False) is False + events.append("fingerprint") + + @asynccontextmanager + async def connection_scope(*_args, **_kwargs): + yield object() + + async def record_document_ids(*_args, **_kwargs): + events.append("checkpoint") + + async def execute_document(*_args, **_kwargs): + events.append("extract-and-write") + return service_module._DocumentOutcome((), TokenUsage(), 0) + + monkeypatch.setattr(service_module.bank_utils, "get_bank_profile", get_bank_profile) + monkeypatch.setattr(service_module, "ensure_bank_embedding_fingerprint", ensure_fingerprint) + monkeypatch.setattr(service_module, "acquire_with_retry", connection_scope) + monkeypatch.setattr(pipeline, "_record_document_ids", record_document_ids) + monkeypatch.setattr(pipeline, "_execute_document", execute_document) + monkeypatch.setattr(pipeline, "_merge_document_result", lambda *_args, **_kwargs: None) + + await pipeline._retain_in_schema(invocation, execution) + + assert events == ["bank", "fingerprint", "checkpoint", "extract-and-write"] + + +@pytest.mark.asyncio +async def test_preflight_mismatch_stops_before_write_and_redacts_bank(monkeypatch) -> None: + bank_id = "sensitive-bank" + callback = AsyncMock() + invocation = RetainInvocation( + bank_id=bank_id, + raw_contents=({"content": "payload", "document_id": "doc"},), + request_context=object(), + outbox_callback=callback, + sanitize_log_identifiers=True, + ) + pipeline = service_module.RetainPipelineService() + plan = SimpleNamespace( + change=SimpleNamespace(kind=DocumentChangeKind.FULL), + recovered_unit_ids=None, + ) + record_document_ids = AsyncMock() + execute_document = AsyncMock() + + monkeypatch.setattr( + service_module, + "normalize_contents", + lambda *_args, **_kwargs: [SimpleNamespace(document_id="doc")], + ) + monkeypatch.setattr(service_module, "plan_documents", lambda *_args, **_kwargs: (object(),)) + monkeypatch.setattr( + pipeline, + "_recover_checkpoint", + AsyncMock(return_value=SimpleNamespace(document_ids=(), core_committed_document_ids=())), + ) + monkeypatch.setattr(pipeline, "_preflight_documents", AsyncMock(return_value=(plan,))) + monkeypatch.setattr( + service_module.bank_utils, + "get_bank_profile", + AsyncMock(return_value={"name": "agent"}), + ) + + @asynccontextmanager + async def connection_scope(*_args, **_kwargs): + yield object() + + async def reject_fingerprint(*_args, **kwargs): + assert kwargs["for_write"] is False + raise EmbeddingFingerprintMismatchError(f"Fingerprint mismatch for {bank_id}") + + monkeypatch.setattr(service_module, "acquire_with_retry", connection_scope) + monkeypatch.setattr(service_module, "ensure_bank_embedding_fingerprint", reject_fingerprint) + monkeypatch.setattr(pipeline, "_record_document_ids", record_document_ids) + monkeypatch.setattr(pipeline, "_execute_document", execute_document) + + with pytest.raises(EmbeddingFingerprintMismatchError) as error: + await pipeline._retain_in_schema(invocation, _execution()) + + assert bank_id not in str(error.value) + assert "" in str(error.value) + record_document_ids.assert_not_awaited() + execute_document.assert_not_awaited() + callback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_all_stale_publication_fence_aborts_without_callback(monkeypatch) -> None: + bank_id = "sensitive-bank" + document_id = "sensitive-document" + operation_id = "sensitive-operation" + callback = AsyncMock() + invocation = RetainInvocation( + bank_id=bank_id, + raw_contents=({"content": "payload", "document_id": document_id},), + request_context=object(), + operation_id=operation_id, + outbox_callback=callback, + sanitize_log_identifiers=True, + ) + pipeline = service_module.RetainPipelineService() + stale_plan = SimpleNamespace( + change=SimpleNamespace(kind=DocumentChangeKind.STALE_SKIP), + recovered_unit_ids=None, + ) + + monkeypatch.setattr( + service_module, + "normalize_contents", + lambda *_args, **_kwargs: [SimpleNamespace(document_id=document_id)], + ) + monkeypatch.setattr(service_module, "plan_documents", lambda *_args, **_kwargs: (object(),)) + monkeypatch.setattr( + pipeline, + "_recover_checkpoint", + AsyncMock(return_value=SimpleNamespace(document_ids=(), core_committed_document_ids=())), + ) + monkeypatch.setattr(pipeline, "_preflight_documents", AsyncMock(return_value=(stale_plan,))) + + with pytest.raises(RetainPublicationAborted, match="superseded before publication") as error: + await pipeline._retain_in_schema(invocation, _execution()) + + for identifier in (bank_id, document_id, operation_id): + assert identifier not in str(error.value) + callback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_mixed_stale_publication_fence_aborts_before_partial_publish(monkeypatch) -> None: + callback = AsyncMock() + invocation = RetainInvocation( + bank_id="bank", + raw_contents=( + {"content": "current payload", "document_id": "current-doc"}, + {"content": "superseded payload", "document_id": "stale-doc"}, + ), + request_context=object(), + outbox_callback=callback, + ) + pipeline = service_module.RetainPipelineService() + current_plan = SimpleNamespace( + change=SimpleNamespace(kind=DocumentChangeKind.FULL), + recovered_unit_ids=None, + ) + stale_plan = SimpleNamespace( + change=SimpleNamespace(kind=DocumentChangeKind.STALE_SKIP), + recovered_unit_ids=None, + ) + execute_document = AsyncMock() + + monkeypatch.setattr( + service_module, + "normalize_contents", + lambda *_args, **_kwargs: [ + SimpleNamespace(document_id="current-doc"), + SimpleNamespace(document_id="stale-doc"), + ], + ) + monkeypatch.setattr(service_module, "plan_documents", lambda *_args, **_kwargs: (object(), object())) + monkeypatch.setattr( + pipeline, + "_recover_checkpoint", + AsyncMock(return_value=SimpleNamespace(document_ids=(), core_committed_document_ids=())), + ) + monkeypatch.setattr( + pipeline, + "_preflight_documents", + AsyncMock(return_value=(current_plan, stale_plan)), + ) + monkeypatch.setattr(pipeline, "_execute_document", execute_document) + + with pytest.raises(RetainPublicationAborted, match="superseded before publication"): + await pipeline._retain_in_schema(invocation, _execution()) + + execute_document.assert_not_awaited() + callback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_all_stale_without_publication_fence_remains_a_noop(monkeypatch) -> None: + invocation = RetainInvocation( + bank_id="bank", + raw_contents=({"content": "payload", "document_id": "doc"},), + request_context=object(), + ) + pipeline = service_module.RetainPipelineService() + stale_plan = SimpleNamespace( + change=SimpleNamespace(kind=DocumentChangeKind.STALE_SKIP), + recovered_unit_ids=None, + intent=SimpleNamespace(items=(object(),)), + ) + execute_document = AsyncMock() + + monkeypatch.setattr( + service_module, + "normalize_contents", + lambda *_args, **_kwargs: [SimpleNamespace(document_id="doc")], + ) + monkeypatch.setattr(service_module, "plan_documents", lambda *_args, **_kwargs: (object(),)) + monkeypatch.setattr( + pipeline, + "_recover_checkpoint", + AsyncMock(return_value=SimpleNamespace(document_ids=(), core_committed_document_ids=())), + ) + monkeypatch.setattr(pipeline, "_preflight_documents", AsyncMock(return_value=(stale_plan,))) + monkeypatch.setattr( + service_module.bank_utils, + "get_bank_profile", + AsyncMock(return_value={"name": "agent"}), + ) + + @asynccontextmanager + async def connection_scope(*_args, **_kwargs): + yield object() + + monkeypatch.setattr(service_module, "acquire_with_retry", connection_scope) + monkeypatch.setattr(service_module, "ensure_bank_embedding_fingerprint", AsyncMock()) + monkeypatch.setattr(pipeline, "_record_document_ids", AsyncMock()) + monkeypatch.setattr(pipeline, "_execute_document", execute_document) + monkeypatch.setattr(pipeline, "_merge_document_result", lambda *_args, **_kwargs: None) + + outcome = await pipeline._retain_in_schema(invocation, _execution()) + + assert outcome.unit_ids_by_input == [[]] + execute_document.assert_not_awaited() + + +class _OperationActivity: + def __init__(self, events: list[str], *, active: bool = True) -> None: + self._events = events + self._active = active + + async def assert_active(self, connection, *, bank_id: str) -> None: + assert connection.in_transaction + assert bank_id == "bank" + self._events.append("operation") + if not self._active: + raise RetainOperationInactiveError("inactive") + + +class _FenceConnection: + def __init__(self, rows: list[dict]) -> None: + self.rows = list(rows) + self.queries: list[str] = [] + self.args: list[tuple[object, ...]] = [] + + async def fetchrow(self, query: str, *args): + self.queries.append(query) + self.args.append(args) + return self.rows.pop(0) + + @staticmethod + def parse_json(value): + return value + + +@pytest.mark.asyncio +async def test_operation_activity_fence_locks_active_child_then_parent() -> None: + child_id = uuid.uuid4() + parent_id = uuid.uuid4() + connection = _FenceConnection( + [ + { + "status": "processing", + "result_metadata": {"parent_operation_id": str(parent_id)}, + }, + {"status": "pending"}, + ] + ) + + await OperationActivityFence(str(child_id), schema="tenant").assert_active( + connection, + bank_id="bank", + ) + + assert len(connection.queries) == 2 + assert all("FOR UPDATE" in query for query in connection.queries) + assert '"tenant".async_operations' in connection.queries[0] + assert connection.args == [ + (child_id, "bank"), + (parent_id, "bank"), + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "rows,match", + [ + ([{"status": "cancelled", "result_metadata": {}}], "operation is no longer active"), + ( + [ + { + "status": "processing", + "result_metadata": {"parent_operation_id": str(uuid.uuid4())}, + }, + {"status": "cancelled"}, + ], + "parent operation is no longer active", + ), + ], +) +async def test_operation_activity_fence_rejects_terminal_child_or_parent(rows, match) -> None: + connection = _FenceConnection(rows) + + with pytest.raises(RetainOperationInactiveError, match=match): + await OperationActivityFence(str(uuid.uuid4())).assert_active( + connection, + bank_id="bank", + ) + + assert all("FOR UPDATE" in query for query in connection.queries) + + +@pytest.mark.asyncio +async def test_inactive_operation_rolls_back_before_fingerprint_or_any_core_write(monkeypatch) -> None: + events: list[str] = [] + writer = _writer( + events, + operation_activity=_OperationActivity(events, active=False), + ) + connection = _Connection(events) + fingerprint = AsyncMock(side_effect=AssertionError("fingerprint must not run")) + callback = AsyncMock() + monkeypatch.setattr(writer_module, "ensure_bank_embedding_fingerprint", fingerprint) + + with pytest.raises(RetainOperationInactiveError, match="inactive"): + await _unit_of_work(writer, connection).execute( + MetadataOnlyWriteRequest( + bank_id="bank", + document_id="doc", + expected_content_hash="hash", + combined_content="payload", + input_slot_count=1, + outbox_callback=callback, + ), + ) + + assert events == ["begin", "operation", "rollback"] + fingerprint.assert_not_awaited() + callback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_unhashed_existing_document_uses_locked_full_replacement(monkeypatch) -> None: + events: list[str] = [] + writer = _writer(events) + connection = _Connection(events) + + async def ensure_fingerprint(*_args, **_kwargs): + events.append("fingerprint") + + async def track_document(*_args, **_kwargs): + events.append("document") + + async def insert_facts(*_args, **_kwargs): + events.append("facts") + return [[]], object() + + monkeypatch.setattr(writer_module, "ensure_bank_embedding_fingerprint", ensure_fingerprint) + monkeypatch.setattr(writer_module.fact_storage, "handle_document_tracking", track_document) + monkeypatch.setattr(writer_module.runtime, "insert_facts_and_links", insert_facts) + monkeypatch.setattr(writer, "_finalize_entity_graph", AsyncMock(return_value=CoreGraphWrite())) + + result = await _unit_of_work(writer, connection).execute( + WriteWindowRequest( + bank_id="bank", + document_id="doc", + document_window=FirstFullWriteWindow( + combined_content="payload", + expects_unhashed_existing_document=True, + ), + contents=(RetainContent(content="payload"),), + ), + ) + + assert result.core.ownership.value == "owned" + assert events == [ + "begin", + "fingerprint", + "unhashed-ownership", + "document", + "facts", + "commit", + ] + + +@pytest.mark.asyncio +async def test_unhashed_existing_document_loses_ownership_before_replacement(monkeypatch) -> None: + events: list[str] = [] + writer = _writer(events, owns=False) + connection = _Connection(events) + track_document = AsyncMock(side_effect=AssertionError("replacement must not run")) + callback = AsyncMock() + + async def ensure_fingerprint(*_args, **_kwargs): + events.append("fingerprint") + + monkeypatch.setattr(writer_module, "ensure_bank_embedding_fingerprint", ensure_fingerprint) + monkeypatch.setattr(writer_module.fact_storage, "handle_document_tracking", track_document) + + result = await _unit_of_work(writer, connection).execute( + WriteWindowRequest( + bank_id="bank", + document_id="doc", + document_window=FirstFullWriteWindow( + combined_content="payload", + expects_unhashed_existing_document=True, + ), + contents=(RetainContent(content="payload"),), + outbox_callback=callback, + ), + ) + + assert result.core.ownership.value == "lost" + assert events == ["begin", "fingerprint", "unhashed-ownership", "commit"] + track_document.assert_not_awaited() + callback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_transaction_guard_covers_zero_fact_full_write(monkeypatch) -> None: + events: list[str] = [] + writer = _writer(events) + connection = _Connection(events) + + async def ensure_fingerprint(*args, **kwargs): + assert args[0] is connection + assert args[1] == "bank" + assert connection.in_transaction + assert kwargs["for_write"] is True + events.append("fingerprint") + + async def track_document(*_args, **_kwargs): + events.append("document") + + async def insert_facts(*_args, **kwargs): + events.append("facts") + await kwargs["outbox_callback"](connection) + return [[]], object() + + async def callback(callback_connection): + assert callback_connection is connection + events.append("outbox") + + monkeypatch.setattr(writer_module, "ensure_bank_embedding_fingerprint", ensure_fingerprint) + monkeypatch.setattr(writer_module.fact_storage, "handle_document_tracking", track_document) + monkeypatch.setattr(writer_module.runtime, "insert_facts_and_links", insert_facts) + monkeypatch.setattr(writer, "_finalize_entity_graph", AsyncMock(return_value=CoreGraphWrite())) + + result = await _unit_of_work(writer, connection).execute( + WriteWindowRequest( + bank_id="bank", + document_id="doc", + document_window=FirstFullWriteWindow(combined_content="payload"), + contents=(RetainContent(content="payload"),), + outbox_callback=callback, + ), + ) + + assert result.core.unit_ids_by_content == ((),) + assert result.core.post_commit_required is False + assert events == ["begin", "fingerprint", "ownership", "document", "facts", "outbox", "commit"] + + +@pytest.mark.asyncio +async def test_transaction_guard_covers_metadata_only_write(monkeypatch) -> None: + events: list[str] = [] + writer = _writer(events) + connection = _Connection(events) + + async def ensure_fingerprint(*_args, **kwargs): + assert connection.in_transaction + assert kwargs["for_write"] is True + events.append("fingerprint") + + async def update_metadata(*_args, **_kwargs): + events.append("metadata") + + async def update_tags(*_args, **_kwargs): + events.append("tags") + + async def callback(_connection): + events.append("outbox") + + monkeypatch.setattr(writer_module, "ensure_bank_embedding_fingerprint", ensure_fingerprint) + monkeypatch.setattr(writer_module.fact_storage, "upsert_document_metadata", update_metadata) + monkeypatch.setattr(writer_module.fact_storage, "update_memory_units_tags", update_tags) + + result = await _unit_of_work(writer, connection).execute( + MetadataOnlyWriteRequest( + bank_id="bank", + document_id="doc", + expected_content_hash="hash", + combined_content="payload", + input_slot_count=1, + outbox_callback=callback, + ), + ) + + assert result.core.unit_ids_by_content == ((),) + assert events == ["begin", "fingerprint", "ownership", "metadata", "tags", "outbox", "commit"] + + +@pytest.mark.asyncio +async def test_fingerprint_failure_prevents_mutation_and_redacts_identifier(monkeypatch) -> None: + bank_id = "sensitive-bank" + document_id = "sensitive-document" + events: list[str] = [] + writer = _writer(events, sanitize=True) + connection = _Connection(events) + callback = AsyncMock() + + async def reject_fingerprint(*_args, **_kwargs): + assert connection.in_transaction + events.append("fingerprint") + raise EmbeddingFingerprintMismatchError(f"Embedding fingerprint mismatch for bank {bank_id!r}") + + monkeypatch.setattr(writer_module, "ensure_bank_embedding_fingerprint", reject_fingerprint) + monkeypatch.setattr( + writer_module.fact_storage, + "upsert_document_metadata", + AsyncMock(side_effect=AssertionError("mutation must not run")), + ) + + with pytest.raises(EmbeddingFingerprintMismatchError) as error: + await _unit_of_work(writer, connection).execute( + MetadataOnlyWriteRequest( + bank_id=bank_id, + document_id=document_id, + expected_content_hash="hash", + combined_content="payload", + input_slot_count=1, + outbox_callback=callback, + ), + ) + + assert bank_id not in str(error.value) + assert "" in str(error.value) + assert events == ["begin", "fingerprint", "rollback"] + callback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_sanitized_fingerprint_warning_hides_bank_identifier(monkeypatch, caplog) -> None: + bank_id = "sensitive-multimodal-bank" + stored = fingerprint_module.build_embedding_fingerprint(_embedding_model()) + current_model = SimpleNamespace( + provider_name="openai", + model="text-embedding-3-small", + dimension=384, + normalization=True, + ) + monkeypatch.setattr( + fingerprint_module, + "_bank_state", + AsyncMock(return_value=(stored, True)), + ) + sanitizer = IdentifierSanitizer.from_values(enabled=True, values=(bank_id,)) + caplog.set_level(logging.WARNING, logger=fingerprint_module.__name__) + + await fingerprint_module.ensure_bank_embedding_fingerprint( + object(), + bank_id, + current_model, + policy="warn", + log_sanitizer=sanitizer, + ) + + assert "" in caplog.text + assert bank_id not in caplog.text + + +@pytest.mark.asyncio +async def test_sanitized_full_replacement_hides_bank_and_document_logs(monkeypatch, caplog) -> None: + bank_id = "sensitive-multimodal-bank" + document_id = "sensitive-multimodal-document" + source_id = uuid.uuid4() + observation_id = uuid.uuid4() + events: list[str] = [] + + class ReplacementConnection(_Connection): + async def fetch(self, query, *_args): + if "unit_entities" in query: + return [] + if "fact_type IN ('experience', 'world')" in query: + return [{"id": source_id}] + if "fact_type = 'observation'" in query: + return [{"id": observation_id, "source_memory_ids": []}] + raise AssertionError(f"Unexpected fetch query: {query}") + + async def execute(self, _query, *_args): + return "DELETE 1" + + async def fetchval(self, _query, *_args): + return None + + connection = ReplacementConnection(events) + ops = SimpleNamespace( + uses_observation_sources_table=False, + refresh_entity_fact_counts=AsyncMock(), + ) + writer = PersistenceWriter( + pool=SimpleNamespace(ops=ops), + embeddings_model=_embedding_model(), + entity_resolver=object(), + config=_config(), + ownership=_Ownership(events), + sanitize_log_identifiers=True, + ) + + async def ensure_fingerprint(*_args, **kwargs): + assert kwargs["log_sanitizer"].identifier(bank_id) == "" + + async def insert_facts(*_args, **_kwargs): + return [[]], object() + + monkeypatch.setattr(writer_module, "ensure_bank_embedding_fingerprint", ensure_fingerprint) + monkeypatch.setattr(fact_storage, "_upsert_document_row", AsyncMock()) + monkeypatch.setattr(writer_module.runtime, "insert_facts_and_links", insert_facts) + monkeypatch.setattr(writer, "_finalize_entity_graph", AsyncMock(return_value=CoreGraphWrite())) + caplog.set_level(logging.INFO, logger=fact_storage.__name__) + + await _unit_of_work(writer, connection).execute( + WriteWindowRequest( + bank_id=bank_id, + document_id=document_id, + document_window=FirstFullWriteWindow(combined_content="canonical multimodal payload"), + contents=(RetainContent(content="canonical multimodal payload"),), + ), + ) + + assert "" in caplog.text + assert bank_id not in caplog.text + assert document_id not in caplog.text + + +@pytest.mark.asyncio +async def test_sanitized_warning_hides_identifiers_and_exception_text( + monkeypatch, + caplog, +) -> None: + bank_id = "sensitive-bank" + operation_id = "sensitive-operation" + document_id = "sensitive-document" + invocation = RetainInvocation( + bank_id=bank_id, + raw_contents=({"content": "payload"},), + request_context=object(), + operation_id=operation_id, + sanitize_log_identifiers=True, + ) + + @asynccontextmanager + async def failing_scope(*_args, **_kwargs): + raise RuntimeError(f"{bank_id}/{operation_id}/{document_id}") + yield # pragma: no cover + + monkeypatch.setattr(service_module, "acquire_with_retry", failing_scope) + caplog.set_level(logging.WARNING, logger=service_module.__name__) + + await service_module.RetainPipelineService()._record_document_ids( + invocation, + _execution(), + (SimpleNamespace(document_id=document_id),), + ) + + assert "" in caplog.text + assert bank_id not in caplog.text + assert operation_id not in caplog.text + assert document_id not in caplog.text + assert caplog.records[-1].exc_info is None + + +@pytest.mark.asyncio +async def test_final_ann_failure_is_reported_as_incomplete(monkeypatch) -> None: + invocation = RetainInvocation( + bank_id="bank", + raw_contents=({"content": "payload"},), + request_context=object(), + ) + plan = SimpleNamespace(intent=SimpleNamespace(document_id="doc")) + monkeypatch.setattr( + service_module, + "run_final_semantic_ann", + AsyncMock(side_effect=RuntimeError("temporary ANN failure")), + ) + + completed = await service_module.RetainPipelineService()._run_full_semantic_ann_best_effort( + invocation, + _execution(), + plan, + ["unit-a"], + ) + + assert completed is False + + +@pytest.mark.asyncio +async def test_recovery_preserves_final_ann_retry_marker_after_failure(monkeypatch) -> None: + invocation = RetainInvocation( + bank_id="bank", + raw_contents=({"content": "payload"},), + request_context=object(), + operation_id="operation", + ) + plan = SimpleNamespace( + recovered_unit_bindings=(SimpleNamespace(unit_id="unit-a"),), + final_ann_pending=True, + intent=SimpleNamespace(document_id="doc"), + ) + pipeline = service_module.RetainPipelineService() + run_final_ann = AsyncMock(return_value=False) + record_completed = AsyncMock() + monkeypatch.setattr(pipeline, "_recovery_result_buckets", lambda _plan: (("unit-a",),)) + monkeypatch.setattr(pipeline, "_run_full_semantic_ann_best_effort", run_final_ann) + monkeypatch.setattr(pipeline, "_record_final_ann_completed_best_effort", record_completed) + + outcome = await pipeline._resume_committed_document(invocation, _execution(), plan) + + assert outcome.unit_ids_by_content == (("unit-a",),) + run_final_ann.assert_awaited_once() + record_completed.assert_not_awaited() diff --git a/core/dataplane/tests/test_ingestion_postgresql_live.py b/core/dataplane/tests/test_ingestion_postgresql_live.py new file mode 100644 index 0000000..8fbefed --- /dev/null +++ b/core/dataplane/tests/test_ingestion_postgresql_live.py @@ -0,0 +1,675 @@ +"""Live PostgreSQL smoke coverage for the Retain ingestion pipeline.""" + +from __future__ import annotations + +import hashlib +import json +import uuid +from unittest.mock import AsyncMock + +import pytest +from hms_api import MemoryEngine, RequestContext +from hms_api.config import clear_config_cache +from hms_api.engine.cross_encoder import RRFPassthroughCrossEncoder +from hms_api.engine.embeddings import Embeddings +from hms_api.engine.ingestion import RetainOperationInactiveError +from hms_api.engine.ingestion.persistence.postgres import PostgresDocumentOwnership +from hms_api.engine.memory_engine import Budget, _RetainOperationCancelled +from hms_api.engine.query_analyzer import DateparserQueryAnalyzer +from hms_api.engine.task_backend import SyncTaskBackend +from hms_api.worker.poller import WorkerPoller + + +class _DeterministicEmbeddings(Embeddings): + """Small deterministic vectors with the production schema dimension.""" + + model_name = "hms-postgresql-live-hash-v1" + + @property + def provider_name(self) -> str: + return "postgresql-live-test" + + @property + def dimension(self) -> int: + return 384 + + async def initialize(self) -> None: + return None + + def encode(self, texts: list[str]) -> list[list[float]]: + vectors: list[list[float]] = [] + for text in texts: + digest = hashlib.sha256(text.encode()).digest() + vectors.append([((digest[index % len(digest)] / 255.0) * 2.0) - 1.0 for index in range(self.dimension)]) + return vectors + + +@pytest.mark.asyncio +async def test_postgresql_retain_fresh_delta_and_recall(pg0_db_url: str) -> None: + """Exercise durable fresh, no-op, replacement, and retrieval paths.""" + + memory = MemoryEngine( + db_url=pg0_db_url, + memory_llm_provider="none", + memory_llm_model="none", + embeddings=_DeterministicEmbeddings(), + cross_encoder=RRFPassthroughCrossEncoder(), + query_analyzer=DateparserQueryAnalyzer(), + pool_min_size=1, + pool_max_size=3, + run_migrations=False, + task_backend=SyncTaskBackend(), + skip_llm_verification=True, + ) + await memory.initialize() + + bank_id = f"postgresql-retain-live-{uuid.uuid4().hex[:12]}" + document_id = "profile" + request_context = RequestContext() + original = "Alice maintains the Atlas search service." + replacement = "Alice maintains the Atlas search service. Bob owns incident response." + + try: + first_ids = await memory.retain_async( + bank_id=bank_id, + content=original, + document_id=document_id, + request_context=request_context, + ) + assert first_ids + + unchanged_ids = await memory.retain_async( + bank_id=bank_id, + content=original, + document_id=document_id, + request_context=request_context, + ) + assert unchanged_ids == [] + + replacement_ids = await memory.retain_async( + bank_id=bank_id, + content=replacement, + document_id=document_id, + request_context=request_context, + ) + assert replacement_ids + + document = await memory.get_document( + document_id, + bank_id, + request_context=request_context, + ) + assert document is not None + assert document["original_text"] == replacement + + recalled = await memory.recall_async( + bank_id=bank_id, + query="Who owns incident response?", + budget=Budget.LOW, + max_tokens=512, + request_context=request_context, + ) + assert recalled.results + assert any("incident response" in result.text for result in recalled.results) + finally: + await memory.delete_bank(bank_id, request_context=request_context) + await memory.close() + + +@pytest.mark.asyncio +async def test_postgresql_retain_replaces_an_existing_unhashed_document(pg0_db_url: str) -> None: + """Upgrade an unhashed row through the same locked full-replacement path.""" + + memory = MemoryEngine( + db_url=pg0_db_url, + memory_llm_provider="none", + memory_llm_model="none", + embeddings=_DeterministicEmbeddings(), + cross_encoder=RRFPassthroughCrossEncoder(), + query_analyzer=DateparserQueryAnalyzer(), + pool_min_size=1, + pool_max_size=3, + run_migrations=False, + task_backend=SyncTaskBackend(), + skip_llm_verification=True, + ) + await memory.initialize() + + bank_id = f"postgresql-retain-unhashed-{uuid.uuid4().hex[:12]}" + document_id = "upgrade-document" + request_context = RequestContext() + original = "Original-only-token belongs to the old document." + replacement = "Replacement-only-token belongs to the upgraded document." + + try: + await memory.retain_async( + bank_id=bank_id, + content=original, + document_id=document_id, + request_context=request_context, + ) + async with memory._pool.acquire() as connection: + await connection.execute( + "UPDATE documents SET content_hash = NULL WHERE id = $1 AND bank_id = $2", + document_id, + bank_id, + ) + + replacement_ids = await memory.retain_async( + bank_id=bank_id, + content=replacement, + document_id=document_id, + request_context=request_context, + ) + + assert replacement_ids + async with memory._pool.acquire() as connection: + document = await connection.fetchrow( + """ + SELECT original_text, content_hash + FROM documents + WHERE id = $1 AND bank_id = $2 + """, + document_id, + bank_id, + ) + facts = await connection.fetch( + """ + SELECT text + FROM memory_units + WHERE document_id = $1 AND bank_id = $2 + """, + document_id, + bank_id, + ) + assert document is not None + assert document["original_text"] == replacement + assert document["content_hash"] == hashlib.sha256(replacement.encode()).hexdigest() + assert facts + assert all("Original-only-token" not in row["text"] for row in facts) + + # Once a competing writer publishes a hash, an older unhashed snapshot + # can no longer acquire replacement ownership. + ownership = PostgresDocumentOwnership() + async with memory._pool.acquire() as connection: + async with connection.transaction(): + await connection.execute( + "UPDATE documents SET content_hash = NULL WHERE id = $1 AND bank_id = $2", + document_id, + bank_id, + ) + assert await ownership.validate_unhashed_window( + connection, + bank_id=bank_id, + document_id=document_id, + ) + await connection.execute( + "UPDATE documents SET content_hash = $3 WHERE id = $1 AND bank_id = $2", + document_id, + bank_id, + "new-owner-hash", + ) + async with memory._pool.acquire() as connection: + async with connection.transaction(): + assert not await ownership.validate_unhashed_window( + connection, + bank_id=bank_id, + document_id=document_id, + ) + finally: + await memory.delete_bank(bank_id, request_context=request_context) + await memory.close() + + +@pytest.mark.asyncio +async def test_postgresql_token_batching_preserves_document_groups( + pg0_db_url: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Keep repeated IDs together while restoring the caller's result order.""" + + monkeypatch.setenv("HMS_API_RETAIN_BATCH_TOKENS", "1") + clear_config_cache() + memory = MemoryEngine( + db_url=pg0_db_url, + memory_llm_provider="none", + memory_llm_model="none", + embeddings=_DeterministicEmbeddings(), + cross_encoder=RRFPassthroughCrossEncoder(), + query_analyzer=DateparserQueryAnalyzer(), + pool_min_size=1, + pool_max_size=3, + run_migrations=False, + task_backend=SyncTaskBackend(), + skip_llm_verification=True, + ) + await memory.initialize() + + bank_id = f"postgresql-retain-groups-{uuid.uuid4().hex[:12]}" + request_context = RequestContext() + + try: + results = await memory.retain_batch_async( + bank_id=bank_id, + contents=[ + {"content": "Alpha first.", "document_id": "doc-alpha"}, + {"content": "Beta only.", "document_id": "doc-beta"}, + {"content": "Alpha second.", "document_id": "doc-alpha"}, + ], + request_context=request_context, + ) + + assert len(results) == 3 + assert all(results) + alpha = await memory.get_document("doc-alpha", bank_id, request_context=request_context) + beta = await memory.get_document("doc-beta", bank_id, request_context=request_context) + assert alpha is not None + assert beta is not None + assert "Alpha first." in alpha["original_text"] + assert "Alpha second." in alpha["original_text"] + assert "Beta only." not in alpha["original_text"] + assert "Beta only." in beta["original_text"] + finally: + await memory.delete_bank(bank_id, request_context=request_context) + await memory.close() + clear_config_cache() + + +@pytest.mark.asyncio +async def test_postgresql_cancellation_is_not_reported_as_completion( + pg0_db_url: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Propagate a between-document cancellation and preserve terminal state.""" + + monkeypatch.setenv("HMS_API_RETAIN_BATCH_TOKENS", "1") + clear_config_cache() + memory = MemoryEngine( + db_url=pg0_db_url, + memory_llm_provider="none", + memory_llm_model="none", + embeddings=_DeterministicEmbeddings(), + cross_encoder=RRFPassthroughCrossEncoder(), + query_analyzer=DateparserQueryAnalyzer(), + pool_min_size=1, + pool_max_size=3, + run_migrations=False, + task_backend=SyncTaskBackend(), + skip_llm_verification=True, + ) + await memory.initialize() + + bank_id = f"postgresql-retain-cancel-{uuid.uuid4().hex[:12]}" + operation_id = uuid.uuid4() + request_context = RequestContext() + outbox_callback = AsyncMock() + + try: + await memory.get_bank_profile(bank_id=bank_id, request_context=request_context) + async with memory._pool.acquire() as connection: + await connection.execute( + """ + INSERT INTO async_operations (operation_id, bank_id, operation_type, status) + VALUES ($1, $2, 'retain', 'processing') + """, + operation_id, + bank_id, + ) + + monkeypatch.setattr(memory, "_check_op_alive", AsyncMock(side_effect=[True, False])) + with pytest.raises(_RetainOperationCancelled, match="cancelled between logical document batches"): + await memory.retain_batch_async( + bank_id=bank_id, + contents=[ + {"content": "Committed before cancellation.", "document_id": "doc-before"}, + {"content": "Must not be committed.", "document_id": "doc-after"}, + ], + request_context=request_context, + operation_id=str(operation_id), + outbox_callback=outbox_callback, + ) + outbox_callback.assert_not_awaited() + + async with memory._pool.acquire() as connection: + await connection.execute( + "UPDATE async_operations SET status = 'cancelled' WHERE operation_id = $1", + operation_id, + ) + await memory._mark_operation_completed(str(operation_id)) + async with memory._pool.acquire() as connection: + status = await connection.fetchval( + "SELECT status FROM async_operations WHERE operation_id = $1", + operation_id, + ) + assert status == "cancelled" + finally: + await memory.delete_bank(bank_id, request_context=request_context) + await memory.close() + clear_config_cache() + + +@pytest.mark.asyncio +async def test_postgresql_parent_cancellation_is_terminal_for_children(pg0_db_url: str) -> None: + """Cancel active child rows atomically and reject late worker completion.""" + + memory = MemoryEngine( + db_url=pg0_db_url, + memory_llm_provider="none", + memory_llm_model="none", + embeddings=_DeterministicEmbeddings(), + cross_encoder=RRFPassthroughCrossEncoder(), + query_analyzer=DateparserQueryAnalyzer(), + pool_min_size=1, + pool_max_size=3, + run_migrations=False, + task_backend=SyncTaskBackend(), + skip_llm_verification=True, + ) + await memory.initialize() + + bank_id = f"postgresql-retain-parent-cancel-{uuid.uuid4().hex[:12]}" + parent_id = uuid.uuid4() + pending_child_id = uuid.uuid4() + processing_child_id = uuid.uuid4() + request_context = RequestContext() + + try: + await memory.get_bank_profile(bank_id=bank_id, request_context=request_context) + async with memory._pool.acquire() as connection: + await connection.execute( + """ + INSERT INTO async_operations + (operation_id, bank_id, operation_type, status, result_metadata) + VALUES ($1, $2, 'batch_retain', 'pending', $3::jsonb) + """, + parent_id, + bank_id, + json.dumps({"is_parent": True, "num_sub_batches": 2, "items_count": 2}), + ) + for child_id, status, index in ( + (pending_child_id, "pending", 1), + (processing_child_id, "processing", 2), + ): + await connection.execute( + """ + INSERT INTO async_operations + (operation_id, bank_id, operation_type, status, result_metadata) + VALUES ($1, $2, 'retain', $3, $4::jsonb) + """, + child_id, + bank_id, + status, + json.dumps( + { + "parent_operation_id": str(parent_id), + "sub_batch_index": index, + "total_sub_batches": 2, + } + ), + ) + + await memory.cancel_operation( + bank_id=bank_id, + operation_id=str(parent_id).upper(), + request_context=request_context, + ) + + poller = WorkerPoller( + backend=memory._backend, + worker_id="postgresql-retain-live", + executor=memory.execute_task, + ) + await poller._mark_completed(str(processing_child_id), None) + await memory._mark_operation_completed(str(pending_child_id)) + + async with memory._pool.acquire() as connection: + rows = await connection.fetch( + """ + SELECT operation_id, status + FROM async_operations + WHERE operation_id = ANY($1::uuid[]) + """, + [parent_id, pending_child_id, processing_child_id], + ) + assert {row["operation_id"]: row["status"] for row in rows} == { + parent_id: "cancelled", + pending_child_id: "cancelled", + processing_child_id: "cancelled", + } + finally: + await memory.delete_bank(bank_id, request_context=request_context) + await memory.close() + + +@pytest.mark.asyncio +async def test_postgresql_accepted_cancellation_blocks_core_retain_and_outbox(pg0_db_url: str) -> None: + """Reject a stale processing child inside its next core write transaction.""" + + memory = MemoryEngine( + db_url=pg0_db_url, + memory_llm_provider="none", + memory_llm_model="none", + embeddings=_DeterministicEmbeddings(), + cross_encoder=RRFPassthroughCrossEncoder(), + query_analyzer=DateparserQueryAnalyzer(), + pool_min_size=1, + pool_max_size=3, + run_migrations=False, + task_backend=SyncTaskBackend(), + skip_llm_verification=True, + ) + await memory.initialize() + + bank_id = f"postgresql-retain-write-fence-{uuid.uuid4().hex[:12]}" + document_id = "must-not-exist" + parent_id = uuid.uuid4() + child_id = uuid.uuid4() + request_context = RequestContext() + outbox_callback = AsyncMock() + + try: + await memory.get_bank_profile(bank_id=bank_id, request_context=request_context) + async with memory._pool.acquire() as connection: + await connection.execute( + """ + INSERT INTO async_operations + (operation_id, bank_id, operation_type, status, result_metadata) + VALUES ($1, $2, 'batch_retain', 'pending', $3::jsonb) + """, + parent_id, + bank_id, + json.dumps({"is_parent": True, "num_sub_batches": 1, "items_count": 1}), + ) + await connection.execute( + """ + INSERT INTO async_operations + (operation_id, bank_id, operation_type, status, result_metadata) + VALUES ($1, $2, 'retain', 'processing', $3::jsonb) + """, + child_id, + bank_id, + json.dumps( + { + "parent_operation_id": str(parent_id), + "sub_batch_index": 1, + "total_sub_batches": 1, + } + ), + ) + + await memory.cancel_operation( + bank_id=bank_id, + operation_id=str(parent_id), + request_context=request_context, + ) + + with pytest.raises(RetainOperationInactiveError, match="no longer active"): + await memory._retain_batch_async_internal( + bank_id=bank_id, + contents=[{"content": "This content must never commit.", "document_id": document_id}], + request_context=request_context, + operation_id=str(child_id), + outbox_callback=outbox_callback, + ) + + outbox_callback.assert_not_awaited() + async with memory._pool.acquire() as connection: + child_status = await connection.fetchval( + "SELECT status FROM async_operations WHERE operation_id = $1", + child_id, + ) + counts = await connection.fetchrow( + """ + SELECT + (SELECT COUNT(*) FROM documents WHERE id = $1 AND bank_id = $2) AS documents, + (SELECT COUNT(*) FROM chunks WHERE document_id = $1 AND bank_id = $2) AS chunks, + (SELECT COUNT(*) FROM memory_units WHERE document_id = $1 AND bank_id = $2) AS memories + """, + document_id, + bank_id, + ) + assert child_status == "cancelled" + assert dict(counts) == {"documents": 0, "chunks": 0, "memories": 0} + finally: + await memory.delete_bank(bank_id, request_context=request_context) + await memory.close() + + +@pytest.mark.asyncio +async def test_postgresql_batch_retry_reopens_only_retryable_work(pg0_db_url: str) -> None: + """Retry a terminal batch or child without reviving completed work.""" + + memory = MemoryEngine( + db_url=pg0_db_url, + memory_llm_provider="none", + memory_llm_model="none", + embeddings=_DeterministicEmbeddings(), + cross_encoder=RRFPassthroughCrossEncoder(), + query_analyzer=DateparserQueryAnalyzer(), + pool_min_size=1, + pool_max_size=3, + run_migrations=False, + task_backend=SyncTaskBackend(), + skip_llm_verification=True, + ) + await memory.initialize() + + bank_id = f"postgresql-retain-batch-retry-{uuid.uuid4().hex[:12]}" + parent_id = uuid.uuid4() + completed_child_id = uuid.uuid4() + failed_child_id = uuid.uuid4() + cancelled_child_id = uuid.uuid4() + request_context = RequestContext() + + try: + await memory.get_bank_profile(bank_id=bank_id, request_context=request_context) + async with memory._pool.acquire() as connection: + await connection.execute( + """ + INSERT INTO async_operations + (operation_id, bank_id, operation_type, status, result_metadata) + VALUES ($1, $2, 'batch_retain', 'failed', $3::jsonb) + """, + parent_id, + bank_id, + json.dumps({"is_parent": True, "num_sub_batches": 3, "items_count": 3}), + ) + for child_id, status, index in ( + (completed_child_id, "completed", 1), + (failed_child_id, "failed", 2), + (cancelled_child_id, "cancelled", 3), + ): + await connection.execute( + """ + INSERT INTO async_operations + (operation_id, bank_id, operation_type, status, result_metadata, task_payload) + VALUES ($1, $2, 'retain', $3, $4::jsonb, $5::jsonb) + """, + child_id, + bank_id, + status, + json.dumps( + { + "parent_operation_id": str(parent_id), + "sub_batch_index": index, + "total_sub_batches": 3, + } + ), + json.dumps( + { + "type": "batch_retain", + "operation_id": str(child_id), + "bank_id": bank_id, + "contents": [], + } + ), + ) + + await memory.retry_operation( + bank_id=bank_id, + operation_id=str(parent_id).upper(), + request_context=request_context, + ) + + async with memory._pool.acquire() as connection: + rows = await connection.fetch( + """ + SELECT operation_id, status + FROM async_operations + WHERE operation_id = ANY($1::uuid[]) + """, + [parent_id, completed_child_id, failed_child_id, cancelled_child_id], + ) + assert {row["operation_id"]: row["status"] for row in rows} == { + parent_id: "pending", + completed_child_id: "completed", + failed_child_id: "pending", + cancelled_child_id: "pending", + } + + # A child-level retry safely reopens a failed/cancelled parent, but + # still preserves completed siblings. + async with memory._pool.acquire() as connection: + await connection.execute( + """ + UPDATE async_operations + SET status = CASE + WHEN operation_id = $1 THEN 'cancelled' + WHEN operation_id = $2 THEN 'cancelled' + WHEN operation_id = $3 THEN 'cancelled' + WHEN operation_id = $4 THEN 'completed' + ELSE status + END + WHERE operation_id = ANY($5::uuid[]) + """, + parent_id, + failed_child_id, + cancelled_child_id, + completed_child_id, + [parent_id, failed_child_id, cancelled_child_id, completed_child_id], + ) + + await memory.retry_operation( + bank_id=bank_id, + operation_id=str(failed_child_id), + request_context=request_context, + ) + + async with memory._pool.acquire() as connection: + rows = await connection.fetch( + """ + SELECT operation_id, status + FROM async_operations + WHERE operation_id = ANY($1::uuid[]) + """, + [parent_id, completed_child_id, failed_child_id, cancelled_child_id], + ) + assert {row["operation_id"]: row["status"] for row in rows} == { + parent_id: "pending", + completed_child_id: "completed", + failed_child_id: "pending", + cancelled_child_id: "cancelled", + } + finally: + await memory.delete_bank(bank_id, request_context=request_context) + await memory.close() diff --git a/core/dataplane/tests/test_link_utils.py b/core/dataplane/tests/test_link_utils.py index cbbf86b..318e0d9 100644 --- a/core/dataplane/tests/test_link_utils.py +++ b/core/dataplane/tests/test_link_utils.py @@ -1,17 +1,19 @@ """Tests for link_utils datetime handling, temporal link computation, and semantic link splitting.""" + +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock, MagicMock + import numpy as np import pytest -from datetime import datetime, timezone, timedelta -from unittest.mock import AsyncMock, MagicMock from hms_api.engine.retain.link_utils import ( - _normalize_datetime, + MAX_TEMPORAL_LINKS_PER_UNIT, _cap_links_per_unit, - compute_temporal_links, - compute_temporal_query_bounds, + _normalize_datetime, compute_semantic_links_ann, compute_semantic_links_within_batch, - MAX_TEMPORAL_LINKS_PER_UNIT, + compute_temporal_links, + compute_temporal_query_bounds, ) @@ -172,8 +174,8 @@ def test_weight_decreases_with_distance(self): links = compute_temporal_links(units, candidates, time_window_hours=24) assert len(links) == 2 - close_link = next(l for l in links if l[1] == "close") - far_link = next(l for l in links if l[1] == "far") + close_link = next(link for link in links if link[1] == "close") + far_link = next(link for link in links if link[1] == "far") assert close_link[3] > far_link[3] @@ -206,8 +208,8 @@ def test_multiple_units_multiple_candidates(self): # unit-1 should link to c1 only # unit-2 should link to c2 only - unit1_links = [l for l in links if l[0] == "unit-1"] - unit2_links = [l for l in links if l[0] == "unit-2"] + unit1_links = [link for link in links if link[0] == "unit-1"] + unit2_links = [link for link in links if link[0] == "unit-2"] assert len(unit1_links) == 1 assert unit1_links[0][1] == "c1" @@ -374,6 +376,7 @@ def test_top_k_limits_per_unit(self): links = compute_semantic_links_within_batch(unit_ids, embs, top_k=3, threshold=0.5) # Each unit should have at most 3 outgoing links from collections import Counter + from_counts = Counter(lnk[0] for lnk in links) for count in from_counts.values(): assert count <= 3 @@ -527,8 +530,33 @@ async def test_uses_set_local_for_ef_search(self, mock_conn): ef_statements = [s for s in executed_sql if "hnsw.ef_search" in s] assert ef_statements, "ef_search must be tuned down for retain ANN" for stmt in ef_statements: - assert stmt.strip().startswith("SET LOCAL"), ( - f"hnsw.ef_search must use SET LOCAL, got: {stmt}" - ) + assert stmt.strip().startswith("SET LOCAL"), f"hnsw.ef_search must use SET LOCAL, got: {stmt}" # And there must not be a RESET — SET LOCAL handles it at commit. assert not any("RESET hnsw.ef_search" in s for s in executed_sql) + + @pytest.mark.asyncio + async def test_read_only_mode_uses_array_cte_without_ddl(self, mock_conn): + """Retain planning must remain valid after ``SET TRANSACTION READ ONLY``.""" + mock_conn.fetch.return_value = [{"from_id": "u1", "to_id": "existing-1", "similarity": 0.85}] + emb = [0.1] * 384 + + result = await compute_semantic_links_ann( + conn=mock_conn, + bank_id="bank-1", + unit_ids=["u1"], + embeddings=[emb], + fact_types=["world"], + threshold=0.7, + read_only=True, + ) + + assert result == [("u1", "existing-1", "semantic", 0.85, None)] + mock_conn.transaction.assert_not_called() + mock_conn.copy_records_to_table.assert_not_called() + mock_conn.execute.assert_awaited_once_with("SET LOCAL hnsw.ef_search = 60") + + query = mock_conn.fetch.await_args.args[0] + assert "WITH seeds" in query + assert "unnest($4::text[], $5::text[])" in query + assert "CREATE" not in query + assert "_ann_seeds" not in query diff --git a/core/dataplane/tests/test_multimodal_admission.py b/core/dataplane/tests/test_multimodal_admission.py index 733aa68..6b699b1 100644 --- a/core/dataplane/tests/test_multimodal_admission.py +++ b/core/dataplane/tests/test_multimodal_admission.py @@ -412,7 +412,7 @@ async def test_parser_chain_order_forms_a_distinct_document_command(monkeypatch) @pytest.mark.asyncio -async def test_legacy_anonymous_file_retain_keeps_random_uuid_identity(monkeypatch) -> None: +async def test_non_multimodal_anonymous_file_retain_uses_random_uuid_identity(monkeypatch) -> None: storage = _CapturingStorage() engine = object.__new__(MemoryEngine) engine._backend = SimpleNamespace() diff --git a/core/dataplane/tests/test_multimodal_engine_bridge.py b/core/dataplane/tests/test_multimodal_engine_bridge.py index 617b06c..8134b59 100644 --- a/core/dataplane/tests/test_multimodal_engine_bridge.py +++ b/core/dataplane/tests/test_multimodal_engine_bridge.py @@ -7,12 +7,13 @@ from unittest.mock import AsyncMock, MagicMock from uuid import UUID -import pytest - import hms_api.engine.memory_engine as memory_engine_module +import pytest +from hms_api.engine.ingestion import RetainOutcome, RetainPipelineService +from hms_api.engine.ingestion.normalization import normalize_contents from hms_api.engine.memory_engine import MemoryEngine from hms_api.engine.parsers import ConvertResult -from hms_api.engine.retain.orchestrator import _build_contents +from hms_api.engine.response_models import TokenUsage from hms_api.worker.exceptions import DeferOperation @@ -601,7 +602,10 @@ async def test_rich_parser_output_is_whitelisted_into_existing_child_retain( assert child_payload["document_tags"] == ["engineering"] assert content["metadata"]["customer_key"] == "kept" assert content["metadata"]["media_kind"] == "image" - merged_content = _build_contents(child_payload["contents"], child_payload["document_tags"])[0] + merged_content = normalize_contents( + child_payload["contents"], + document_tags=child_payload["document_tags"], + )[0] assert set(merged_content.tags) == {"project-a", "engineering"} assert content["entities"] == [{"text": "Python", "type": "CONCEPT"}] assert content["event_date"] is None @@ -651,7 +655,6 @@ async def test_explicit_public_strategy_cannot_override_canonical_chunks(monkeyp @pytest.mark.asyncio async def test_trusted_chunks_override_is_applied_after_public_strategy(monkeypatch) -> None: import hms_api.config_resolver - from hms_api.engine.retain import orchestrator engine = object.__new__(MemoryEngine) resolved = SimpleNamespace( @@ -687,12 +690,13 @@ def apply_strategy(config, strategy): captured = {} - async def retain_batch(**kwargs): - captured.update(kwargs) - return [[]], SimpleNamespace(), 0 + async def retain_pipeline(_service, invocation, execution): + captured["invocation"] = invocation + captured["execution"] = execution + return RetainOutcome([[]], TokenUsage(), 0) monkeypatch.setattr(hms_api.config_resolver, "apply_strategy", apply_strategy) - monkeypatch.setattr(orchestrator, "retain_batch", retain_batch) + monkeypatch.setattr(RetainPipelineService, "retain", retain_pipeline) monkeypatch.setattr(memory_engine_module, "create_operation_span", lambda *args, **kwargs: nullcontext()) await MemoryEngine._retain_batch_async_internal( @@ -704,9 +708,10 @@ async def retain_batch(**kwargs): _retain_extraction_mode="chunks", ) - assert captured["config"].retain_extraction_mode == "chunks" - assert captured["config"].enable_observations is False - assert captured["config"].retain_chunk_size == 2_400 + assert captured["execution"].resolved_config.retain_extraction_mode == "chunks" + assert captured["execution"].resolved_config.enable_observations is False + assert captured["execution"].resolved_config.retain_chunk_size == 2_400 + assert captured["invocation"].sanitize_log_identifiers is True @pytest.mark.asyncio @@ -843,7 +848,7 @@ async def get_backend(): @pytest.mark.asyncio -async def test_legacy_child_failure_keeps_existing_parent_propagation(monkeypatch) -> None: +async def test_standard_child_failure_keeps_parent_propagation(monkeypatch) -> None: operation_id = "00000000-0000-0000-0000-000000000022" connection = _Connection( [ @@ -869,8 +874,8 @@ async def get_backend(): return SimpleNamespace() engine._get_backend = get_backend - await MemoryEngine._mark_operation_failed(engine, operation_id, "legacy failure", "legacy traceback") + await MemoryEngine._mark_operation_failed(engine, operation_id, "retain failure", "retain traceback") - legacy_update = next(item for item in connection.executions if "status <> 'cancelled'" in item[0]) - assert legacy_update[1][0] == UUID(operation_id) + status_update = next(item for item in connection.executions if "status IN ('pending', 'processing')" in item[0]) + assert status_update[1][0] == UUID(operation_id) engine._maybe_update_parent_operation.assert_awaited_once_with(operation_id, connection) diff --git a/core/dataplane/tests/test_multimodal_security.py b/core/dataplane/tests/test_multimodal_security.py index 654ac0e..16f98fd 100644 --- a/core/dataplane/tests/test_multimodal_security.py +++ b/core/dataplane/tests/test_multimodal_security.py @@ -24,6 +24,7 @@ import hms_api.engine.memory_engine as memory_engine_module import hms_api.engine.parsers.openai_multimodal as parser_module from hms_api.engine.memory_engine import MemoryEngine +from hms_api.engine.ingestion.redaction import IdentifierSanitizer from hms_api.engine.multimodal import ( GroundedStatement, MediaAsset, @@ -51,7 +52,6 @@ MultimodalParserConfig, OpenAIMultimodalParser, ) -from hms_api.engine.retain.orchestrator import _RetainLogBuffer from hms_api.metrics import MetricsCollector _FRAME_BYTES = b"HMS_FRAME_BYTES_SENTINEL_ignore_previous_instructions" @@ -277,25 +277,29 @@ def _assert_sentinels_absent(surface: object) -> None: assert sentinel not in rendered -def test_retain_log_buffer_redacts_identifiers_only_when_enabled() -> None: +def test_retain_identifier_sanitizer_redacts_only_when_enabled() -> None: bank_id = "bank-private-security-sentinel" document_id = "document-private-security-sentinel" - sanitized = _RetainLogBuffer(sanitized=True, secrets=(bank_id,)) - sanitized.append(f"bank={bank_id}") - sanitized.add_secret(document_id) - sanitized.append(f"bank={bank_id} document={document_id}") - - sanitized_output = "\n".join(sanitized) + sanitized = IdentifierSanitizer.from_values(enabled=True, values=(bank_id,)) + sanitized_output = "\n".join( + ( + sanitized.text(f"bank={bank_id}"), + sanitized.text( + f"bank={bank_id} document={document_id}", + extra_identifiers=(document_id,), + ), + ) + ) assert bank_id not in sanitized_output assert document_id not in sanitized_output assert sanitized_output.count("") == 3 - legacy = _RetainLogBuffer(sanitized=False, secrets=(bank_id,)) - legacy.add_secret(document_id) - legacy.append(f"bank={bank_id} document={document_id}") - - assert legacy == [f"bank={bank_id} document={document_id}"] + unsanitized = IdentifierSanitizer.from_values(enabled=False, values=(bank_id,)) + assert unsanitized.text( + f"bank={bank_id} document={document_id}", + extra_identifiers=(document_id,), + ) == f"bank={bank_id} document={document_id}" @pytest.mark.asyncio diff --git a/core/dataplane/tests/test_observation_invalidation.py b/core/dataplane/tests/test_observation_invalidation.py index 23d2a7a..529f9b2 100644 --- a/core/dataplane/tests/test_observation_invalidation.py +++ b/core/dataplane/tests/test_observation_invalidation.py @@ -328,11 +328,11 @@ async def test_upsert_document_removes_observations_from_outgoing_memories( ) # Trigger the upsert path directly. ``handle_document_tracking`` is - # what the retain orchestrator calls on every document re-ingest. + # what the Retain pipeline calls on every document re-ingest. # Pass ops=memory._backend.ops so the inner observation-cleanup query # selects the PG (native array) read path instead of falling back to # the Oracle junction-table path (which would query a non-existent - # public.observation_sources relation under PG). The orchestrator + # public.observation_sources relation under PG). The pipeline # call sites in _streaming_retain_batch already do this via pool.ops. async with pool.acquire() as conn: async with conn.transaction(): diff --git a/core/dataplane/tests/test_op_cancellation.py b/core/dataplane/tests/test_op_cancellation.py index 0f6d87a..f99ac97 100644 --- a/core/dataplane/tests/test_op_cancellation.py +++ b/core/dataplane/tests/test_op_cancellation.py @@ -8,20 +8,58 @@ - Retain checkpoint: stops between sub-batches if op was deleted """ +import hashlib +import json import uuid from unittest.mock import AsyncMock, patch import pytest import pytest_asyncio - -from hms_api.engine.memory_engine import MemoryEngine - +from hms_api.engine.cross_encoder import RRFPassthroughCrossEncoder +from hms_api.engine.embeddings import Embeddings +from hms_api.engine.memory_engine import MemoryEngine, _RetainOperationCancelled pytestmark = pytest.mark.xdist_group("op_cancellation_tests") _BANK_PREFIX = "test-op-cancel" +class _DeterministicEmbeddings(Embeddings): + """Provide stable vectors without loading a local model.""" + + model_name = "hms-operation-cancellation-test-hash-v1" + + @property + def provider_name(self) -> str: + return "operation-cancellation-test" + + @property + def dimension(self) -> int: + return 384 + + async def initialize(self) -> None: + return None + + def encode(self, texts: list[str]) -> list[list[float]]: + vectors: list[list[float]] = [] + for text in texts: + digest = hashlib.sha256(text.encode()).digest() + vectors.append([((digest[index % len(digest)] / 255.0) * 2.0) - 1.0 for index in range(self.dimension)]) + return vectors + + +@pytest.fixture(scope="session") +def embeddings() -> Embeddings: + """Use deterministic embeddings that need no optional ML dependencies.""" + return _DeterministicEmbeddings() + + +@pytest.fixture(scope="session") +def cross_encoder() -> RRFPassthroughCrossEncoder: + """Use the dependency-free reciprocal-rank fusion reranker.""" + return RRFPassthroughCrossEncoder() + + @pytest_asyncio.fixture async def pool(pg0_db_url): import asyncpg @@ -49,15 +87,22 @@ async def _insert_bank(pool, bank_id: str): ) -async def _insert_op(pool, bank_id: str, op_id: uuid.UUID | None = None) -> uuid.UUID: +async def _insert_op( + pool, + bank_id: str, + op_id: uuid.UUID | None = None, + *, + operation_type: str = "consolidation", +) -> uuid.UUID: op_id = op_id or uuid.uuid4() await pool.execute( """ INSERT INTO async_operations (operation_id, bank_id, operation_type, status) - VALUES ($1, $2, 'consolidation', 'processing') + VALUES ($1, $2, $3, 'processing') """, op_id, bank_id, + operation_type, ) return op_id @@ -172,6 +217,156 @@ async def test_returns_false_after_bank_cascade_delete(self, memory: MemoryEngin assert await memory._check_op_alive(str(op_id)) is False +# --------------------------------------------------------------------------- +# Batch parent/child cancellation +# --------------------------------------------------------------------------- + + +class TestBatchCancellation: + @pytest.mark.asyncio + async def test_parent_cancellation_atomically_cancels_active_children( + self, + memory: MemoryEngine, + pool, + request_context, + ): + bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}" + await _insert_bank(pool, bank_id) + parent_id = uuid.uuid4() + pending_child_id = uuid.uuid4() + processing_child_id = uuid.uuid4() + + await pool.execute( + """ + INSERT INTO async_operations + (operation_id, bank_id, operation_type, status, result_metadata) + VALUES ($1, $2, 'batch_retain', 'pending', $3::jsonb) + """, + parent_id, + bank_id, + json.dumps({"is_parent": True, "num_sub_batches": 2, "items_count": 2}), + ) + for child_id, status, index in ( + (pending_child_id, "pending", 1), + (processing_child_id, "processing", 2), + ): + await pool.execute( + """ + INSERT INTO async_operations + (operation_id, bank_id, operation_type, status, result_metadata) + VALUES ($1, $2, 'retain', $3, $4::jsonb) + """, + child_id, + bank_id, + status, + json.dumps( + { + "parent_operation_id": str(parent_id), + "sub_batch_index": index, + "total_sub_batches": 2, + } + ), + ) + + await memory.cancel_operation( + bank_id=bank_id, + operation_id=str(parent_id), + request_context=request_context, + ) + + rows = await pool.fetch( + """ + SELECT operation_id, status + FROM async_operations + WHERE operation_id = ANY($1::uuid[]) + """, + [parent_id, pending_child_id, processing_child_id], + ) + assert {row["operation_id"]: row["status"] for row in rows} == { + parent_id: "cancelled", + pending_child_id: "cancelled", + processing_child_id: "cancelled", + } + + # A late worker completion must lose the terminal-state CAS. + await memory._mark_operation_completed(str(processing_child_id)) + assert ( + await pool.fetchval( + "SELECT status FROM async_operations WHERE operation_id = $1", + processing_child_id, + ) + == "cancelled" + ) + assert ( + await pool.fetchval( + "SELECT status FROM async_operations WHERE operation_id = $1", + parent_id, + ) + == "cancelled" + ) + + @pytest.mark.asyncio + async def test_direct_child_cancellation_resolves_parent(self, memory: MemoryEngine, pool, request_context): + bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}" + await _insert_bank(pool, bank_id) + parent_id = uuid.uuid4() + completed_child_id = uuid.uuid4() + pending_child_id = uuid.uuid4() + + await pool.execute( + """ + INSERT INTO async_operations + (operation_id, bank_id, operation_type, status, result_metadata) + VALUES ($1, $2, 'batch_retain', 'pending', $3::jsonb) + """, + parent_id, + bank_id, + json.dumps({"is_parent": True, "num_sub_batches": 2, "items_count": 2}), + ) + for child_id, status, index in ( + (completed_child_id, "completed", 1), + (pending_child_id, "pending", 2), + ): + await pool.execute( + """ + INSERT INTO async_operations + (operation_id, bank_id, operation_type, status, result_metadata) + VALUES ($1, $2, 'retain', $3, $4::jsonb) + """, + child_id, + bank_id, + status, + json.dumps( + { + "parent_operation_id": str(parent_id), + "sub_batch_index": index, + "total_sub_batches": 2, + } + ), + ) + + await memory.cancel_operation( + bank_id=bank_id, + operation_id=str(pending_child_id), + request_context=request_context, + ) + + assert ( + await pool.fetchval( + "SELECT status FROM async_operations WHERE operation_id = $1", + pending_child_id, + ) + == "cancelled" + ) + assert ( + await pool.fetchval( + "SELECT status FROM async_operations WHERE operation_id = $1", + parent_id, + ) + == "cancelled" + ) + + # --------------------------------------------------------------------------- # _mark_operation_completed / _mark_operation_failed graceful no-op # --------------------------------------------------------------------------- @@ -184,15 +379,31 @@ async def test_mark_completed_does_not_raise_when_row_missing(self, memory: Memo missing_id = str(uuid.uuid4()) await memory._mark_operation_completed(missing_id) # no exception + @pytest.mark.asyncio + async def test_mark_completed_does_not_overwrite_cancelled(self, memory: MemoryEngine, pool): + bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}" + await _insert_bank(pool, bank_id) + operation_id = await _insert_op(pool, bank_id, operation_type="retain") + await pool.execute( + "UPDATE async_operations SET status = 'cancelled' WHERE operation_id = $1", + operation_id, + ) + + await memory._mark_operation_completed(str(operation_id)) + + status = await pool.fetchval( + "SELECT status FROM async_operations WHERE operation_id = $1", + operation_id, + ) + assert status == "cancelled" + @pytest.mark.asyncio async def test_mark_failed_does_not_raise_when_row_missing(self, memory: MemoryEngine): missing_id = str(uuid.uuid4()) await memory._mark_operation_failed(missing_id, "some error", "traceback here") # no exception @pytest.mark.asyncio - async def test_mark_completed_and_fire_webhook_does_not_raise_when_row_missing( - self, memory: MemoryEngine - ): + async def test_mark_completed_and_fire_webhook_does_not_raise_when_row_missing(self, memory: MemoryEngine): missing_id = str(uuid.uuid4()) await memory._mark_operation_completed_and_fire_webhook( operation_id=missing_id, @@ -265,10 +476,8 @@ async def _fake_check(operation_id: str) -> bool: class TestRetainCheckpoint: @pytest.mark.asyncio - async def test_retain_stops_between_sub_batches_when_cancelled( - self, memory: MemoryEngine, request_context - ): - """retain_batch_async returns partial results if _check_op_alive is False between sub-batches.""" + async def test_retain_stops_between_sub_batches_when_cancelled(self, memory: MemoryEngine, request_context, pool): + """retain_batch_async raises instead of reporting partial success after cancellation.""" from hms_api.config import _get_raw_config bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}" @@ -281,7 +490,10 @@ async def test_retain_stops_between_sub_batches_when_cancelled( config.retain_batch_tokens = 1 try: - op_id = str(uuid.uuid4()) + # Retain persists its exact recovery checkpoint in the tracked + # operation's core transaction. Keep this cancellation fixture + # faithful to the worker path by creating that operation first. + op_id = str(await _insert_op(pool, bank_id, operation_type="retain")) check_calls = 0 async def _fake_check(operation_id: str) -> bool: @@ -291,21 +503,22 @@ async def _fake_check(operation_id: str) -> bool: return check_calls <= 1 contents = [ - {"content": f"Memory item {i} about something interesting."} for i in range(4) + { + "content": f"Memory item {i} about something interesting.", + "document_id": f"document-{i}", + } + for i in range(4) ] with patch.object(memory, "_check_op_alive", side_effect=_fake_check): - result = await memory.retain_batch_async( - bank_id=bank_id, - contents=contents, - request_context=request_context, - operation_id=op_id, - ) + with pytest.raises(_RetainOperationCancelled): + await memory.retain_batch_async( + bank_id=bank_id, + contents=contents, + request_context=request_context, + operation_id=op_id, + ) - # Should have stopped early: fewer results than total items - assert len(result) < len(contents), ( - f"Expected early stop but got {len(result)}/{len(contents)} results" - ) assert check_calls >= 1 finally: config.retain_batch_tokens = original_tokens diff --git a/core/dataplane/tests/test_retain_orchestrator_mapping.py b/core/dataplane/tests/test_retain_orchestrator_mapping.py deleted file mode 100644 index c6f48c3..0000000 --- a/core/dataplane/tests/test_retain_orchestrator_mapping.py +++ /dev/null @@ -1,526 +0,0 @@ -"""Unit tests for retain orchestrator mapping and embeddings length guarantee. - -Regression coverage for issue #1037: a silent length mismatch between the -extracted facts and the generated embeddings caused -`_map_results_to_contents` to raise IndexError during batch_retain. -""" - -from __future__ import annotations - -import asyncio -from contextlib import asynccontextmanager -from datetime import UTC, datetime, timedelta -from types import SimpleNamespace -from unittest.mock import AsyncMock, MagicMock - -import pytest - -from hms_api.engine.retain import embedding_utils, orchestrator -from hms_api.engine.retain.orchestrator import ( - RetainPublicationAborted, - _consume_streaming_batches, - _map_results_to_contents, -) -from hms_api.engine.retain.types import ProcessedFact, RetainContent - - -def _make_processed_fact(content_index: int, text: str = "fact") -> ProcessedFact: - return ProcessedFact( - fact_text=text, - fact_type="world", - embedding=[0.0, 0.0, 0.0], - occurred_start=None, - occurred_end=None, - mentioned_at=datetime(2026, 1, 1), - context="", - metadata={}, - content_index=content_index, - ) - - -def _make_content(text: str = "x") -> RetainContent: - return RetainContent(content=text) - - -class _FakeTransaction: - def __init__(self, connection: "_FakeConnection") -> None: - self.connection = connection - - async def __aenter__(self): - assert not self.connection.in_transaction - self.connection.in_transaction = True - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - self.connection.in_transaction = False - return False - - -class _FakeConnection: - def __init__(self) -> None: - self.in_transaction = False - self.fetchrow = AsyncMock(return_value=None) - self.fetchval = AsyncMock(return_value=None) - self.fetch = AsyncMock(return_value=[]) - self.execute = AsyncMock() - - def transaction(self) -> _FakeTransaction: - return _FakeTransaction(self) - - -def _install_streaming_mocks(monkeypatch: pytest.MonkeyPatch, connection: _FakeConnection): - @asynccontextmanager - async def acquire(_pool): - yield connection - - extract = AsyncMock(return_value=([], [], [], orchestrator.TokenUsage())) - track_document = AsyncMock() - update_document = AsyncMock() - monkeypatch.setattr(orchestrator, "acquire_with_retry", acquire) - monkeypatch.setattr(orchestrator, "ensure_bank_embedding_fingerprint", AsyncMock()) - monkeypatch.setattr(orchestrator, "_extract_and_embed", extract) - monkeypatch.setattr(orchestrator.fact_storage, "handle_document_tracking", track_document) - monkeypatch.setattr(orchestrator.fact_storage, "upsert_document_metadata", update_document) - return extract, track_document, update_document - - -async def _run_streaming_publication( - *, - callback: AsyncMock, - contents_dicts: list[dict], - chunks: list[str], - operation_id: str | None = None, -): - return await orchestrator._streaming_retain_batch( - pool=SimpleNamespace(ops=object()), - embeddings_model=object(), - llm_config=None, - entity_resolver=MagicMock(), - format_date_fn=None, - bank_id="bank", - contents_dicts=contents_dicts, - contents=[_make_content("")], - config=SimpleNamespace( - embedding_fingerprint_policy="strict", - embedding_fingerprint_legacy_attestation=None, - ), - document_id="document", - is_first_batch=True, - fact_type_override=None, - document_tags=None, - agent_name="test-agent", - log_buffer=[], - start_time=datetime.now(UTC).timestamp(), - all_pre_chunks=list(chunks), - chunk_to_content=[0 for _ in chunks], - chunk_batch_size=1, - operation_id=operation_id, - outbox_callback=callback, - ) - - -class TestMapResultsToContents: - def test_groups_unit_ids_by_content_index(self): - contents = [_make_content("a"), _make_content("b"), _make_content("c")] - processed = [ - _make_processed_fact(0, "a1"), - _make_processed_fact(0, "a2"), - _make_processed_fact(2, "c1"), - ] - unit_ids = ["u-a1", "u-a2", "u-c1"] - - result = _map_results_to_contents(contents, processed, unit_ids) - - assert result == [["u-a1", "u-a2"], [], ["u-c1"]] - - def test_handles_out_of_range_content_index(self): - contents = [_make_content("a"), _make_content("b")] - processed = [ - _make_processed_fact(-1, "f1"), - _make_processed_fact(99, "f2"), - ] - unit_ids = ["u1", "u2"] - - result = _map_results_to_contents(contents, processed, unit_ids) - - assert result == [["u1"], ["u2"]] - - def test_empty_inputs(self): - assert _map_results_to_contents([], [], []) == [] - - def test_length_mismatch_raises(self): - # Regression for #1037: previously the function silently overran unit_ids. - contents = [_make_content("a")] - processed = [_make_processed_fact(0), _make_processed_fact(0)] - unit_ids = ["u1"] # one fewer than processed_facts - - with pytest.raises(ValueError, match="length mismatch"): - _map_results_to_contents(contents, processed, unit_ids) - - def test_unit_ids_assigned_by_processed_fact_position(self): - # Even if processed_facts are interleaved across contents, each unit_id - # must follow its corresponding processed_fact (positional alignment). - contents = [_make_content("a"), _make_content("b")] - processed = [ - _make_processed_fact(1, "b1"), - _make_processed_fact(0, "a1"), - _make_processed_fact(1, "b2"), - ] - unit_ids = ["u-b1", "u-a1", "u-b2"] - - result = _map_results_to_contents(contents, processed, unit_ids) - - assert result == [["u-a1"], ["u-b1", "u-b2"]] - - -class TestEmbeddingsBatchLengthGuarantee: - def test_raises_when_backend_returns_fewer_embeddings(self): - # Regression for #1037: backends that silently truncate must not pass - # through — `zip(extracted_facts, embeddings)` would otherwise drop - # facts and break unit_id alignment downstream. - backend = MagicMock() - backend.encode.return_value = [[0.1, 0.2]] # only 1 vector for 3 inputs - - with pytest.raises(RuntimeError, match="returned 1 vectors for 3 input texts"): - asyncio.run(embedding_utils.generate_embeddings_batch(backend, ["a", "b", "c"])) - - def test_raises_when_backend_returns_more_embeddings(self): - backend = MagicMock() - backend.encode.return_value = [[0.1], [0.2], [0.3]] - - with pytest.raises(RuntimeError, match="returned 3 vectors for 2 input texts"): - asyncio.run(embedding_utils.generate_embeddings_batch(backend, ["a", "b"])) - - def test_passes_through_aligned_embeddings(self): - backend = MagicMock() - backend.encode.return_value = [[0.1], [0.2]] - - result = asyncio.run(embedding_utils.generate_embeddings_batch(backend, ["a", "b"])) - - assert result == [[0.1], [0.2]] - - -class TestStreamingPublicationBoundary: - @pytest.mark.asyncio - async def test_exact_multiple_marks_only_real_final_batch_as_last(self): - queue: asyncio.Queue = asyncio.Queue() - queue.put_nowait(("first",)) - queue.put_nowait(("second",)) - queue.put_nowait(None) - calls: list[tuple[list[tuple], int, bool]] = [] - - async def process_batch(batch: list[tuple], batch_index: int, is_last: bool) -> None: - calls.append((list(batch), batch_index, is_last)) - - await _consume_streaming_batches( - queue, - chunk_batch_size=1, - process_batch=process_batch, - producer_error=[], - pipeline_aborted=[False], - ) - - assert calls == [ - ([("first",)], 0, False), - ([("second",)], 1, True), - ] - assert sum(is_last for _, _, is_last in calls) == 1 - - @pytest.mark.asyncio - async def test_producer_error_suppresses_pending_final_batch(self): - queue: asyncio.Queue = asyncio.Queue() - queue.put_nowait(("pending",)) - queue.put_nowait(None) - calls: list[tuple[list[tuple], int, bool]] = [] - - async def process_batch(batch: list[tuple], batch_index: int, is_last: bool) -> None: - calls.append((list(batch), batch_index, is_last)) - - await _consume_streaming_batches( - queue, - chunk_batch_size=1, - process_batch=process_batch, - producer_error=[RuntimeError("provider failed")], - pipeline_aborted=[False], - ) - - assert calls == [] - - @pytest.mark.asyncio - async def test_takeover_drains_queue_and_raises_without_final_batch(self): - queue: asyncio.Queue = asyncio.Queue() - queue.put_nowait(("first",)) - queue.put_nowait(("discarded",)) - queue.put_nowait(None) - pipeline_aborted = [False] - calls: list[tuple[list[tuple], int, bool]] = [] - - async def process_batch(batch: list[tuple], batch_index: int, is_last: bool) -> None: - calls.append((list(batch), batch_index, is_last)) - pipeline_aborted[0] = True - - with pytest.raises(RetainPublicationAborted, match="ownership was lost"): - await _consume_streaming_batches( - queue, - chunk_batch_size=1, - process_batch=process_batch, - producer_error=[], - pipeline_aborted=pipeline_aborted, - ) - - assert calls == [([("first",)], 0, False)] - assert not any(is_last for _, _, is_last in calls) - assert queue.empty() - - -class TestStreamingPublicationRecovery: - @pytest.mark.asyncio - async def test_zero_fact_final_batch_publishes_once_in_tracking_transaction(self, monkeypatch): - connection = _FakeConnection() - _extract, track_document, update_document = _install_streaming_mocks(monkeypatch, connection) - - async def assert_transactional_callback(callback_connection) -> None: - assert callback_connection is connection - assert connection.in_transaction - - callback = AsyncMock(side_effect=assert_transactional_callback) - unit_ids, _usage, processed_tokens = await _run_streaming_publication( - callback=callback, - contents_dicts=[{"content": "chunk"}], - chunks=["chunk"], - ) - - assert unit_ids == [[]] - assert processed_tokens is None - callback.assert_awaited_once_with(connection) - track_document.assert_awaited_once() - update_document.assert_not_awaited() - - @pytest.mark.asyncio - async def test_all_recovery_chunks_skipped_still_publishes_once(self, monkeypatch): - connection = _FakeConnection() - extract, track_document, update_document = _install_streaming_mocks(monkeypatch, connection) - content = "already committed chunk" - sanitized = orchestrator.fact_extraction._sanitize_text(content) or "" - document_hash = orchestrator.hashlib.sha256(sanitized.encode()).hexdigest() - connection.fetchrow.return_value = {"content_hash": document_hash} - connection.fetchval.return_value = document_hash - monkeypatch.setattr( - orchestrator.chunk_storage, - "load_existing_chunks", - AsyncMock( - return_value=[ - SimpleNamespace(content_hash=orchestrator.chunk_storage.compute_chunk_hash(content)), - ] - ), - ) - - async def assert_transactional_callback(callback_connection) -> None: - assert callback_connection is connection - assert connection.in_transaction - - callback = AsyncMock(side_effect=assert_transactional_callback) - unit_ids, _usage, processed_tokens = await _run_streaming_publication( - callback=callback, - contents_dicts=[{"content": content}], - chunks=[content], - ) - - assert unit_ids == [[]] - assert processed_tokens is None - extract.assert_not_awaited() - callback.assert_awaited_once_with(connection) - track_document.assert_not_awaited() - update_document.assert_awaited_once() - - @pytest.mark.asyncio - async def test_committed_operation_recovery_revalidates_owner_before_publish(self, monkeypatch): - connection = _FakeConnection() - extract, track_document, update_document = _install_streaming_mocks(monkeypatch, connection) - content = "committed document" - sanitized = orchestrator.fact_extraction._sanitize_text(content) or "" - document_hash = orchestrator.hashlib.sha256(sanitized.encode()).hexdigest() - connection.fetchrow.side_effect = [ - {"content_hash": document_hash}, - { - "result_metadata": { - "document_ids": ["document"], - "facts_committed_document_ids": ["document"], - "unit_ids_count": 1, - } - }, - ] - connection.fetchval.return_value = document_hash - connection.fetch.return_value = [{"id": "unit-1"}] - monkeypatch.setattr(orchestrator.chunk_storage, "load_existing_chunks", AsyncMock(return_value=[])) - final_ann = AsyncMock() - monkeypatch.setattr(orchestrator, "_run_final_semantic_ann", final_ann) - - async def assert_transactional_callback(callback_connection) -> None: - assert callback_connection is connection - assert connection.in_transaction - - callback = AsyncMock(side_effect=assert_transactional_callback) - unit_ids, _usage, processed_tokens = await _run_streaming_publication( - callback=callback, - contents_dicts=[{"content": content}], - chunks=[content], - operation_id="00000000-0000-0000-0000-000000000001", - ) - - assert unit_ids == [["unit-1"]] - assert processed_tokens is None - extract.assert_not_awaited() - callback.assert_awaited_once_with(connection) - track_document.assert_not_awaited() - update_document.assert_not_awaited() - final_ann.assert_awaited_once() - - -def _install_stale_retain_mocks(monkeypatch: pytest.MonkeyPatch) -> SimpleNamespace: - connection = AsyncMock() - connection.fetchrow.return_value = {"updated_at": datetime.now(UTC) + timedelta(minutes=1)} - - @asynccontextmanager - async def acquire(_pool): - yield connection - - monkeypatch.setattr(orchestrator, "acquire_with_retry", acquire) - monkeypatch.setattr( - orchestrator.bank_utils, - "get_bank_profile", - AsyncMock(return_value={"name": "test-agent"}), - ) - monkeypatch.setattr(orchestrator, "ensure_bank_embedding_fingerprint", AsyncMock()) - return SimpleNamespace( - embedding_fingerprint_policy="strict", - embedding_fingerprint_legacy_attestation=None, - ) - - -class TestStaleRetainPublication: - @pytest.mark.asyncio - async def test_stale_retain_with_callback_is_not_reported_as_success(self, monkeypatch): - config = _install_stale_retain_mocks(monkeypatch) - callback = AsyncMock() - - with pytest.raises(RetainPublicationAborted, match="superseded before publication"): - await orchestrator.retain_batch( - pool=object(), - embeddings_model=object(), - llm_config=None, - entity_resolver=None, - format_date_fn=None, - bank_id="bank", - contents_dicts=[{"content": "older content"}], - config=config, - document_id="document", - outbox_callback=callback, - ) - - callback.assert_not_awaited() - - @pytest.mark.asyncio - async def test_stale_retain_without_callback_preserves_legacy_noop(self, monkeypatch): - config = _install_stale_retain_mocks(monkeypatch) - - unit_ids, _usage, processed_tokens = await orchestrator.retain_batch( - pool=object(), - embeddings_model=object(), - llm_config=None, - entity_resolver=None, - format_date_fn=None, - bank_id="bank", - contents_dicts=[{"content": "older content"}], - config=config, - document_id="document", - ) - - assert unit_ids == [[]] - assert processed_tokens == 0 - - -class TestDeltaRetainPublicationFallback: - @pytest.mark.parametrize("with_callback", [False, True]) - @pytest.mark.asyncio - async def test_concurrent_hash_change_never_returns_false_success(self, monkeypatch, with_callback): - old_hash = "old-document-hash" - load_connection = AsyncMock() - load_connection.fetchval.return_value = old_hash - - write_connection = AsyncMock() - write_connection.fetchval.return_value = "newer-document-hash" - - @asynccontextmanager - async def transaction(): - yield - - write_connection.transaction = transaction - connections = iter((load_connection, write_connection)) - - @asynccontextmanager - async def acquire(_pool): - yield next(connections) - - same_hash = orchestrator.chunk_storage.compute_chunk_hash("same") - old_chunk_hash = orchestrator.chunk_storage.compute_chunk_hash("old") - existing_chunks = [ - SimpleNamespace(chunk_index=0, content_hash=same_hash, chunk_id="chunk-0"), - SimpleNamespace(chunk_index=1, content_hash=old_chunk_hash, chunk_id="chunk-1"), - ] - monkeypatch.setattr(orchestrator, "acquire_with_retry", acquire) - monkeypatch.setattr( - orchestrator.chunk_storage, - "load_existing_chunks", - AsyncMock(return_value=existing_chunks), - ) - monkeypatch.setattr( - orchestrator, - "_chunk_contents_for_delta", - lambda _contents, _config: {0: "same", 1: "changed"}, - ) - monkeypatch.setattr( - orchestrator, - "_extract_and_embed", - AsyncMock(return_value=([], [], [], orchestrator.TokenUsage())), - ) - monkeypatch.setattr( - orchestrator, - "_pre_resolve_phase1", - AsyncMock(return_value=SimpleNamespace()), - ) - - entity_resolver = MagicMock() - callback = AsyncMock() if with_callback else None - config = SimpleNamespace(write_semantic_links=True) - contents = [_make_content("whole document")] - - async def run_delta(): - return await orchestrator._try_delta_retain( - pool=SimpleNamespace(ops=object()), - embeddings_model=object(), - llm_config=None, - entity_resolver=entity_resolver, - format_date_fn=None, - bank_id="bank", - contents_dicts=[{"content": "whole document"}], - contents=contents, - config=config, - document_id="document", - fact_type_override=None, - document_tags=None, - agent_name="test-agent", - log_buffer=[], - start_time=0.0, - operation_id=None, - schema=None, - outbox_callback=callback, - ) - - if callback is None: - assert await run_delta() is None - else: - with pytest.raises(RetainPublicationAborted, match="changed during retain publication"): - await run_delta() - callback.assert_not_awaited() - assert entity_resolver.discard_pending_stats.call_count == 2 diff --git a/deploy/containers/standalone/Dockerfile b/deploy/containers/standalone/Dockerfile index 5dff65d..9a74a09 100644 --- a/deploy/containers/standalone/Dockerfile +++ b/deploy/containers/standalone/Dockerfile @@ -49,6 +49,8 @@ RUN apt-get update && apt-get install -y \ # Copy dependency files and README (required by pyproject.toml) COPY core/dataplane/pyproject.toml ./api/ COPY core/dataplane/README.md ./api/ +COPY core/dataplane/LICENSE ./api/ +COPY core/dataplane/THIRD_PARTY_NOTICES.md ./api/ WORKDIR /app/api diff --git a/docker-compose.yml b/docker-compose.yml index f24ffab..03e7341 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -45,6 +45,7 @@ services: HMS_API_RETAIN_LLM_MODEL: ${HMS_API_RETAIN_LLM_MODEL:-gpt-4o} HMS_API_RETAIN_LLM_API_KEY: ${HMS_API_RETAIN_LLM_API_KEY:-openai_key_change_me} HMS_API_RETAIN_LLM_BASE_URL: ${HMS_API_RETAIN_LLM_BASE_URL:-} + HMS_API_RETAIN_EMBEDDING_FAILURE_POLICY: ${HMS_API_RETAIN_EMBEDDING_FAILURE_POLICY:-store_without_embedding} HMS_API_EMBEDDINGS_PROVIDER: ${HMS_API_EMBEDDINGS_PROVIDER:-openai} HMS_API_EMBEDDINGS_OPENAI_MODEL: ${HMS_API_EMBEDDINGS_OPENAI_MODEL:-text-embedding-3-small} HMS_API_EMBEDDINGS_OPENAI_API_KEY: ${HMS_API_EMBEDDINGS_OPENAI_API_KEY:-openai_key_change_me} diff --git a/docs/multimodal_memory.md b/docs/multimodal_memory.md index bc9ebdf..5a0b2be 100644 --- a/docs/multimodal_memory.md +++ b/docs/multimodal_memory.md @@ -463,8 +463,8 @@ SHA-256. The raw digest and filename are not exposed in that ID. A byte-for-byte retry inside the same tenant and bank therefore converges on one logical document command, while the same bytes in another tenant or bank do not merge. Use distinct explicit IDs when the same media must intentionally appear as -multiple logical documents. Legacy anonymous non-multimodal file retain keeps -its historical random-ID behavior. +multiple logical documents. Anonymous non-multimodal file retain keeps its +existing random-ID behavior. The document-command identity also includes only upload hints that can change media validation: normalized declared MIME and the recognized final-extension diff --git a/docs/system_architecture_and_multimodal.md b/docs/system_architecture_and_multimodal.md index e6fd2d6..1252e13 100644 --- a/docs/system_architecture_and_multimodal.md +++ b/docs/system_architecture_and_multimodal.md @@ -852,7 +852,7 @@ Oracle 多模态要进入 runtime-supported matrix,至少需要真实 Oracle 2 | HTTP app 与 file endpoint | `core/dataplane/hms_api/api/http.py`:`create_app`、file retain、operation status、version capability | | HTTP/MCP 统一入口 | `core/dataplane/hms_api/api/__init__.py` | | 核心编排 | `core/dataplane/hms_api/engine/memory_engine.py`:`MemoryEngine` | -| Retain | `engine/retain/orchestrator.py`、`fact_extraction.py`、`fact_storage.py`、`link_utils.py` | +| Retain | `engine/ingestion/service.py`, `engine/ingestion/persistence/`, `engine/retain/fact_extraction.py`, `engine/retain/fact_storage.py`, `engine/retain/link_utils.py` | | Recall | `engine/search/retrieval.py`、`fusion.py`、`reranking.py`、`link_expansion_retrieval.py` | | Reflect | `engine/reflect/`、`MemoryEngine.reflect_async` | | Async/worker | `engine/task_backend.py`、`worker/poller.py`、`worker/main.py` | From 70c1f866266586c2b9e74ff2466f9972c5b9f411 Mon Sep 17 00:00:00 2001 From: Dannong Xu Date: Sun, 26 Jul 2026 20:08:51 +0800 Subject: [PATCH 2/8] feat(evaluation): add reproducible LongMemEval pipeline Add a pinned and integrity-checked Retain-to-Judge benchmark workflow with resumable checkpoints, explicit model roles, bounded concurrency, privacy-safe manifests, and public reproduction documentation. Refs #1 --- .aaaSCRIPT/run_benchmark.sh | 201 ++ .env.example | 18 + .gitignore | 1 + README.md | 25 +- lab/evaluation/README.md | 8 +- lab/evaluation/benchmarks/README.md | 15 + lab/evaluation/benchmarks/__init__.py | 1 + lab/evaluation/benchmarks/common/__init__.py | 1 + .../benchmarks/common/benchmark_runner.py | 2485 ++++++++++++++++ .../common/test_benchmark_runner.py | 374 +++ .../benchmarks/longmemeval/README.md | 218 ++ .../benchmarks/longmemeval/__init__.py | 1 + .../longmemeval/evidence_bundles.py | 316 ++ .../longmemeval/longmemeval.env.example | 51 + .../longmemeval/longmemeval_benchmark.py | 2639 +++++++++++++++++ .../benchmarks/longmemeval/source_backfill.py | 504 ++++ .../longmemeval/test_evidence_bundles.py | 178 ++ .../longmemeval/test_release_integrity.py | 308 ++ .../longmemeval/test_source_backfill.py | 253 ++ .../test_source_context_integration.py | 161 + lab/evaluation/pyproject.toml | 3 +- uv.lock | 2 + 22 files changed, 7760 insertions(+), 3 deletions(-) create mode 100755 .aaaSCRIPT/run_benchmark.sh create mode 100644 lab/evaluation/benchmarks/README.md create mode 100644 lab/evaluation/benchmarks/__init__.py create mode 100644 lab/evaluation/benchmarks/common/__init__.py create mode 100644 lab/evaluation/benchmarks/common/benchmark_runner.py create mode 100644 lab/evaluation/benchmarks/common/test_benchmark_runner.py create mode 100644 lab/evaluation/benchmarks/longmemeval/README.md create mode 100644 lab/evaluation/benchmarks/longmemeval/__init__.py create mode 100644 lab/evaluation/benchmarks/longmemeval/evidence_bundles.py create mode 100644 lab/evaluation/benchmarks/longmemeval/longmemeval.env.example create mode 100644 lab/evaluation/benchmarks/longmemeval/longmemeval_benchmark.py create mode 100644 lab/evaluation/benchmarks/longmemeval/source_backfill.py create mode 100644 lab/evaluation/benchmarks/longmemeval/test_evidence_bundles.py create mode 100644 lab/evaluation/benchmarks/longmemeval/test_release_integrity.py create mode 100644 lab/evaluation/benchmarks/longmemeval/test_source_backfill.py create mode 100644 lab/evaluation/benchmarks/longmemeval/test_source_context_integration.py diff --git a/.aaaSCRIPT/run_benchmark.sh b/.aaaSCRIPT/run_benchmark.sh new file mode 100755 index 0000000..fd51c8d --- /dev/null +++ b/.aaaSCRIPT/run_benchmark.sh @@ -0,0 +1,201 @@ +#!/usr/bin/env bash +set -euo pipefail + +ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +ENV_FILE_EXPLICIT=0 +if [ "${HMS_ENV_FILE+x}" = "x" ]; then + ENV_FILE="$HMS_ENV_FILE" + ENV_FILE_EXPLICIT=1 +else + ENV_FILE="$ROOT_DIR/.env" +fi + +if [ -f "$ENV_FILE" ]; then + set -a + # shellcheck disable=SC1091 + source "$ENV_FILE" + set +a +elif [ "$ENV_FILE_EXPLICIT" = "1" ]; then + echo "HMS_ENV_FILE does not exist: $ENV_FILE" >&2 + exit 2 +fi + +if [ "${HMS_BENCHMARK:-longmemeval}" != "longmemeval" ]; then + echo "This launcher supports only HMS_BENCHMARK=longmemeval." >&2 + exit 2 +fi + +if [ -n "${HMS_BENCHMARK_DATABASE_URL:-}" ]; then + export HMS_API_DATABASE_URL="$HMS_BENCHMARK_DATABASE_URL" +fi + +if [ -z "${HMS_API_DATABASE_URL:-}" ]; then + echo "HMS_API_DATABASE_URL is required. See lab/evaluation/benchmarks/longmemeval/README.md." >&2 + exit 2 +fi + +case "$HMS_API_DATABASE_URL" in + *@postgres:*) + echo "HMS_API_DATABASE_URL points to the Compose-only host 'postgres'." >&2 + echo "Use HMS_BENCHMARK_DATABASE_URL with a host-reachable address such as 127.0.0.1." >&2 + exit 2 + ;; +esac + +DATA_DIR="${HMS_DATA_DIR:-$ROOT_DIR/.aaaDATA}" +LOG_DIR="${HMS_LOG_DIR:-$ROOT_DIR/.aaaLOG}" +RESULT_DIR="${HMS_RESULT_DIR:-$ROOT_DIR/.aaaRESULT}" +mkdir -p "$DATA_DIR" "$LOG_DIR" "$RESULT_DIR" + +TIMESTAMP="$(date -u +%Y%m%d_%H%M%S)" +LOG_FILE="${HMS_BENCHMARK_LOG:-$LOG_DIR/longmemeval_${TIMESTAMP}.log}" + +RESUME_REQUESTED="${HMS_RESUME:-0}" +CLI_RESULTS_FILENAME="" +EXPECT_RESULTS_FILENAME=0 +for argument in "$@"; do + if [ "$EXPECT_RESULTS_FILENAME" = "1" ]; then + CLI_RESULTS_FILENAME="$argument" + EXPECT_RESULTS_FILENAME=0 + continue + fi + if [ "$argument" = "--resume" ]; then + RESUME_REQUESTED=1 + fi + case "$argument" in + --results-filename) + EXPECT_RESULTS_FILENAME=1 + ;; + --results-filename=*) + CLI_RESULTS_FILENAME="${argument#--results-filename=}" + ;; + esac +done + +if [ "$RESUME_REQUESTED" = "1" ] && [ -z "${HMS_RESULTS_FILENAME:-}" ] && [ -z "$CLI_RESULTS_FILENAME" ]; then + echo "Resume requires HMS_RESULTS_FILENAME or --results-filename to name the existing result artifact." >&2 + exit 2 +fi + +RESULTS_FILENAME="${CLI_RESULTS_FILENAME:-${HMS_RESULTS_FILENAME:-longmemeval_${TIMESTAMP}.json}}" +CONTEXT_FORMAT="${HMS_CONTEXT_FORMAT:-structured_source}" +PARALLEL="${HMS_PARALLEL:-1}" +MAX_CONCURRENT_QUESTIONS="${HMS_MAX_CONCURRENT_QUESTIONS:-$PARALLEL}" +EVAL_SEMAPHORE_SIZE="${HMS_EVAL_SEMAPHORE_SIZE:-$PARALLEL}" +THINKING_BUDGET="${HMS_THINKING_BUDGET:-500}" +MAX_TOKENS="${HMS_MAX_TOKENS:-8192}" + +require_positive_integer() { + local name="$1" + local value="$2" + case "$value" in + ""|*[!0-9]*|0) + echo "$name must be a positive integer, got: $value" >&2 + exit 2 + ;; + esac +} + +require_positive_integer HMS_PARALLEL "$PARALLEL" +require_positive_integer HMS_MAX_CONCURRENT_QUESTIONS "$MAX_CONCURRENT_QUESTIONS" +require_positive_integer HMS_EVAL_SEMAPHORE_SIZE "$EVAL_SEMAPHORE_SIZE" +require_positive_integer HMS_THINKING_BUDGET "$THINKING_BUDGET" +require_positive_integer HMS_MAX_TOKENS "$MAX_TOKENS" +if [ -n "${HMS_MAX_INSTANCES:-}" ]; then + require_positive_integer HMS_MAX_INSTANCES "$HMS_MAX_INSTANCES" +fi +if [ -n "${HMS_MAX_QUESTIONS:-}" ]; then + require_positive_integer HMS_MAX_QUESTIONS "$HMS_MAX_QUESTIONS" +fi + +LONGMEMEVAL_ARGS=( + --results-dir "$RESULT_DIR" + --results-filename "$RESULTS_FILENAME" + --context-format "$CONTEXT_FORMAT" + --parallel "$PARALLEL" + --max-concurrent-questions "$MAX_CONCURRENT_QUESTIONS" + --eval-semaphore-size "$EVAL_SEMAPHORE_SIZE" + --thinking-budget "$THINKING_BUDGET" + --max-tokens "$MAX_TOKENS" +) + +if [ -n "${HMS_MAX_INSTANCES:-}" ]; then + LONGMEMEVAL_ARGS+=(--max-instances "$HMS_MAX_INSTANCES") +fi + +if [ -n "${HMS_MAX_QUESTIONS:-}" ]; then + LONGMEMEVAL_ARGS+=(--max-questions "$HMS_MAX_QUESTIONS") +fi + +if [ -n "${HMS_DATASET_PATH:-}" ]; then + LONGMEMEVAL_ARGS+=(--dataset-path "$HMS_DATASET_PATH") +fi + +if [ "${HMS_RETRIEVAL_ONLY:-0}" = "1" ]; then + LONGMEMEVAL_ARGS+=(--skip-ingestion) +fi + +if [ "${HMS_ENABLE_QUERY_EXPANSION:-0}" = "1" ]; then + LONGMEMEVAL_ARGS+=(--enable-query-expansion) + LONGMEMEVAL_ARGS+=(--query-rewriting-strategy "${HMS_QUERY_REWRITING_STRATEGY:-llm_driven}") +fi + +if [ -n "${HMS_SESSION_EXPANSION_WEIGHT:-}" ]; then + LONGMEMEVAL_ARGS+=(--session-expansion-weight "$HMS_SESSION_EXPANSION_WEIGHT") +fi + +case "${HMS_PIPELINE:-ledger}" in + ledger) + LONGMEMEVAL_ARGS+=(--oracle-planner-v26) + ;; + self_evolution) + LONGMEMEVAL_ARGS+=(--oracle-planner-v220) + ;; + standard) + ;; + *) + echo "Unsupported HMS_PIPELINE: ${HMS_PIPELINE:-}" >&2 + echo "Supported values: standard, ledger, self_evolution" >&2 + exit 2 + ;; +esac + +if [ "$RESUME_REQUESTED" = "1" ] && [ "${HMS_RESUME:-0}" = "1" ]; then + LONGMEMEVAL_ARGS+=(--resume) +fi + +PYTHON_BIN="${HMS_PYTHON_BIN:-}" +if [ -n "$PYTHON_BIN" ]; then + if [ ! -x "$PYTHON_BIN" ]; then + echo "HMS_PYTHON_BIN is not executable: $PYTHON_BIN" >&2 + exit 2 + fi + export PYTHONPATH="$ROOT_DIR/core/dataplane:$ROOT_DIR/lab/evaluation${PYTHONPATH:+:$PYTHONPATH}" + CMD=("$PYTHON_BIN" -m benchmarks.longmemeval.longmemeval_benchmark "${LONGMEMEVAL_ARGS[@]}" "$@") +else + if ! command -v uv >/dev/null 2>&1; then + echo "The 'uv' command is required. Install uv or set HMS_PYTHON_BIN." >&2 + exit 2 + fi + CMD=( + uv run + --project "$ROOT_DIR/lab/evaluation" + python -m benchmarks.longmemeval.longmemeval_benchmark + "${LONGMEMEVAL_ARGS[@]}" + "$@" + ) +fi + +echo "LongMemEval pipeline: Retain -> Recall -> Answer -> Judge" +echo "Configuration source: $ENV_FILE" +echo "Pipeline profile: ${HMS_PIPELINE:-ledger}" +echo "Context format: $CONTEXT_FORMAT" +echo "Parallel items: $PARALLEL" +echo "Results: $RESULT_DIR/$RESULTS_FILENAME" +echo "Log: $LOG_FILE" + +{ + echo "[$(date -u '+%Y-%m-%dT%H:%M:%SZ')] Running: ${CMD[*]}" + cd "$ROOT_DIR" + "${CMD[@]}" +} 2>&1 | tee "$LOG_FILE" diff --git a/.env.example b/.env.example index 2100d09..20ba9da 100644 --- a/.env.example +++ b/.env.example @@ -105,6 +105,24 @@ HMS_API_RETAIN_LLM_MODEL=gpt-4o HMS_API_RETAIN_LLM_API_KEY=openai_key_change_me HMS_API_RETAIN_LLM_BASE_URL=https://api.openai.com/v1 +# Store facts without embeddings by default when an embedding batch fails. +# Set to "raise" to reject the retain request instead. +HMS_API_RETAIN_EMBEDDING_FAILURE_POLICY=store_without_embedding + +# Answer model used after recall to generate the benchmark response. If omitted, +# the benchmark falls back to HMS_API_LLM_*. +HMS_API_ANSWER_LLM_PROVIDER=openai +HMS_API_ANSWER_LLM_MODEL=gpt-4o +HMS_API_ANSWER_LLM_API_KEY=openai_key_change_me +HMS_API_ANSWER_LLM_BASE_URL=https://api.openai.com/v1 + +# Judge model used to compare the generated response with the gold answer. If +# omitted, the benchmark falls back to HMS_API_LLM_*. +HMS_API_JUDGE_LLM_PROVIDER=openai +HMS_API_JUDGE_LLM_MODEL=gpt-4o +HMS_API_JUDGE_LLM_API_KEY=openai_key_change_me +HMS_API_JUDGE_LLM_BASE_URL=https://api.openai.com/v1 + # Embeddings and lightweight RRF reranking. HMS_API_EMBEDDINGS_PROVIDER=openai HMS_API_EMBEDDINGS_OPENAI_MODEL=text-embedding-3-small diff --git a/.gitignore b/.gitignore index 551d5cf..6c24c63 100644 --- a/.gitignore +++ b/.gitignore @@ -9,6 +9,7 @@ __pycache__/ # Local generated data, dumps, caches, and reports artifacts*/ .tmp/ +/tmp/ .hf_cache/ torchinductor_root/ README_*_VERSION.md diff --git a/README.md b/README.md index dc485fb..0582408 100644 --- a/README.md +++ b/README.md @@ -178,6 +178,27 @@ export HMS_API_MILVUS_URI=./hms_milvus.db # Milvus Lite After enabling Milvus for an existing database, rebuild its projection with `hms-admin rebuild-vector-index --yes`. Milvus Lite is intended for a single HMS process; use Milvus Server or Zilliz Cloud for multi-worker deployments. See [the dataplane README](core/dataplane/README.md#optional-milvus-semantic-index) for all settings and consistency guidance. +## LongMemEval Reproduction + +The code-only LongMemEval adapter runs the complete +`Retain -> Recall -> Answer -> Judge` workflow: + +```bash +cp lab/evaluation/benchmarks/longmemeval/longmemeval.env.example .env.longmemeval +chmod 600 .env.longmemeval +# Fill in the database and model credentials, then: +HMS_ENV_FILE=.env.longmemeval \ +HMS_MAX_INSTANCES=1 \ +HMS_RESULTS_FILENAME=longmemeval-smoke.json \ +bash .aaaSCRIPT/run_benchmark.sh +``` + +The runner downloads and verifies a pinned dataset revision. Datasets, +credentials, retained banks, logs, and generated results are not included in +the repository. See the +[LongMemEval reproduction guide](lab/evaluation/benchmarks/longmemeval/README.md) +for database setup, concurrency, resume behavior, and full-run validation. + ## Security Notes - Keep `.env`, private keys, tokens, and populated credentials out of Git. @@ -187,4 +208,6 @@ After enabling Milvus for an existing database, rebuild its projection with `hms ## License -See [LICENSE](LICENSE). +See the repository [MIT License](LICENSE), any component-specific package +metadata, and [THIRD_PARTY_NOTICES.md](THIRD_PARTY_NOTICES.md) for notices +covering included third-party code. diff --git a/lab/evaluation/README.md b/lab/evaluation/README.md index db84693..77d566e 100644 --- a/lab/evaluation/README.md +++ b/lab/evaluation/README.md @@ -1 +1,7 @@ -# Memory Dev/Utils scripts \ No newline at end of file +# Evaluation utilities + +This package contains development utilities, compatibility checks, and public +benchmark adapters. + +See the [LongMemEval reproduction guide](benchmarks/longmemeval/README.md) for +the complete `Retain -> Recall -> Answer -> Judge` workflow. diff --git a/lab/evaluation/benchmarks/README.md b/lab/evaluation/benchmarks/README.md new file mode 100644 index 0000000..bc35d16 --- /dev/null +++ b/lab/evaluation/benchmarks/README.md @@ -0,0 +1,15 @@ +# Evaluation adapters + +This directory contains source code for benchmark ingestion, recall, answer +generation, and judging. The repository does not include datasets, retained +memory banks, generated answers, logs, or result artifacts. + +The public benchmark adapter currently supports LongMemEval: + +- [LongMemEval reproduction guide](longmemeval/README.md) +- Entry point: `bash .aaaSCRIPT/run_benchmark.sh` +- Python module: `python -m benchmarks.longmemeval.longmemeval_benchmark` + +The launcher performs the complete `Retain -> Recall -> Answer -> Judge` +pipeline by default. Set `HMS_RETRIEVAL_ONLY=1` only when the configured +database already contains every required memory bank. diff --git a/lab/evaluation/benchmarks/__init__.py b/lab/evaluation/benchmarks/__init__.py new file mode 100644 index 0000000..0d05362 --- /dev/null +++ b/lab/evaluation/benchmarks/__init__.py @@ -0,0 +1 @@ +"""Benchmarks package for memory system evaluation.""" diff --git a/lab/evaluation/benchmarks/common/__init__.py b/lab/evaluation/benchmarks/common/__init__.py new file mode 100644 index 0000000..784540e --- /dev/null +++ b/lab/evaluation/benchmarks/common/__init__.py @@ -0,0 +1 @@ +"""Common benchmark framework.""" diff --git a/lab/evaluation/benchmarks/common/benchmark_runner.py b/lab/evaluation/benchmarks/common/benchmark_runner.py new file mode 100644 index 0000000..79aebf2 --- /dev/null +++ b/lab/evaluation/benchmarks/common/benchmark_runner.py @@ -0,0 +1,2485 @@ +""" +Common benchmark runner framework. + +This module provides a unified interface for running memory benchmarks with: +- Batch ingestion for speed +- Parallel question processing with semaphores +- Parallel LLM judging with rate limiting +- Progress tracking with Rich +- Comprehensive metrics collection +- Support for both traditional (search + LLM) and integrated (think API) approaches + +The framework supports two answer generation patterns: +1. Traditional: Benchmark runner performs search, then passes results to answer generator +2. Integrated: Answer generator performs its own retrieval (e.g., think API) + - Indicated by needs_external_search() returning False + - Skips the search step for efficiency +""" + +import asyncio +import hashlib +import json +import logging +import os +import re +from abc import ABC, abstractmethod +from dataclasses import asdict, dataclass, field +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Callable, Dict, Iterable, List, Mapping, Optional, Tuple, Union + +import pydantic +from hms_api import MemoryEngine +from hms_api.config import DEFAULT_LLM_MODEL, PROVIDER_DEFAULT_MODELS, get_config + +# Configure logging from environment variable +get_config().configure_logging() +from hms_api.engine.memory_engine import Budget +from hms_api.engine.schema import fq_table +from hms_api.models import RequestContext +from openai import AsyncOpenAI +from rich import box +from rich.console import Console +from rich.progress import BarColumn, Progress, SpinnerColumn, TextColumn +from rich.table import Table + +console = Console() + + +def _endpoint_fingerprint(value: Optional[str]) -> str: + """Return a non-secret identity for an endpoint used in result compatibility checks.""" + + normalized = (value or "").strip().rstrip("/") + digest = hashlib.sha256(normalized.encode("utf-8")).hexdigest() + return f"sha256:{digest[:16]}" + + +def _embedding_runtime_config() -> Dict[str, str]: + """Describe the effective embedding backend without exposing credentials.""" + + provider = os.getenv("HMS_API_EMBEDDINGS_PROVIDER", "local").lower() + model_settings = { + "local": ("HMS_API_EMBEDDINGS_LOCAL_MODEL", "BAAI/bge-small-en-v1.5"), + "tei": (None, ""), + "openai": ("HMS_API_EMBEDDINGS_OPENAI_MODEL", "text-embedding-3-small"), + "openrouter": ("HMS_API_EMBEDDINGS_OPENROUTER_MODEL", "perplexity/pplx-embed-v1-0.6b"), + "cohere": ("HMS_API_EMBEDDINGS_COHERE_MODEL", "embed-english-v3.0"), + "litellm": ("HMS_API_EMBEDDINGS_LITELLM_MODEL", "text-embedding-3-small"), + "litellm-sdk": ("HMS_API_EMBEDDINGS_LITELLM_SDK_MODEL", "cohere/embed-english-v3.0"), + "google": ("HMS_API_EMBEDDINGS_GEMINI_MODEL", "gemini-embedding-001"), + } + endpoint_settings = { + "tei": ("HMS_API_EMBEDDINGS_TEI_URL", None), + "openai": ("HMS_API_EMBEDDINGS_OPENAI_BASE_URL", None), + "openrouter": (None, "https://openrouter.ai/api/v1"), + "cohere": ("HMS_API_EMBEDDINGS_COHERE_BASE_URL", None), + "litellm": ("HMS_API_EMBEDDINGS_LITELLM_API_BASE", "http://localhost:4000"), + "litellm-sdk": ("HMS_API_EMBEDDINGS_LITELLM_SDK_API_BASE", None), + } + + model_env, model_default = model_settings.get(provider, (None, "")) + model = os.getenv(model_env, model_default) if model_env else model_default + endpoint_env, endpoint_default = endpoint_settings.get(provider, (None, None)) + endpoint = os.getenv(endpoint_env, endpoint_default) if endpoint_env else endpoint_default + if provider == "google": + endpoint = "|".join( + ( + "google", + os.getenv("HMS_API_EMBEDDINGS_VERTEXAI_PROJECT_ID", ""), + os.getenv("HMS_API_EMBEDDINGS_VERTEXAI_REGION", "us-central1"), + ) + ) + + return { + "provider": provider, + "model": model, + "fingerprint_policy": os.getenv("HMS_API_EMBEDDING_FINGERPRINT_POLICY", "strict"), + "endpoint_fingerprint": _endpoint_fingerprint(endpoint), + } + + +def _reranker_runtime_config() -> Dict[str, str]: + """Describe the effective reranker backend without exposing credentials.""" + + provider = os.getenv("HMS_API_RERANKER_PROVIDER", "local").lower() + model_settings = { + "local": ("HMS_API_RERANKER_LOCAL_MODEL", "cross-encoder/ms-marco-MiniLM-L-6-v2"), + "tei": (None, ""), + "cohere": ("HMS_API_RERANKER_COHERE_MODEL", "rerank-english-v3.0"), + "openrouter": ("HMS_API_RERANKER_OPENROUTER_MODEL", "cohere/rerank-v3.5"), + "flashrank": ("HMS_API_RERANKER_FLASHRANK_MODEL", "ms-marco-MiniLM-L-12-v2"), + "litellm": ("HMS_API_RERANKER_LITELLM_MODEL", "cohere/rerank-english-v3.0"), + "litellm-sdk": ("HMS_API_RERANKER_LITELLM_SDK_MODEL", "cohere/rerank-english-v3.0"), + "zeroentropy": ("HMS_API_RERANKER_ZEROENTROPY_MODEL", "zerank-2"), + "siliconflow": ("HMS_API_RERANKER_SILICONFLOW_MODEL", "BAAI/bge-reranker-v2-m3"), + "google": ("HMS_API_RERANKER_GOOGLE_MODEL", "semantic-ranker-default-004"), + "qwen3-reranker": ("HMS_API_RERANKER_QWEN3_MODEL_PATH", ""), + "rrf": (None, ""), + "jina-mlx": (None, ""), + } + endpoint_settings = { + "tei": ("HMS_API_RERANKER_TEI_URL", None), + "cohere": ("HMS_API_RERANKER_COHERE_BASE_URL", None), + "openrouter": (None, "https://openrouter.ai/api/v1/rerank"), + "litellm": ("HMS_API_RERANKER_LITELLM_API_BASE", "http://localhost:4000"), + "litellm-sdk": ("HMS_API_RERANKER_LITELLM_SDK_API_BASE", None), + "siliconflow": ("HMS_API_RERANKER_SILICONFLOW_BASE_URL", "https://api.siliconflow.cn/v1"), + } + + model_env, model_default = model_settings.get(provider, (None, "")) + model = os.getenv(model_env, model_default) if model_env else model_default + endpoint_env, endpoint_default = endpoint_settings.get(provider, (None, None)) + endpoint = os.getenv(endpoint_env, endpoint_default) if endpoint_env else endpoint_default + if provider == "google": + endpoint = f"google|{os.getenv('HMS_API_RERANKER_GOOGLE_PROJECT_ID', '')}" + + return { + "provider": provider, + "model": model, + "endpoint_fingerprint": _endpoint_fingerprint(endpoint), + } + + +def _default_model_for_provider(provider: str) -> str: + """Mirror the core configuration's provider-specific model defaults.""" + + return PROVIDER_DEFAULT_MODELS.get(provider.lower(), DEFAULT_LLM_MODEL) + + +class IngestionIntegrityError(RuntimeError): + """Raised when retained source documents are not durably queryable.""" + + def __init__(self, report: Dict[str, Any]): + self.report = report + item_id = report.get("item_id", "unknown") + missing = report.get("missing_documents", []) + empty = report.get("documents_without_chunks", []) + super().__init__( + f"Durable ingestion audit failed for item {item_id!r}: " + f"missing_documents={missing}, documents_without_chunks={empty}" + ) + + +def _write_json_atomic(payload: Dict[str, Any], output_path: Path) -> None: + """Write a JSON result without exposing a partially written resume file.""" + + output_path.parent.mkdir(parents=True, exist_ok=True) + temporary_path = output_path.with_name(f".{output_path.name}.tmp") + try: + with temporary_path.open("w", encoding="utf-8") as handle: + json.dump(payload, handle, ensure_ascii=False, indent=2, default=str) + handle.write("\n") + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary_path, output_path) + finally: + temporary_path.unlink(missing_ok=True) + + +def _result_is_resume_complete(result: Mapping[str, Any]) -> bool: + """Return whether an existing item can be safely skipped by resume mode.""" + + metrics = result.get("metrics") + if not isinstance(metrics, Mapping): + return False + if int(metrics.get("total", 0) or 0) < 1 or int(metrics.get("invalid", 0) or 0) > 0: + return False + + details = metrics.get("detailed_results", []) + if not isinstance(details, list) or not details: + return False + for detail in details: + if not isinstance(detail, Mapping): + return False + if detail.get("is_invalid") or detail.get("error"): + return False + predicted = str(detail.get("predicted_answer", "")) + reasoning = str(detail.get("correctness_reasoning", "")) + if predicted.startswith("Error generating answer:") or reasoning.startswith("Error:"): + return False + return True + + +def get_model_config() -> Dict[str, Dict[str, str]]: + """ + Get the non-secret model configuration for the benchmark runtime. + + Reads directly from environment variables without instantiating LLM clients. + + Returns: + Provider and model identifiers for each runtime role. + """ + # Memory/HMS config (base config) + memory_provider = os.getenv("HMS_API_LLM_PROVIDER", "groq") + memory_model = os.getenv("HMS_API_LLM_MODEL", "openai/gpt-oss-120b") + + # Retain config (falls back to memory config). + retain_provider = os.getenv("HMS_API_RETAIN_LLM_PROVIDER", memory_provider) + retain_model = os.getenv("HMS_API_RETAIN_LLM_MODEL") + if retain_model is None: + retain_model = ( + _default_model_for_provider(retain_provider) + if "HMS_API_RETAIN_LLM_PROVIDER" in os.environ + else memory_model + ) + memory_base_url = os.getenv("HMS_API_LLM_BASE_URL") or None + retain_base_url = os.getenv("HMS_API_RETAIN_LLM_BASE_URL") or memory_base_url + + # Answer generation config (falls back to memory config) + answer_provider = os.getenv("HMS_API_ANSWER_LLM_PROVIDER", memory_provider) + answer_model = os.getenv("HMS_API_ANSWER_LLM_MODEL", memory_model) + answer_base_url = os.getenv("HMS_API_ANSWER_LLM_BASE_URL") or memory_base_url + + # Judge config (falls back to memory config) + judge_provider = os.getenv("HMS_API_JUDGE_LLM_PROVIDER", memory_provider) + judge_model = os.getenv("HMS_API_JUDGE_LLM_MODEL", memory_model) + judge_base_url = os.getenv("HMS_API_JUDGE_LLM_BASE_URL") or memory_base_url + + return { + "hms": { + "provider": memory_provider, + "model": memory_model, + "endpoint_fingerprint": _endpoint_fingerprint(memory_base_url), + }, + "retain": { + "provider": retain_provider, + "model": retain_model, + "endpoint_fingerprint": _endpoint_fingerprint(retain_base_url), + }, + "answer_generation": { + "provider": answer_provider, + "model": answer_model, + "endpoint_fingerprint": _endpoint_fingerprint(answer_base_url), + }, + "judge": { + "provider": judge_provider, + "model": judge_model, + "endpoint_fingerprint": _endpoint_fingerprint(judge_base_url), + }, + "embeddings": _embedding_runtime_config(), + "reranker": _reranker_runtime_config(), + } + + +def print_model_config(): + """Print the model configuration to console.""" + config = get_model_config() + + console.print("\n[bold cyan]Model Configuration:[/bold cyan]") + console.print(f" HMS: {config['hms']['provider']}/{config['hms']['model']}") + console.print(f" Retain: {config['retain']['provider']}/{config['retain']['model']}") + console.print( + f" Answer Generation: {config['answer_generation']['provider']}/{config['answer_generation']['model']}" + ) + console.print(f" LLM Judge: {config['judge']['provider']}/{config['judge']['model']}") + console.print() + + +async def create_memory_engine() -> MemoryEngine: + """ + Create and initialize a MemoryEngine instance from environment variables. + + Reads configuration from: + - HMS_API_DATABASE_URL (default: "pg0") + - HMS_API_LLM_PROVIDER (default: "groq") + - HMS_API_LLM_API_KEY + - HMS_API_LLM_MODEL (default: "openai/gpt-oss-120b") + - HMS_API_LLM_BASE_URL (optional) + + Returns: + Initialized MemoryEngine instance + """ + memory = MemoryEngine( + db_url=os.getenv("HMS_API_DATABASE_URL", "pg0"), + memory_llm_provider=os.getenv("HMS_API_LLM_PROVIDER", "groq"), + memory_llm_api_key=os.getenv("HMS_API_LLM_API_KEY"), + memory_llm_model=os.getenv("HMS_API_LLM_MODEL", "openai/gpt-oss-120b"), + memory_llm_base_url=os.getenv("HMS_API_LLM_BASE_URL") or None, # Use None to get provider defaults + ) + await memory.initialize() + return memory + + +class BenchmarkDataset(ABC): + """Abstract base class for benchmark datasets.""" + + @abstractmethod + def load(self, path: Path, max_items: Optional[int] = None) -> List[Dict[str, Any]]: + """ + Load dataset from file. + + Returns: + List of dataset items + """ + pass + + @abstractmethod + def get_item_id(self, item: Dict) -> str: + """Get unique identifier for an item.""" + pass + + @abstractmethod + def prepare_sessions_for_ingestion(self, item: Dict) -> List[Dict[str, Any]]: + """ + Prepare conversation sessions for batch ingestion. + + Returns: + List of session dicts with keys: 'content', 'context', 'event_date' + """ + pass + + @abstractmethod + def get_qa_pairs(self, item: Dict) -> List[Dict[str, Any]]: + """ + Extract QA pairs from an item. + + Returns: + List of QA dicts with keys: 'question', 'answer', 'category' (optional) + """ + pass + + +class LLMAnswerGenerator(ABC): + """Abstract base class for LLM-based answer generation.""" + + def needs_external_search(self) -> bool: + """ + Whether this generator needs external search to be performed. + + Returns: + True if the benchmark runner should perform search before calling generate_answer. + False if the generator does its own retrieval (e.g., integrated think API). + """ + return True + + @abstractmethod + async def generate_answer( + self, + question: str, + recall_result: Dict[str, Any], + question_date: Optional[datetime] = None, + question_type: Optional[str] = None, + bank_id: Optional[str] = None, + ) -> Tuple[str, str, Optional[List[Dict[str, Any]]]]: + """ + Generate answer from retrieved memories. + + Args: + question: The question text + recall_result: Full RecallResult dict containing results, entities, chunks, and trace + question_date: Optional date when the question was asked (for temporal context) + question_type: Optional question category/type (e.g., 'multi-session', 'temporal-reasoning') + bank_id: Optional bank ID for generators that need it (e.g., ReflectAnswerGenerator) + + Returns: + Tuple of (answer, reasoning, retrieved_memories_override) + - answer: The generated answer text + - reasoning: Explanation of how the answer was derived + - retrieved_memories_override: Optional list of memories to include in results + - None: Use memories from recall_result (traditional mode) + - List: Use these memories instead (integrated mode like think API) + """ + pass + + +class JudgeResponse(pydantic.BaseModel): + """Judge response format.""" + + correct: bool + reasoning: str + + +class LLMAnswerEvaluator: + """LLM-based answer evaluator with configurable provider.""" + + def __init__(self): + """Initialize with LLM configuration for judge/evaluator.""" + import os + + from hms_api.engine.llm_wrapper import LLMConfig + + self.llm_config = LLMConfig( + provider=os.getenv("HMS_API_JUDGE_LLM_PROVIDER", os.getenv("HMS_API_LLM_PROVIDER", "openai")), + api_key=os.getenv("HMS_API_JUDGE_LLM_API_KEY", os.getenv("HMS_API_LLM_API_KEY", "")), + base_url=os.getenv("HMS_API_JUDGE_LLM_BASE_URL", os.getenv("HMS_API_LLM_BASE_URL", "")), + model=os.getenv("HMS_API_JUDGE_LLM_MODEL", os.getenv("HMS_API_LLM_MODEL", "gpt-4o-mini")), + reasoning_effort="high", + ) + self.client = self.llm_config._client + self.model = self.llm_config.model + + async def judge_answer( + self, + question: str, + correct_answer: str, + predicted_answer: str, + semaphore: asyncio.Semaphore, + category: Optional[str] = None, + max_retries: int = 3, + ) -> Tuple[bool, str, float]: + """ + Evaluate predicted answer using LLM-as-judge with category-specific prompts. + + Args: + question: The question + correct_answer: Gold/correct answer + predicted_answer: Predicted answer + semaphore: Semaphore for rate limiting + category: Question category for LongMemEval-specific evaluation + max_retries: Maximum retry attempts for validation errors + + Returns: + Tuple of (is_correct, reasoning, judge_time) + """ + async with semaphore: + import time + + judge_start_time = time.time() + + for attempt in range(max_retries): + try: + # LongMemEval-specific evaluation prompts + if category in ["single-session-user", "single-session-assistant", "multi-session"]: + prompt_content = f"""Evaluate if the model response contains the correct answer to the question. + +I will give you a question, a correct answer, and a response from a model. +Please set correct=true if the response contains the correct answer. Otherwise, set correct=no. +If the response is equivalent to the correct answer or contains all the intermediate steps to get the correct answer, you should also set correct=true. +If the response only contains a subset of the information required by the answer, set correct=false + +Question: {question} + +Correct Answer: {correct_answer} + +Model Response: {predicted_answer} + +Evaluation criteria: +- Set correct=true if the response contains the correct answer +- Set correct=true if the response is equivalent to the correct answer or contains intermediate steps +- Set correct=false if the response is incorrect or missing key information + +Provide your evaluation as JSON with: +- reasoning: One sentence explanation +- correct: true or false""" + + elif category == "temporal-reasoning": + prompt_content = """ +I will give you a question, a correct answer, and a response from a model. +Please set correct=true if the response contains the correct answer. Otherwise, set correct=false. +If the response is equivalent to the correct answer or contains all the intermediate steps to get the correct answer, you should also set correct=true. +If the response only contains a subset of the information required by the answer, answer correct=false. +In addition, do not penalize off-by-one errors for the number of days. If the question asks for the number of days/weeks/months, etc., and the model makes off-by-one errors (e.g., predicting 19 days when the answer is 18), the model's response is still correct. +""" + + elif category == "knowledge-update": + prompt_content = """ +I will give you a question, a correct answer, and a response from a model. +Please set correct=true if the response contains the correct answer. Otherwise, set correct=false. +If the response contains some previous information along with an updated answer, the response should be considered as correct as long as the updated answer is the required answer. +""" + + elif category == "single-session-preference": + prompt_content = """ +I will give you a question, a answer for desired personalized response, and a response from a model. +Please set correct=true if the response satisfies the desired response. Otherwise, set correct=false. +The model does not need to reflect all the points in the desired response. The response is correct as long as it recalls and utilizes the user's personal information correctly. +""" + + else: + # Default short-form answer evaluation. + prompt_content = """Your task is to label an answer to a question as 'CORRECT' or 'WRONG'. You will be given the following data: + (1) a question (posed by one user to another user), + (2) a 'gold' (ground truth) answer, + (3) a generated answer + which you will score as CORRECT/WRONG. + + The point of the question is to ask about something one user should know about the other user based on their prior conversations. + The gold answer will usually be a concise and short answer that includes the referenced topic, for example: + Question: Do you remember which keepsake I chose at the fictional Harborlight fair? + Gold answer: A copper compass pin + The generated answer might be much longer, but you should be generous with your grading - as long as it touches on the same topic as the gold answer, it should be counted as CORRECT. + + For time related questions, the gold answer will be a specific date, month, year, etc. The generated answer might be much longer or use relative time references (like "last Tuesday" or "next month"), but you should be generous with your grading - as long as it refers to the same date or time period as the gold answer, it should be counted as CORRECT. Even if the format differs (e.g., "May 7th" vs "7 May"), consider it CORRECT if it's the same date. + There's an edge case where the actual answer can't be found in the data and in that case the gold answer will say so (e.g. 'You did not mention this information.'); if the generated answer says that it cannot be answered or it doesn't know all the details, it should be counted as CORRECT. +""" + + judgement = await self.llm_config.call( + messages=[ + { + "role": "user", + "content": f"""{prompt_content} + + +Question: {question} +Gold answer: {correct_answer} +Generated answer: {predicted_answer} +First, provide a short (one sentence) explanation of your reasoning. Short reasoning is preferred. +If it's correct, set correct=true. +""", + } + ], + response_format=JudgeResponse, + scope="judge", + temperature=0, + max_completion_tokens=4096, + ) + + judge_time = time.time() - judge_start_time + console.print(f" [cyan]Answer judged in {judge_time:.1f}s[/cyan]") + + return judgement.correct, judgement.reasoning, judge_time + + except Exception as e: + # Check if it's a validation error (LLM returned malformed JSON) + error_str = str(e) + is_validation_error = "ValidationError" in error_str or "Field required" in error_str + + # Retry on validation errors, fail immediately on other errors + if is_validation_error and attempt < max_retries - 1: + print(f"Judge validation error on attempt {attempt + 1}/{max_retries}, retrying...") + await asyncio.sleep(0.5) # Small delay before retry + continue + + # Provider and parsing failures are not benchmark judgments. + # Propagate them so the caller records this question as + # invalid instead of silently scoring it as an ordinary + # incorrect answer. + raise RuntimeError(f"Judge failed after {attempt + 1} attempt(s): {e}") from e + + +@dataclass +class CoarseSearchCandidate: + """Candidate document produced by the coarse retrieval stage.""" + + rank: int + document_id: str + text: str + context: Optional[str] = None + occurred_start: Optional[str] = None + fact_type: str = "unknown" + rrf_score: float = 0.0 + proof_count: Optional[int] = None + + +@dataclass +class CoarseSearchResults: + """Results produced by the coarse retrieval stage.""" + + total_candidates: int + candidates: list[CoarseSearchCandidate] + + +@dataclass +class RerankedCandidate: + """Candidate document after reranking.""" + + original_rank: int + document_id: str + text: str + cross_encoder_score: float + combined_score: float + final_rank: int + + +@dataclass +class RerankedResults: + """Results produced by reranking.""" + + reranker_model: str + reranker_provider: str + reranked_candidates: list[RerankedCandidate] + + +@dataclass +class RecallPlan: + """Per-question retrieval controls selected by a planner.""" + + name: str = "default" + session_expansion_weight: Optional[float] = None + query_rewriting_enabled: Optional[bool] = None + query_rewriting_strategy_name: Optional[str] = None + max_tokens: Optional[int] = None + include_chunks: Optional[bool] = None + max_chunk_tokens: Optional[int] = None + evidence_appendix_mode: Optional[str] = None + + +RetrievalPlanner = Callable[[str, Optional[str], Optional[datetime]], RecallPlan] + + +def _appendix_session_key(fact: Dict[str, Any]) -> str: + document_id = fact.get("document_id") + if document_id: + return str(document_id) + + context = fact.get("context") or "" + match = re.search(r"Session\s+([^\s]+)", context) + if match: + return match.group(1) + + return "__unknown_session__" + + +def _truncate_appendix_text(text: str, max_chars: int = 360) -> str: + normalized = " ".join((text or "").replace("<|endoftext|>", " ").split()) + if len(normalized) <= max_chars: + return normalized + return normalized[: max_chars - 3].rstrip() + "..." + + +def add_cross_session_evidence_appendix( + recall_result: Dict[str, Any], + *, + max_sessions: int = 6, + per_session_facts: int = 3, + max_chars: int = 360, + instruction: Optional[str] = None, +) -> Dict[str, Any]: + """Attach grouped cross-session evidence without changing recall ordering.""" + results = recall_result.get("results") + if not isinstance(results, list) or not results: + return recall_result + + grouped: Dict[str, list[Dict[str, Any]]] = {} + session_order: list[str] = [] + for fact in results: + if not isinstance(fact, dict): + continue + session_key = _appendix_session_key(fact) + if session_key not in grouped: + grouped[session_key] = [] + session_order.append(session_key) + if len(grouped[session_key]) < per_session_facts: + grouped[session_key].append(fact) + + if len(session_order) <= 1: + return recall_result + + sessions = [] + for session_key in session_order[:max_sessions]: + evidence = [] + for fact in grouped.get(session_key, []): + evidence.append( + { + "id": fact.get("id"), + "text": _truncate_appendix_text(fact.get("text") or "", max_chars=max_chars), + "fact_type": fact.get("fact_type"), + "occurred_start": fact.get("occurred_start"), + "mentioned_at": fact.get("mentioned_at"), + "document_id": fact.get("document_id"), + "entities": (fact.get("entities") or [])[:8], + } + ) + if evidence: + sessions.append({"session_id": session_key, "evidence": evidence}) + + if not sessions: + return recall_result + + updated = dict(recall_result) + updated["cross_session_evidence"] = { + "mode": "appendix", + "instruction": instruction + or ( + "Supplementary grouped evidence for cross-session comparison. " + "Use this only to compare facts across sessions; preserve the main retrieved results as primary evidence." + ), + "max_sessions": max_sessions, + "per_session_facts": per_session_facts, + "sessions": sessions, + } + return updated + + +@dataclass +class RetrievalCache: + """Diagnostic retrieval cache for an incorrectly answered question.""" + + question_id: str + question: str + question_date: Optional[str] + category: str + correct_answer: str + generated_answer: str + is_correct: bool + judge_reasoning: str + retrieval_timestamp: str + coarse_search_results: CoarseSearchResults + reranked_results: RerankedResults + + +class BenchmarkRunner: + """ + Common benchmark runner for retain, recall, answer, and judge evaluation. + + Optimizations: + - Batch ingestion (put_batch_async) + - Parallel question processing with rate limiting + - Parallel LLM judging with rate limiting + - Progress tracking + """ + + def __init__( + self, + dataset: BenchmarkDataset, + answer_generator: LLMAnswerGenerator, + answer_evaluator: LLMAnswerEvaluator, + memory: Optional[MemoryEngine] = None, + query_rewriting_strategy_name: str = "noop", + query_rewriting_enabled: bool = False, + session_expansion_weight: float = 0.3, + retrieval_planner: Optional[RetrievalPlanner] = None, + ): + """ + Initialize benchmark runner. + + Args: + dataset: Dataset implementation + answer_generator: Answer generator implementation + answer_evaluator: Answer evaluator implementation + memory: Memory system instance (creates new if None) + query_rewriting_strategy_name: Name of query rewriting strategy to use + query_rewriting_enabled: Whether to enable query rewriting + session_expansion_weight: Weight for session-based node expansion (default 0.3) + retrieval_planner: Optional per-question planner for dynamic retrieval controls + """ + import os + + self.dataset = dataset + self.answer_generator = answer_generator + self.answer_evaluator = answer_evaluator + self.template_path: Optional[str] = None + self.query_rewriting_strategy_name = query_rewriting_strategy_name + self.query_rewriting_enabled = query_rewriting_enabled + self.session_expansion_weight = session_expansion_weight + self.retrieval_planner = retrieval_planner + self._diagnostic_cache_dir: Optional[Path] = None + self.memory = memory or MemoryEngine( + db_url=os.getenv("HMS_API_DATABASE_URL", "pg0"), + memory_llm_provider=os.getenv("HMS_API_LLM_PROVIDER", "groq"), + memory_llm_api_key=os.getenv("HMS_API_LLM_API_KEY"), + memory_llm_model=os.getenv("HMS_API_LLM_MODEL", "openai/gpt-oss-20b"), + memory_llm_base_url=os.getenv("HMS_API_LLM_BASE_URL") or None, + ) + + def _save_retrieval_cache(self, cache: RetrievalCache, cache_dir: Path) -> bool: + """Persist diagnostic retrieval data without affecting judge validity.""" + try: + cache_dir.mkdir(parents=True, exist_ok=True) + timestamp = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S") + filename = f"retrieval_cache_wrong_{cache.question_id}_{timestamp}.json" + filepath = cache_dir / filename + + with open(filepath, "w", encoding="utf-8") as f: + json.dump(asdict(cache), f, ensure_ascii=False, indent=2) + return True + except OSError as exc: + logging.warning( + "Could not write optional retrieval cache for %s to %s: %s", + cache.question_id, + cache_dir, + exc, + ) + return False + + def calculate_data_stats(self, items: List[Dict[str, Any]]) -> Dict[str, Any]: + """ + Calculate statistics about the data to be ingested. + + Returns: + Dict with statistics: total_sessions, total_chars, avg_session_length, etc. + """ + total_sessions = 0 + total_chars = 0 + session_lengths = [] + + for item in items: + batch_contents = self.dataset.prepare_sessions_for_ingestion(item) + total_sessions += len(batch_contents) + + for session in batch_contents: + content_len = len(session["content"]) + total_chars += content_len + session_lengths.append(content_len) + + avg_length = total_chars / total_sessions if total_sessions > 0 else 0 + + return { + "total_sessions": total_sessions, + "total_chars": total_chars, + "total_items": len(items), + "avg_session_length": avg_length, + "min_session_length": min(session_lengths) if session_lengths else 0, + "max_session_length": max(session_lengths) if session_lengths else 0, + } + + async def apply_template(self, bank_id: str, manifest_path: str) -> None: + """Apply a bank template manifest to a bank before ingestion. + + Reads the manifest JSON file and applies config overrides, creates + mental models and directives — same logic as the /import API endpoint. + """ + from hms_api.api.http import BankTemplateManifest + from hms_api.models import RequestContext + + raw = json.loads(Path(manifest_path).read_text()) + manifest = BankTemplateManifest.model_validate(raw) + + request_context = RequestContext() + await self.memory.get_bank_profile(bank_id, request_context=request_context) + + # Apply bank config overrides + if manifest.bank: + config_updates = manifest.bank.get_config_updates() + if config_updates: + await self.memory._config_resolver.update_bank_config(bank_id, config_updates, request_context) + + # Create directives + for directive in manifest.directives or []: + await self.memory.create_directive( + bank_id=bank_id, + name=directive.name, + content=directive.content, + priority=directive.priority, + is_active=directive.is_active, + tags=directive.tags if directive.tags else None, + request_context=request_context, + ) + + # Create mental models (async content generation) + for mm in manifest.mental_models or []: + mental_model = await self.memory.create_mental_model( + bank_id=bank_id, + name=mm.name, + source_query=mm.source_query, + content="Generating content...", + mental_model_id=mm.id, + tags=mm.tags if mm.tags else None, + max_tokens=mm.max_tokens, + trigger=mm.trigger.model_dump() if mm.trigger else None, + request_context=request_context, + ) + await self.memory.submit_async_refresh_mental_model( + bank_id=bank_id, + mental_model_id=mental_model["id"], + request_context=request_context, + ) + + async def ingest_conversation( + self, item: Dict[str, Any], agent_id: str, wait_for_consolidation: bool = False + ) -> int: + """ + Ingest conversation into memory using batch ingestion. + + Uses put_batch_async for maximum efficiency. + + Args: + item: Dataset item to ingest + agent_id: Agent/bank ID to ingest into + wait_for_consolidation: If True, wait for consolidation to complete after ingestion + + Returns: + Number of sessions ingested + """ + batch_contents = self.dataset.prepare_sessions_for_ingestion(item) + + if batch_contents: + max_retries = 3 + retry_delay = 2.0 + + for attempt in range(max_retries): + try: + await self.memory.retain_batch_async( + bank_id=agent_id, + contents=batch_contents, + request_context=RequestContext(), + ) + break + except ValueError as e: + error_msg = str(e) + if "Batch contains duplicate document_ids" in error_msg: + console.print( + f" [yellow]⚠[/yellow] Duplicate document_id error on attempt {attempt + 1}, " + f"deduplicating batch contents..." + ) + unique_contents = [] + seen_doc_ids = set() + for content in batch_contents: + doc_id = content.get("document_id") + if doc_id and doc_id in seen_doc_ids: + continue + unique_contents.append(content) + if doc_id: + seen_doc_ids.add(doc_id) + batch_contents = unique_contents + if attempt == max_retries - 1: + raise + await asyncio.sleep(retry_delay) + retry_delay *= 1.5 + continue + else: + raise + except Exception as e: + console.print(f" [yellow]⚠[/yellow] Ingestion error: {e}") + if attempt == max_retries - 1: + raise + await asyncio.sleep(retry_delay) + retry_delay *= 1.5 + + if wait_for_consolidation and batch_contents: + await self._wait_for_consolidation(agent_id) + + return len(batch_contents) + + async def _get_pending_consolidation_count(self, bank_id: str) -> int: + """ + Get the count of memories pending consolidation. + + Returns: + Number of memories not yet processed by the consolidation job + """ + pool = await self.memory._get_pool() + from hms_api.engine.memory_engine import fq_table + + async with pool.acquire() as conn: + result = await conn.fetchrow( + f""" + SELECT COUNT(*) as count + FROM {fq_table("memory_units")} + WHERE bank_id = $1 AND consolidated_at IS NULL AND fact_type IN ('experience', 'world') + """, + bank_id, + ) + return result["count"] if result else 0 + + async def _wait_for_consolidation(self, bank_id: str, poll_interval: float = 2.0, timeout: float = 3000.0) -> None: + """ + Wait for consolidation to complete (pending_consolidation reaches 0). + + Args: + bank_id: Bank ID to check + poll_interval: Seconds between polls + timeout: Maximum seconds to wait + + Raises: + TimeoutError: If consolidation doesn't complete within timeout + """ + import time + + start_time = time.time() + console.print(" [yellow]Waiting for consolidation to complete...[/yellow]") + + while True: + elapsed = time.time() - start_time + if elapsed > timeout: + raise TimeoutError(f"Consolidation did not complete within {timeout}s") + + pending = await self._get_pending_consolidation_count(bank_id) + if pending == 0: + console.print(" [green]✓[/green] Consolidation complete") + return + + # Still pending, wait and poll again + await asyncio.sleep(poll_interval) + + async def answer_question( + self, + agent_id: str, + question: str, + thinking_budget: int = 500, + max_tokens: int = 4096, + question_date: Optional[datetime] = None, + question_type: Optional[str] = None, + ) -> Tuple[str, str, List[Dict], Dict[str, Dict], Optional[Dict[str, Any]]]: + """ + Answer a question using memory retrieval. + + Args: + agent_id: Agent ID + question: Question text + thinking_budget: Thinking budget for search + max_tokens: Maximum tokens to retrieve + question_date: Date when the question was asked (for temporal filtering) + question_type: Question category/type (e.g., 'multi-session', 'temporal-reasoning') + + Returns: + Tuple of (answer, reasoning, retrieved_memories, chunks, retrieval_details) + - retrieval_details: Optional dict with coarse_search_results and reranked_results + """ + retrieval_details = None + + # Check if generator needs external search + if self.answer_generator.needs_external_search(): + # Traditional flow: search then generate + # Use MemoryEngine directly + # Map thinking_budget to budget level + budget = Budget.LOW if thinking_budget <= 30 else Budget.MID if thinking_budget <= 70 else Budget.HIGH + + import time + + recall_start_time = time.time() + plan = ( + self.retrieval_planner(question, question_type, question_date) + if self.retrieval_planner is not None + else RecallPlan() + ) + recall_max_tokens = plan.max_tokens if plan.max_tokens is not None else max_tokens + recall_include_chunks = plan.include_chunks if plan.include_chunks is not None else True + recall_max_chunk_tokens = plan.max_chunk_tokens if plan.max_chunk_tokens is not None else 8192 + recall_query_rewriting_enabled = ( + plan.query_rewriting_enabled + if plan.query_rewriting_enabled is not None + else self.query_rewriting_enabled + ) + recall_query_rewriting_strategy = ( + plan.query_rewriting_strategy_name + if plan.query_rewriting_strategy_name is not None + else self.query_rewriting_strategy_name + ) + recall_session_expansion_weight = ( + plan.session_expansion_weight + if plan.session_expansion_weight is not None + else self.session_expansion_weight + ) + + # Use default fact types (no filtering) + search_result = await self.memory.recall_async( + bank_id=agent_id, + query=question, + budget=budget, + max_tokens=recall_max_tokens, + question_date=question_date, + include_entities=True, + max_entity_tokens=2048, + include_chunks=recall_include_chunks, + max_chunk_tokens=recall_max_chunk_tokens, + request_context=RequestContext(), + query_rewriting_strategy_name=recall_query_rewriting_strategy, + query_rewriting_enabled=recall_query_rewriting_enabled, + session_expansion_weight=recall_session_expansion_weight, + ) + recall_time = time.time() - recall_start_time + + # Log recall stats + num_results = len(search_result.results) if search_result.results else 0 + num_chunks = len(search_result.chunks) if search_result.chunks else 0 + num_entities = len(search_result.entities) if search_result.entities else 0 + + # Convert entire RecallResult to dictionary for answer generation + recall_result_dict = search_result.model_dump() + if plan.evidence_appendix_mode == "cross_session": + recall_result_dict = add_cross_session_evidence_appendix(recall_result_dict) + elif plan.evidence_appendix_mode == "cross_session_compact": + recall_result_dict = add_cross_session_evidence_appendix( + recall_result_dict, + max_sessions=4, + per_session_facts=2, + max_chars=240, + instruction=( + "Compact supplementary evidence for questions that require counting, ordering, or comparing " + "facts across sessions. Use it only when it directly supports a multi-session calculation; " + "do not override more specific main evidence, and answer that information is insufficient " + "when neither source directly supports the answer." + ), + ) + + # Extract retrieval details from trace for caching + trace = search_result.trace or {} + rrf_results = trace.get("rrf_merged", []) + reranked_results = trace.get("reranked", []) + + # Extract cross-encoder config info + cross_encoder_config = {} + if hasattr(self.memory, "_cross_encoder") and self.memory._cross_encoder: + ce = self.memory._cross_encoder + cross_encoder_config["model"] = getattr(ce, "model_name", "unknown") + cross_encoder_config["provider"] = getattr(ce, "provider_name", "unknown") + else: + cross_encoder_config["model"] = "unknown" + cross_encoder_config["provider"] = "unknown" + + # Build coarse search results from RRF merge + coarse_candidates = [] + for rrf_item in rrf_results: + if isinstance(rrf_item, dict): + node_id = rrf_item.get("node_id", "") + text = rrf_item.get("text", "") + # Find matching result for context and metadata + matching_result = None + if search_result.results: + for result in search_result.results: + if result.id == node_id or (result.document_id and node_id in result.document_id): + matching_result = result + break + + coarse_candidates.append( + CoarseSearchCandidate( + rank=rrf_item.get("final_rrf_rank", 0), + document_id=node_id, + text=text, + context=matching_result.context if matching_result else None, + occurred_start=matching_result.occurred_start if matching_result else None, + fact_type=matching_result.fact_type if matching_result else "unknown", + rrf_score=rrf_item.get("rrf_score", 0.0), + proof_count=matching_result.metadata.get("proof_count") + if matching_result and matching_result.metadata + else None, + ) + ) + + coarse_search_results = CoarseSearchResults( + total_candidates=len(coarse_candidates), + candidates=coarse_candidates, + ) + + # Build reranked results + reranked_candidates = [] + for rerank_item in reranked_results: + if isinstance(rerank_item, dict): + node_id = rerank_item.get("node_id", "") + text = rerank_item.get("text", "") + score_components = rerank_item.get("score_components", {}) + + reranked_candidates.append( + RerankedCandidate( + original_rank=rerank_item.get("rrf_rank", 0), + document_id=node_id, + text=text, + cross_encoder_score=rerank_item.get("rerank_score", 0.0), + combined_score=score_components.get("combined_score", 0.0), + final_rank=rerank_item.get("rerank_rank", 0), + ) + ) + + reranked_results_obj = RerankedResults( + reranker_model=cross_encoder_config["model"], + reranker_provider=cross_encoder_config["provider"], + reranked_candidates=reranked_candidates, + ) + + retrieval_details = { + "coarse_search_results": coarse_search_results, + "reranked_results": reranked_results_obj, + } + + # Extract chunks from search result + chunks = {} + if search_result.chunks: + for chunk_key, chunk_info in search_result.chunks.items(): + chunks[chunk_key] = chunk_info.model_dump() + + # Check if we have any results + if not search_result.results: + return ( + "I don't have enough information to answer that question.", + "No relevant memories found.", + [], + {}, + None, + ) + + # Generate answer using LLM - pass entire recall result + answer, reasoning, memories_override = await self.answer_generator.generate_answer( + question, recall_result_dict, question_date, question_type, bank_id=agent_id + ) + + # Use override if provided, otherwise use the results from recall + final_memories = ( + memories_override + if memories_override is not None + else recall_result_dict.get("results", [fact.model_dump() for fact in search_result.results]) + ) + + return answer, reasoning, final_memories, chunks, retrieval_details + else: + # Integrated flow: generator does its own search (e.g., reflect API) + # Pass empty recall result since generator doesn't need them + answer, reasoning, memories_override = await self.answer_generator.generate_answer( + question, {"results": []}, question_date, question_type, bank_id=agent_id + ) + + # Use memories from generator (should not be None for integrated mode) + final_memories = memories_override if memories_override is not None else [] + + return answer, reasoning, final_memories, {}, None + + async def evaluate_qa_task( + self, + agent_id: str, + qa_pairs: List[Dict], + item_id: str, + thinking_budget: int, + max_tokens: int, + max_questions: Optional[int] = None, + semaphore: asyncio.Semaphore = None, + ) -> List[Dict]: + """ + Evaluate QA task with parallel question processing. + + Args: + semaphore: Semaphore to limit concurrent question processing + + Returns: + List of QA results + """ + # Filter out questions without answers (category 5) + # First, identify and log category 5 questions that will be skipped + category_5_questions = [pair for pair in qa_pairs if pair.get("category") == 5] + if category_5_questions: + logging.info(f"Skipping {len(category_5_questions)} category=5 questions for {item_id}") + for q in category_5_questions: + logging.debug(f" Skipped category=5 question: {q.get('question', 'N/A')[:100]}") + + # Filter out category 5 and questions without answers, preserving original indices + indexed_pairs = [ + (orig_idx, pair) + for orig_idx, pair in enumerate(qa_pairs) + if pair.get("category") != 5 and pair.get("answer") + ] + indexed_pairs_to_eval = indexed_pairs[:max_questions] if max_questions else indexed_pairs + + # Progress output disabled for cleaner logs + async def process_question(orig_idx: int, qa: dict): + async with semaphore: + question = qa["question"] + correct_answer = qa["answer"] + category = qa.get("category", 0) + question_date = qa.get("question_date") + + import time + + start_time = time.time() + + try: + ( + predicted_answer, + reasoning, + retrieved_memories, + chunks, + retrieval_details, + ) = await self.answer_question( + agent_id, + question, + thinking_budget, + max_tokens, + question_date, + category, + ) + + answer_time = time.time() - start_time + + memories_without_embeddings = [ + {k: v for k, v in mem.items() if k != "embedding"} for mem in retrieved_memories + ] + + return { + "question_index": orig_idx, + "question": question, + "correct_answer": correct_answer, + "predicted_answer": predicted_answer, + "reasoning": reasoning, + "category": category, + "retrieved_memories": memories_without_embeddings, + "is_invalid": False, + "error": None, + "answer_time": answer_time, + "retrieval_details": retrieval_details, + } + except Exception as e: + logging.exception(f"Failed to answer question: {question[:100]}") + return { + "question_index": orig_idx, + "question": question, + "correct_answer": correct_answer, + "predicted_answer": "ERROR: Failed to generate answer", + "reasoning": f"Error: {str(e)}", + "category": category, + "retrieved_memories": [], + "is_invalid": True, + "error": str(e), + "retrieval_details": None, + } + + question_tasks = [process_question(orig_idx, qa) for orig_idx, qa in indexed_pairs_to_eval] + + results = await asyncio.gather(*question_tasks, return_exceptions=True) + results = [r if not isinstance(r, Exception) else {"error": str(r)} for r in results] + + return results + + async def calculate_metrics( + self, + results: List[Dict], + eval_semaphore: asyncio.Semaphore, + item_id: Optional[str] = None, + ) -> Dict: + """ + Calculate evaluation metrics using parallel LLM-as-judge. + + Args: + results: QA results to evaluate + eval_semaphore: Run-scoped semaphore for all LLM judge requests + item_id: Optional item ID for retrieval cache naming + + Returns: + Dict with evaluation metrics + """ + total = len(results) + + # Progress output disabled for cleaner logs + async def judge_single(result): + # Skip judging if already marked as invalid + if result.get("is_invalid", False): + result["is_correct"] = None + result["correctness_reasoning"] = ( + f"Question invalid due to error: {result.get('error', 'Unknown error')}" + ) + return result + + try: + is_correct, eval_reasoning, judge_time = await self.answer_evaluator.judge_answer( + result["question"], + result["correct_answer"], + result["predicted_answer"], + eval_semaphore, + category=result.get("category"), + ) + result["is_correct"] = is_correct + result["correctness_reasoning"] = eval_reasoning + result["judge_time"] = judge_time + + if not is_correct: + retrieval_details = result.get("retrieval_details") + if retrieval_details: + cache = RetrievalCache( + question_id=f"{item_id}_{result.get('question_index', 0)}", + question=result["question"], + question_date=result.get("question_date"), + category=str(result.get("category", "unknown")), + correct_answer=result["correct_answer"], + generated_answer=result["predicted_answer"], + is_correct=is_correct, + judge_reasoning=eval_reasoning, + retrieval_timestamp=datetime.now(timezone.utc).isoformat(), + coarse_search_results=retrieval_details.get("coarse_search_results"), + reranked_results=retrieval_details.get("reranked_results"), + ) + if self._diagnostic_cache_dir is not None: + self._save_retrieval_cache(cache, self._diagnostic_cache_dir) + + return result + except Exception as e: + logging.exception(f"Failed to judge answer for question: {result.get('question', 'unknown')[:100]}") + result["is_invalid"] = True + result["is_correct"] = None + result["correctness_reasoning"] = f"Judge error: {str(e)}" + result["error"] = str(e) + return result + + judgment_tasks = [judge_single(result) for result in results] + judged_results = await asyncio.gather(*judgment_tasks, return_exceptions=True) + judged_results = [r if not isinstance(r, Exception) else {"error": str(r)} for r in judged_results] + + # Calculate stats + correct = sum(1 for r in judged_results if r.get("is_correct", False)) + invalid = sum(1 for r in judged_results if r.get("is_invalid", False)) + valid_total = total - invalid + category_stats = {} + + for result in judged_results: + category = result.get("category", "unknown") + if category not in category_stats: + category_stats[category] = {"correct": 0, "total": 0, "invalid": 0} + category_stats[category]["total"] += 1 + if result.get("is_invalid", False): + category_stats[category]["invalid"] += 1 + elif result.get("is_correct", False): + category_stats[category]["correct"] += 1 + + # Provider and processing failures count as incorrect in the primary + # benchmark metric. Keep valid-only accuracy as a diagnostic. + accuracy = (correct / total * 100) if total > 0 else 0 + valid_accuracy = (correct / valid_total * 100) if valid_total > 0 else 0 + + return { + "accuracy": accuracy, + "correct": correct, + "total": total, + "invalid": invalid, + "valid_total": valid_total, + "valid_accuracy": valid_accuracy, + "category_stats": category_stats, + "detailed_results": judged_results, + } + + async def _agent_has_data(self, agent_id: str) -> bool: + """ + Check if an agent has any indexed memory units. + + Args: + agent_id: Agent ID to check + + Returns: + True if agent has at least one memory unit, False otherwise + """ + try: + # A bank is reusable only when it has durable source chunks. A + # document may legitimately produce no extracted facts. + pool = await self.memory._get_pool() + async with pool.acquire() as conn: + result = await conn.fetchval( + f"SELECT EXISTS(SELECT 1 FROM {fq_table('chunks')} WHERE bank_id = $1)", + agent_id, + ) + return bool(result) + except Exception as e: + console.print(f" [red]Warning: Error checking agent data: {e}[/red]") + return False + + async def _audit_durable_ingestion(self, item: Dict[str, Any], agent_id: str) -> Dict[str, Any]: + """Verify that every input document has at least one durable source chunk. + + Fact extraction is lossy by design, so a document with zero facts is + reported but remains valid. Missing documents and zero-chunk documents + are integrity failures because recall cannot recover their source text. + """ + + prepared = self.dataset.prepare_sessions_for_ingestion(item) + expected_document_ids = list( + dict.fromkeys(str(content["document_id"]) for content in prepared if content.get("document_id") is not None) + ) + report: Dict[str, Any] = { + "item_id": self.dataset.get_item_id(item), + "bank_id": agent_id, + "expected_documents": len(expected_document_ids), + "durable_documents": 0, + "missing_documents": [], + "documents_without_chunks": [], + "documents_without_facts": [], + } + if not expected_document_ids: + return report + + pool = await self.memory._get_pool() + async with pool.acquire() as conn: + rows = await conn.fetch( + f""" + SELECT d.id, + COUNT(DISTINCT c.chunk_id) AS chunk_count, + COUNT(DISTINCT m.id) AS fact_count + FROM {fq_table("documents")} AS d + LEFT JOIN {fq_table("chunks")} AS c + ON c.bank_id = d.bank_id AND c.document_id = d.id + LEFT JOIN {fq_table("memory_units")} AS m + ON m.bank_id = d.bank_id AND m.document_id = d.id + WHERE d.bank_id = $1 AND d.id = ANY($2::text[]) + GROUP BY d.id + """, + agent_id, + expected_document_ids, + ) + + by_id = {str(row["id"]): row for row in rows} + report["durable_documents"] = len(by_id) + report["missing_documents"] = [document_id for document_id in expected_document_ids if document_id not in by_id] + report["documents_without_chunks"] = [ + document_id + for document_id in expected_document_ids + if document_id in by_id and int(by_id[document_id]["chunk_count"]) == 0 + ] + report["documents_without_facts"] = [ + document_id + for document_id in expected_document_ids + if document_id in by_id and int(by_id[document_id]["fact_count"]) == 0 + ] + + if report["missing_documents"] or report["documents_without_chunks"]: + raise IngestionIntegrityError(report) + return report + + async def process_single_item( + self, + item: Dict, + agent_id: str, + i: int, + total_items: int, + thinking_budget: int, + max_tokens: int, + max_questions_per_item: Optional[int], + skip_ingestion: bool, + question_semaphore: asyncio.Semaphore, + eval_semaphore: asyncio.Semaphore, + clear_this_agent: bool = True, + wait_consolidation: bool = False, + ingest_only: bool = False, + skip_if_already_ingested: bool = False, + force_reingest: bool = False, + ) -> Dict: + """ + Process a single item (ingest + evaluate). + + Args: + clear_this_agent: Whether to clear this agent's data before ingesting. + Set to False to skip clearing (e.g., when agent_id is shared and already cleared) + wait_consolidation: If True, wait for consolidation to complete before evaluating QA. + skip_if_already_ingested: If True and agent already has data, skip ingestion entirely. + + Returns: + Result dict with metrics + """ + item_id = self.dataset.get_item_id(item) + + console.print(f"\n[bold blue]Item {i}/{total_items}[/bold blue] (ID: {item_id})") + + step = 1 + num_sessions = 0 + if not skip_ingestion: + # Check if already ingested (for smart resume) + already_ingested = False + if skip_if_already_ingested and not force_reingest: + already_ingested = await self._agent_has_data(agent_id) + if already_ingested: + console.print(f" [{step}] [yellow]⊘[/yellow] Skipping - already ingested") + + if not already_ingested or force_reingest: + # Clear agent data before ingesting (always clear when force_reingest) + if clear_this_agent or force_reingest: + console.print(f" [{step}] Clearing previous agent data...") + await self.memory.delete_bank(agent_id, request_context=RequestContext()) + console.print(f" [green]✓[/green] Cleared '{agent_id}' agent data") + + # Apply template if configured + if self.template_path: + step += 1 + console.print(f" [{step}] Applying bank template...") + await self.apply_template(agent_id, self.template_path) + console.print(" [green]✓[/green] Template applied") + + # Ingest conversation + step += 1 + console.print(f" [{step}] Ingesting conversation (batch mode)...") + num_sessions = await self.ingest_conversation(item, agent_id, wait_for_consolidation=False) + console.print(f" [green]✓[/green] Ingested {num_sessions} sessions") + + step += 1 + console.print(f" [{step}] Auditing durable ingestion...") + ingestion_audit = await self._audit_durable_ingestion(item, agent_id) + console.print( + " [green]✓[/green] " + f"{ingestion_audit['durable_documents']}/{ingestion_audit['expected_documents']} " + "documents have durable chunks" + ) + + # Wait for consolidation before evaluating if requested + if wait_consolidation: + step += 1 + console.print(f" [{step}] Waiting for consolidation...") + await self._wait_for_consolidation(agent_id) + + # Ingest-only mode: skip evaluation + if ingest_only: + console.print(" [green]✓[/green] Ingest complete (skipping evaluation)") + return { + "item_id": item_id, + "metrics": {"correct": 0, "total": 0, "invalid": 0, "accuracy": 0.0}, + "num_sessions": num_sessions, + "ingest_only": True, + "ingestion_audit": ingestion_audit, + } + + # Evaluate QA + step += 1 + qa_pairs = self.dataset.get_qa_pairs(item) + console.print(f" [{step}] Evaluating {len(qa_pairs)} QA pairs (parallel)...") + qa_results = await self.evaluate_qa_task( + agent_id, + qa_pairs, + item_id, + thinking_budget, + max_tokens, + max_questions_per_item, + question_semaphore, + ) + + # Calculate metrics + step += 1 + console.print(f" [{step}] Calculating metrics...") + metrics = await self.calculate_metrics(qa_results, eval_semaphore, item_id) + + console.print( + f" [green]✓[/green] Accuracy: {metrics['accuracy']:.2f}% ({metrics['correct']}/{metrics['total']})" + ) + + return { + "item_id": item_id, + "metrics": metrics, + "num_sessions": num_sessions, + "ingestion_audit": ingestion_audit, + } + + async def run( + self, + dataset_path: Path, + agent_id: str, + max_items: Optional[int] = None, + max_questions_per_item: Optional[int] = None, + thinking_budget: int = 500, + max_tokens: int = 4096, + skip_ingestion: bool = False, + max_concurrent_questions: int = 10, + eval_semaphore_size: int = 10, + clear_agent_per_item: bool = False, + specific_item: Optional[Union[str, Iterable[str]]] = None, + separate_ingestion_phase: bool = False, + filln: bool = False, + max_concurrent_items: int = 1, # Max concurrent items (conversations) to process in parallel + output_path: Optional[Path] = None, # Path to save results incrementally + merge_with_existing: bool = False, # Whether to merge with existing results + wait_consolidation: bool = False, # Wait for consolidation to complete before evaluating QA + template_path: Optional[str] = None, # Path to a bank template manifest to apply before ingestion + ingest_only: bool = False, # Only ingest, skip evaluation + force_reingest: bool = False, # If True, always re-ingest even if data already exists + rerun_invalid_existing: bool = False, # Resume mode: rerun invalid existing item results + run_manifest: Optional[Dict[str, Any]] = None, # Stable metadata included in every checkpoint + ) -> Dict[str, Any]: + """ + Run the full benchmark evaluation. + + Args: + dataset_path: Path to dataset file + agent_id: Agent ID to use + max_items: Maximum number of items to evaluate + max_questions_per_item: Maximum questions per item + thinking_budget: Thinking budget for search + max_tokens: Maximum tokens to retrieve from memories + skip_ingestion: Skip ingestion and use existing data + max_concurrent_questions: Max concurrent question processing + eval_semaphore_size: Max concurrent LLM judge requests + clear_agent_per_item: Use unique agent ID per item for isolation (deprecated when separate_ingestion_phase=True) + specific_item: If provided, only run this specific item ID (e.g., conversation) + separate_ingestion_phase: If True, ingest all data first, then evaluate all questions (single agent) + filln: If True, skip item IDs already complete in the result artifact + max_concurrent_items: Max concurrent items to process in parallel (requires clear_agent_per_item=True) + + Returns: + Dict with complete benchmark results + """ + console.print("\n[bold cyan]Benchmark Evaluation[/bold cyan]") + console.print("=" * 80) + + for name, value in ( + ("max_concurrent_items", max_concurrent_items), + ("max_concurrent_questions", max_concurrent_questions), + ("eval_semaphore_size", eval_semaphore_size), + ): + if value < 1: + raise ValueError(f"{name} must be a positive integer, got {value}") + + self._run_manifest = run_manifest + self._diagnostic_cache_dir = output_path.parent / "retrieval_cache" if output_path is not None else None + + # Print model configuration + print_model_config() + + # Load dataset + console.print(f"\n[1] Loading dataset from {dataset_path}...") + items = self.dataset.load(dataset_path, max_items) + + # Filter for specific item(s) if requested + if specific_item is not None: + target_ids = {specific_item} if isinstance(specific_item, str) else set(specific_item) + items = [item for item in items if self.dataset.get_item_id(item) in target_ids] + if not items: + console.print(f" [red]✗[/red] No item found with ID(s): {sorted(target_ids)}") + raise ValueError(f"No items matching ID(s) {sorted(target_ids)} found in dataset") + console.print(f" [green]✓[/green] Filtering to {len(items)} item(s): {sorted(target_ids)}") + + console.print(f" [green]✓[/green] Loaded {len(items)} items") + + # Initialize memory system + console.print("\n[2] Initializing memory system...") + if template_path: + self.template_path = template_path + console.print(f" Bank template: {template_path}") + console.print(" [green]✓[/green] Memory system initialized") + + # Start a background worker poller when we need to wait for consolidation. + # Consolidation is submitted as an async task by retain_batch_async, but + # without a running worker those tasks sit in the queue forever. + poller_task = None + poller = None + if wait_consolidation: + from hms_api.worker.poller import WorkerPoller + + poller = WorkerPoller( + backend=self.memory._backend, + worker_id="benchmark-runner-worker", + executor=self.memory.execute_task, + poll_interval_ms=500, + max_slots=4, + ) + poller_task = asyncio.create_task(poller.run()) + console.print(" [green]✓[/green] Background worker started (for consolidation)") + + try: + return await self._run_inner( + items, + agent_id, + thinking_budget, + max_tokens, + skip_ingestion, + max_questions_per_item, + max_concurrent_questions, + eval_semaphore_size, + clear_agent_per_item, + specific_item, + separate_ingestion_phase, + filln, + max_concurrent_items, + output_path, + merge_with_existing, + wait_consolidation, + ingest_only, + force_reingest, + rerun_invalid_existing, + ) + finally: + if poller and poller_task: + await poller.shutdown_graceful(timeout=60.0) + poller_task.cancel() + try: + await poller_task + except asyncio.CancelledError: + pass + console.print(" [green]✓[/green] Background worker stopped") + + async def _run_inner( + self, + items: List[Dict[str, Any]], + agent_id: str, + thinking_budget: int, + max_tokens: int, + skip_ingestion: bool, + max_questions_per_item: Optional[int], + max_concurrent_questions: int, + eval_semaphore_size: int, + clear_agent_per_item: bool, + specific_item: Any, + separate_ingestion_phase: bool, + filln: bool, + max_concurrent_items: int, + output_path: Optional[Path], + merge_with_existing: bool, + wait_consolidation: bool, + ingest_only: bool = False, + force_reingest: bool = False, + rerun_invalid_existing: bool = False, + ) -> Dict[str, Any]: + if separate_ingestion_phase: + # New two-phase approach: ingest all, then evaluate all + return await self._run_two_phase( + items, + agent_id, + thinking_budget, + max_tokens, + skip_ingestion, + max_questions_per_item, + max_concurrent_questions, + eval_semaphore_size, + output_path, + merge_with_existing, + ) + else: + # Original approach: process each item independently + return await self._run_single_phase( + items, + agent_id, + thinking_budget, + max_tokens, + skip_ingestion, + max_questions_per_item, + max_concurrent_questions, + eval_semaphore_size, + clear_agent_per_item, + filln, + max_concurrent_items, + output_path, + merge_with_existing, + wait_consolidation, + ingest_only, + force_reingest, + rerun_invalid_existing, + ) + + async def _run_single_phase( + self, + items: List[Dict[str, Any]], + agent_id: str, + thinking_budget: int, + max_tokens: int, + skip_ingestion: bool, + max_questions_per_item: Optional[int], + max_concurrent_questions: int, + eval_semaphore_size: int, + clear_agent_per_item: bool, + filln: bool = False, + max_concurrent_items: int = 1, + output_path: Optional[Path] = None, + merge_with_existing: bool = False, + wait_consolidation: bool = False, + ingest_only: bool = False, + force_reingest: bool = False, + rerun_invalid_existing: bool = False, + ) -> Dict[str, Any]: + """Original single-phase approach: process each item independently.""" + # Create semaphore for question processing + question_semaphore = asyncio.Semaphore(max_concurrent_questions) + # This semaphore is deliberately shared by every item in the run. + eval_semaphore = asyncio.Semaphore(eval_semaphore_size) + + # Process items - either in parallel or sequentially + if max_concurrent_items > 1 and clear_agent_per_item: + # Parallel item processing (requires unique agent IDs) + all_results = await self._process_items_parallel( + items, + agent_id, + thinking_budget, + max_tokens, + skip_ingestion, + max_questions_per_item, + question_semaphore, + eval_semaphore, + filln, + max_concurrent_items, + output_path, + merge_with_existing, + wait_consolidation, + ingest_only, + force_reingest, + rerun_invalid_existing, + ) + else: + # Sequential item processing (original behavior) + all_results = await self._process_items_sequential( + items, + agent_id, + thinking_budget, + max_tokens, + skip_ingestion, + max_questions_per_item, + question_semaphore, + eval_semaphore, + clear_agent_per_item, + filln, + output_path, + merge_with_existing, + wait_consolidation, + ingest_only, + force_reingest, + rerun_invalid_existing, + ) + + # Calculate overall metrics + total_correct = sum(r["metrics"]["correct"] for r in all_results) + total_questions = sum(r["metrics"]["total"] for r in all_results) + total_invalid = sum(r["metrics"].get("invalid", 0) for r in all_results) + total_valid = total_questions - total_invalid + overall_accuracy = (total_correct / total_questions * 100) if total_questions > 0 else 0 + valid_only_accuracy = (total_correct / total_valid * 100) if total_valid > 0 else 0 + + return { + "overall_accuracy": overall_accuracy, + "total_correct": total_correct, + "total_questions": total_questions, + "total_invalid": total_invalid, + "total_valid": total_valid, + "valid_only_accuracy": valid_only_accuracy, + "num_items": len(all_results), + "model_config": get_model_config(), + "item_results": all_results, + } + + async def _process_items_sequential( + self, + items: List[Dict[str, Any]], + agent_id: str, + thinking_budget: int, + max_tokens: int, + skip_ingestion: bool, + max_questions_per_item: Optional[int], + question_semaphore: asyncio.Semaphore, + eval_semaphore: asyncio.Semaphore, + clear_agent_per_item: bool, + filln: bool, + output_path: Optional[Path] = None, + merge_with_existing: bool = False, + wait_consolidation: bool = False, + ingest_only: bool = False, + force_reingest: bool = False, + rerun_invalid_existing: bool = False, + ) -> List[Dict]: + """Process items sequentially (original behavior).""" + all_results = [] + existing_item_ids = set() + resume_complete_item_ids = set() + + # Load existing results if merge_with_existing is True + if merge_with_existing and output_path and output_path.exists(): + with open(output_path, "r", encoding="utf-8") as f: + existing_data = json.load(f) + if "item_results" in existing_data: + all_results = existing_data["item_results"] + existing_item_ids = {r["item_id"] for r in all_results} + resume_complete_item_ids = {r["item_id"] for r in all_results if _result_is_resume_complete(r)} + console.print(f"[cyan]Loaded {len(all_results)} existing results from {output_path}[/cyan]") + + # Pre-load durable banks for reuse and ingest-only fill modes. + ingested_item_ids = set() + if skip_ingestion or ingest_only: + console.print("[cyan]Checking which items are already ingested...[/cyan]") + for item in items: + item_id = self.dataset.get_item_id(item) + item_agent_id = f"{agent_id}_{item_id}" if clear_agent_per_item else agent_id + if await self._agent_has_data(item_agent_id): + ingested_item_ids.add(item_id) + if ingested_item_ids: + console.print(f"[cyan]Found {len(ingested_item_ids)} items already ingested[/cyan]") + else: + console.print("[cyan]No items found with existing ingest data[/cyan]") + if skip_ingestion: + requested_item_ids = {self.dataset.get_item_id(item) for item in items} + missing_item_ids = sorted(requested_item_ids - ingested_item_ids) + if missing_item_ids: + preview = ", ".join(missing_item_ids[:10]) + suffix = " ..." if len(missing_item_ids) > 10 else "" + raise RuntimeError( + "Retrieval-only mode requires durable retained chunks for every selected item; " + f"missing {len(missing_item_ids)} bank(s): {preview}{suffix}" + ) + + for i, item in enumerate(items, 1): + item_id = self.dataset.get_item_id(item) + # Use unique agent ID per item if requested (for isolation in benchmarks like LongMemEval) + # This avoids deadlocks from deleting agent data + if clear_agent_per_item: + item_agent_id = f"{agent_id}_{item_id}" + # Always clear for unique agents (each agent_id is used only once) + clear_this_agent = True + else: + item_agent_id = agent_id + # Only clear on first item for shared agent_id + clear_this_agent = i == 1 + + # Skip items without existing ingest data (only applies when using --skip-ingestion) + # When not using --skip-ingestion, we should ingest the items + if skip_ingestion and item_id not in ingested_item_ids and not ingest_only: + raise RuntimeError(f"Retrieval-only bank disappeared before evaluation: {item_agent_id}") + + # For ingest_only mode, skip already ingested items + if ingest_only and not skip_ingestion and item_id in ingested_item_ids: + console.print(f"\n[bold blue]Item {i}/{len(items)}[/bold blue] (ID: {item_id})") + console.print(" [yellow]⊘[/yellow] Skipping - already ingested") + continue + + # Then check fill status (results file) + if filln: + skip_item_ids = resume_complete_item_ids if rerun_invalid_existing else existing_item_ids + if item_id in skip_item_ids: + console.print(f"\n[bold blue]Item {i}/{len(items)}[/bold blue] (ID: {item_id})") + console.print(" [yellow]⊘[/yellow] Skipping - already has results in output file") + continue + + result = await self.process_single_item( + item, + item_agent_id, + i, + len(items), + thinking_budget, + max_tokens, + max_questions_per_item, + skip_ingestion, + question_semaphore, + eval_semaphore, + clear_this_agent, + wait_consolidation, + ingest_only, + skip_if_already_ingested=False, + force_reingest=force_reingest, + ) + + # Replace existing result or append new one + result_item_id = result["item_id"] + if result_item_id in existing_item_ids: + # Replace existing result + all_results = [r for r in all_results if r["item_id"] != result_item_id] + console.print(f" [cyan]↻[/cyan] Updating existing result for {result_item_id}") + all_results.append(result) + existing_item_ids.add(result_item_id) + + # Save results incrementally after each item + if output_path: + self._save_incremental_results(all_results, output_path) + + return all_results + + async def _process_items_parallel( + self, + items: List[Dict[str, Any]], + agent_id: str, + thinking_budget: int, + max_tokens: int, + skip_ingestion: bool, + max_questions_per_item: Optional[int], + question_semaphore: asyncio.Semaphore, + eval_semaphore: asyncio.Semaphore, + filln: bool, + max_concurrent_items: int, + output_path: Optional[Path] = None, + merge_with_existing: bool = False, + wait_consolidation: bool = False, + ingest_only: bool = False, + force_reingest: bool = False, + rerun_invalid_existing: bool = False, + ) -> List[Dict]: + """Process items in parallel (requires unique agent IDs per item).""" + # Load existing results if merge_with_existing is True + all_results = [] + existing_item_ids = set() + resume_complete_item_ids = set() + result_lock = asyncio.Lock() # Lock for thread-safe updates to all_results + + if merge_with_existing and output_path and output_path.exists(): + with open(output_path, "r", encoding="utf-8") as f: + existing_data = json.load(f) + if "item_results" in existing_data: + all_results = existing_data["item_results"] + existing_item_ids = {r["item_id"] for r in all_results} + resume_complete_item_ids = {r["item_id"] for r in all_results if _result_is_resume_complete(r)} + console.print(f"[cyan]Loaded {len(all_results)} existing results from {output_path}[/cyan]") + + # Pre-load durable banks when evaluating retained data or filling an + # ingest-only run. + ingested_item_ids = set() + if skip_ingestion or ingest_only: + console.print("[cyan]Checking which items are already ingested...[/cyan]") + for item in items: + item_id = self.dataset.get_item_id(item) + item_agent_id = f"{agent_id}_{item_id}" + if await self._agent_has_data(item_agent_id): + ingested_item_ids.add(item_id) + if ingested_item_ids: + console.print(f"[cyan]Found {len(ingested_item_ids)} items already ingested[/cyan]") + else: + console.print("[cyan]No items found with existing ingest data[/cyan]") + if skip_ingestion: + requested_item_ids = {self.dataset.get_item_id(item) for item in items} + missing_item_ids = sorted(requested_item_ids - ingested_item_ids) + if missing_item_ids: + preview = ", ".join(missing_item_ids[:10]) + suffix = " ..." if len(missing_item_ids) > 10 else "" + raise RuntimeError( + "Retrieval-only mode requires durable retained chunks for every selected item; " + f"missing {len(missing_item_ids)} bank(s): {preview}{suffix}" + ) + + # Create semaphore for item-level parallelism + item_semaphore = asyncio.Semaphore(max_concurrent_items) + + async def process_item_wrapper(i: int, item: Dict) -> Optional[Dict]: + """Wrapper to process a single item with semaphore control.""" + async with item_semaphore: + item_id = self.dataset.get_item_id(item) + item_agent_id = f"{agent_id}_{item_id}" + + # Only reuse mode requires data to exist before this item runs. + # A fresh parallel run must proceed to ingestion. + if skip_ingestion and item_id not in ingested_item_ids and not ingest_only: + raise RuntimeError(f"Retrieval-only bank disappeared before evaluation: {item_agent_id}") + + # For ingest_only mode, skip already ingested items + if ingest_only and not skip_ingestion and item_id in ingested_item_ids: + console.print(f"\n[bold blue]Item {i}/{len(items)}[/bold blue] (ID: {item_id})") + console.print(" [yellow]⊘[/yellow] Skipping - already ingested") + return None + + # Then check fill status (results file) + if filln: + skip_item_ids = resume_complete_item_ids if rerun_invalid_existing else existing_item_ids + if item_id in skip_item_ids: + console.print(f"\n[bold blue]Item {i}/{len(items)}[/bold blue] (ID: {item_id})") + console.print(" [yellow]⊘[/yellow] Skipping - already has results in output file") + return None + + # Process the item + result = await self.process_single_item( + item, + item_agent_id, + i, + len(items), + thinking_budget, + max_tokens, + max_questions_per_item, + skip_ingestion, + question_semaphore, + eval_semaphore, + clear_this_agent=True, + wait_consolidation=wait_consolidation, + ingest_only=ingest_only, + skip_if_already_ingested=False, + force_reingest=force_reingest, + ) + return result + + # Create all tasks + tasks = [process_item_wrapper(i, item) for i, item in enumerate(items, 1)] + + # Run in parallel and collect results incrementally + for completed_task in asyncio.as_completed(tasks): + result = await completed_task + if result is not None: + async with result_lock: + # Replace existing result or append new one + result_item_id = result["item_id"] + if result_item_id in existing_item_ids: + # Replace existing result + all_results = [r for r in all_results if r["item_id"] != result_item_id] + console.print(f" [cyan]↻[/cyan] Updating existing result for {result_item_id}") + all_results.append(result) + existing_item_ids.add(result_item_id) + + # Save results incrementally after each item completes + if output_path: + self._save_incremental_results(all_results, output_path) + + return all_results + + async def _run_two_phase( + self, + items: List[Dict[str, Any]], + agent_id: str, + thinking_budget: int, + max_tokens: int, + skip_ingestion: bool, + max_questions_per_item: Optional[int], + max_concurrent_questions: int, + eval_semaphore_size: int, + output_path: Optional[Path] = None, + merge_with_existing: bool = False, + ) -> Dict[str, Any]: + """ + Two-phase approach: ingest all data into single agent, then evaluate all questions. + + More realistic scenario where agent accumulates memories over time. + """ + # Phase 1: Ingestion + if not skip_ingestion: + # Calculate and display data statistics + console.print("\n[3] Analyzing data to be ingested...") + stats = self.calculate_data_stats(items) + console.print(f" [cyan]Total items:[/cyan] {stats['total_items']}") + console.print(f" [cyan]Total sessions:[/cyan] {stats['total_sessions']}") + console.print(f" [cyan]Total characters:[/cyan] {stats['total_chars']:,}") + console.print(f" [cyan]Avg session length:[/cyan] {stats['avg_session_length']:.0f} chars") + console.print( + f" [cyan]Session length range:[/cyan] {stats['min_session_length']}-{stats['max_session_length']} chars" + ) + + console.print(f"\n[4] Phase 1: Ingesting all data into agent '{agent_id}'...") + console.print(" [yellow]Clearing previous agent data...[/yellow]") + await self.memory.delete_bank(agent_id, request_context=RequestContext()) + console.print(" [green]✓[/green] Cleared agent data") + + # Apply template if configured + if self.template_path: + console.print(" [yellow]Applying bank template...[/yellow]") + await self.apply_template(agent_id, self.template_path) + console.print(" [green]✓[/green] Template applied") + + # Collect all sessions and send in one batch (with auto-chunking) + console.print(" [yellow]Collecting sessions from all items...[/yellow]") + all_sessions = [] + for item in items: + item_sessions = self.dataset.prepare_sessions_for_ingestion(item) + all_sessions.extend(item_sessions) + + console.print(f" [cyan]Collected {len(all_sessions)} sessions from {len(items)} items[/cyan]") + console.print(" [yellow]Ingesting in one batch (auto-chunks if needed)...[/yellow]") + + max_retries = 3 + retry_delay = 2.0 + ingest_success = False + + for attempt in range(max_retries): + try: + await self.memory.retain_batch_async( + bank_id=agent_id, contents=all_sessions, request_context=RequestContext() + ) + ingest_success = True + break + except ValueError as e: + error_msg = str(e) + if "Batch contains duplicate document_ids" in error_msg: + console.print( + f" [yellow]⚠[/yellow] Duplicate document_id error on attempt {attempt + 1}, " + f"deduplicating session batch..." + ) + unique_sessions = [] + seen_doc_ids = set() + for session in all_sessions: + doc_id = session.get("document_id") + if doc_id and doc_id in seen_doc_ids: + continue + unique_sessions.append(session) + if doc_id: + seen_doc_ids.add(doc_id) + all_sessions = unique_sessions + console.print(f" [cyan]Deduplicated to {len(all_sessions)} sessions[/cyan]") + if attempt == max_retries - 1: + raise + await asyncio.sleep(retry_delay) + retry_delay *= 1.5 + continue + else: + raise + except Exception as e: + console.print(f" [yellow]⚠[/yellow] Ingestion error: {e}") + if attempt == max_retries - 1: + raise + await asyncio.sleep(retry_delay) + retry_delay *= 1.5 + + if not ingest_success: + raise RuntimeError("Ingestion exhausted all retry attempts without completing") + console.print(f" [green]✓[/green] Ingested {len(all_sessions)} sessions from {len(items)} items") + else: + console.print("\n[3] Skipping ingestion (using existing data)") + + # Phase 2: Evaluation + console.print("\n[5] Phase 2: Evaluating all questions...") + + # Create semaphore for question processing + question_semaphore = asyncio.Semaphore(max_concurrent_questions) + eval_semaphore = asyncio.Semaphore(eval_semaphore_size) + + all_results = [] + for i, item in enumerate(items, 1): + item_id = self.dataset.get_item_id(item) + console.print(f"\n[bold blue]Item {i}/{len(items)}[/bold blue] (ID: {item_id})") + + # Get QA pairs + qa_pairs = self.dataset.get_qa_pairs(item) + console.print(f" Evaluating {len(qa_pairs)} QA pairs (parallel)...") + + qa_results = await self.evaluate_qa_task( + agent_id, + qa_pairs, + item_id, + thinking_budget, + max_tokens, + max_questions_per_item, + question_semaphore, + ) + + # Calculate metrics + metrics = await self.calculate_metrics(qa_results, eval_semaphore, item_id) + console.print( + f" [green]✓[/green] Accuracy: {metrics['accuracy']:.2f}% ({metrics['correct']}/{metrics['total']})" + ) + + all_results.append( + { + "item_id": item_id, + "metrics": metrics, + "num_sessions": -1, # Not tracked in two-phase mode + } + ) + + # Calculate overall metrics + total_correct = sum(r["metrics"]["correct"] for r in all_results) + total_questions = sum(r["metrics"]["total"] for r in all_results) + total_invalid = sum(r["metrics"].get("invalid", 0) for r in all_results) + total_valid = total_questions - total_invalid + overall_accuracy = (total_correct / total_questions * 100) if total_questions > 0 else 0 + valid_only_accuracy = (total_correct / total_valid * 100) if total_valid > 0 else 0 + + return { + "overall_accuracy": overall_accuracy, + "total_correct": total_correct, + "total_questions": total_questions, + "total_invalid": total_invalid, + "total_valid": total_valid, + "valid_only_accuracy": valid_only_accuracy, + "num_items": len(all_results), + "item_results": all_results, + } + + def display_results(self, results: Dict[str, Any]): + """Display benchmark results in a formatted table.""" + console.print("\n[bold green]✓ Benchmark Complete![/bold green]\n") + + # Display model configuration + if "model_config" in results: + config = results["model_config"] + console.print("[bold cyan]Model Configuration:[/bold cyan]") + console.print(f" HMS: {config['hms']['provider']}/{config['hms']['model']}") + if "retain" in config: + console.print(f" Retain: {config['retain']['provider']}/{config['retain']['model']}") + console.print( + f" Answer Generation: {config['answer_generation']['provider']}/{config['answer_generation']['model']}" + ) + console.print(f" LLM Judge: {config['judge']['provider']}/{config['judge']['model']}") + console.print() + + # Display results table + table = Table(title="Benchmark Results", box=box.ROUNDED) + table.add_column("Item ID", style="cyan") + table.add_column("Sessions", justify="right", style="yellow") + table.add_column("Questions", justify="right", style="blue") + table.add_column("Correct", justify="right", style="green") + table.add_column("Invalid", justify="right", style="red") + table.add_column("Accuracy", justify="right", style="magenta") + + for result in results["item_results"]: + metrics = result["metrics"] + invalid_count = metrics.get("invalid", 0) + invalid_str = str(invalid_count) if invalid_count > 0 else "-" + table.add_row( + result["item_id"], + str(result["num_sessions"]), + str(metrics["total"]), + str(metrics["correct"]), + invalid_str, + f"{metrics['accuracy']:.1f}%", + ) + + overall_invalid = results.get("total_invalid", 0) + invalid_str = str(overall_invalid) if overall_invalid > 0 else "-" + table.add_row( + "[bold]OVERALL[/bold]", + "-", + f"[bold]{results['total_questions']}[/bold]", + f"[bold]{results['total_correct']}[/bold]", + f"[bold]{invalid_str}[/bold]", + f"[bold]{results['overall_accuracy']:.1f}%[/bold]", + ) + + console.print(table) + + # Display note about invalid questions if any + if overall_invalid > 0: + console.print( + f"\n[yellow]Note: {overall_invalid} question(s) failed during processing and count as incorrect " + f"in the primary accuracy. Valid-only accuracy: {results.get('valid_only_accuracy', 0.0):.1f}%.[/yellow]" + ) + + def merge_results(self, new_results: Dict[str, Any], existing_results: Dict[str, Any]) -> Dict[str, Any]: + """ + Merge new results into existing results. + + Updates or adds item results, then recalculates overall metrics. + + Args: + new_results: New results to merge (typically from a specific item run) + existing_results: Existing results to merge into + + Returns: + Merged results with updated overall metrics + """ + # Start with existing item results + merged_item_results = existing_results.get("item_results", []) + + # Update or add new item results + for new_item in new_results["item_results"]: + item_id = new_item["item_id"] + + # Find if item already exists + found = False + for i, existing_item in enumerate(merged_item_results): + if existing_item["item_id"] == item_id: + # Replace existing item result + merged_item_results[i] = new_item + found = True + console.print(f" [yellow]→[/yellow] Updated results for item: {item_id}") + break + + if not found: + # Add new item result + merged_item_results.append(new_item) + console.print(f" [green]+[/green] Added results for item: {item_id}") + + # Recalculate overall metrics from all item results + total_correct = sum(r["metrics"]["correct"] for r in merged_item_results) + total_questions = sum(r["metrics"]["total"] for r in merged_item_results) + total_invalid = sum(r["metrics"].get("invalid", 0) for r in merged_item_results) + total_valid = total_questions - total_invalid + overall_accuracy = (total_correct / total_questions * 100) if total_questions > 0 else 0 + valid_only_accuracy = (total_correct / total_valid * 100) if total_valid > 0 else 0 + + return { + "overall_accuracy": overall_accuracy, + "total_correct": total_correct, + "total_questions": total_questions, + "total_invalid": total_invalid, + "total_valid": total_valid, + "valid_only_accuracy": valid_only_accuracy, + "num_items": len(merged_item_results), + "item_results": merged_item_results, + } + + def _save_incremental_results(self, all_results: List[Dict], output_path: Path): + """ + Save results incrementally to JSON file. + + Args: + all_results: Current list of all item results + output_path: Path to save results to + """ + ordered_results = sorted(all_results, key=lambda result: str(result.get("item_id", ""))) + total_correct = sum(r["metrics"]["correct"] for r in ordered_results) + total_questions = sum(r["metrics"]["total"] for r in ordered_results) + total_invalid = sum(r["metrics"].get("invalid", 0) for r in ordered_results) + total_valid = total_questions - total_invalid + overall_accuracy = (total_correct / total_questions * 100) if total_questions > 0 else 0 + valid_only_accuracy = (total_correct / total_valid * 100) if total_valid > 0 else 0 + + results_dict = { + "overall_accuracy": overall_accuracy, + "total_correct": total_correct, + "total_questions": total_questions, + "total_invalid": total_invalid, + "total_valid": total_valid, + "valid_only_accuracy": valid_only_accuracy, + "num_items": len(ordered_results), + "model_config": get_model_config(), + "item_results": ordered_results, + } + run_manifest = getattr(self, "_run_manifest", None) + if run_manifest is not None: + results_dict["run_manifest"] = run_manifest + + _write_json_atomic(results_dict, output_path) + + def save_results(self, results: Dict[str, Any], output_path: Path, merge_with_existing: bool = False): + """ + Save results to JSON file. + + Args: + results: Results to save + output_path: Path to save results to + merge_with_existing: If True, merge with existing results file if it exists + """ + if merge_with_existing and output_path.exists(): + # Load existing results + with open(output_path, "r", encoding="utf-8") as f: + existing_results = json.load(f) + + console.print(f"\n[cyan]Merging with existing results from {output_path}...[/cyan]") + results = self.merge_results(results, existing_results) + + if "item_results" in results: + results["item_results"] = sorted( + results["item_results"], + key=lambda result: str(result.get("item_id", "")), + ) + _write_json_atomic(results, output_path) + console.print(f"\n[green]✓[/green] Results saved to {output_path}") diff --git a/lab/evaluation/benchmarks/common/test_benchmark_runner.py b/lab/evaluation/benchmarks/common/test_benchmark_runner.py new file mode 100644 index 0000000..19c8d9a --- /dev/null +++ b/lab/evaluation/benchmarks/common/test_benchmark_runner.py @@ -0,0 +1,374 @@ +import asyncio +import json +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock + +import pytest + +from benchmarks.common.benchmark_runner import ( + BenchmarkRunner, + LLMAnswerEvaluator, + _embedding_runtime_config, + _endpoint_fingerprint, + _reranker_runtime_config, + _write_json_atomic, + get_model_config, +) + + +class _Dataset: + def get_item_id(self, item): + return item["id"] + + def prepare_sessions_for_ingestion(self, item): + return [{"content": item["id"], "document_id": f"document-{item['id']}"}] + + def get_qa_pairs(self, item): + return [{"question": item["id"], "answer": "answer", "category": "test"}] + + +def test_embedding_runtime_config_uses_provider_specific_identity(monkeypatch): + monkeypatch.setenv("HMS_API_EMBEDDINGS_PROVIDER", "cohere") + monkeypatch.setenv("HMS_API_EMBEDDINGS_COHERE_MODEL", "embed-v4.0") + monkeypatch.setenv("HMS_API_EMBEDDINGS_COHERE_BASE_URL", "https://cohere.example/v2") + monkeypatch.setenv("HMS_API_EMBEDDINGS_OPENAI_MODEL", "must-not-be-used") + monkeypatch.setenv("HMS_API_EMBEDDINGS_OPENAI_BASE_URL", "https://openai.example/v1") + + config = _embedding_runtime_config() + + assert config == { + "provider": "cohere", + "model": "embed-v4.0", + "fingerprint_policy": "strict", + "endpoint_fingerprint": _endpoint_fingerprint("https://cohere.example/v2"), + } + + +def test_reranker_runtime_config_uses_provider_specific_identity(monkeypatch): + monkeypatch.setenv("HMS_API_RERANKER_PROVIDER", "litellm-sdk") + monkeypatch.setenv("HMS_API_RERANKER_LITELLM_SDK_MODEL", "vendor/rerank-v2") + monkeypatch.setenv("HMS_API_RERANKER_LITELLM_SDK_API_BASE", "https://gateway.example/v1") + monkeypatch.setenv("HMS_API_RERANKER_LOCAL_MODEL", "must-not-be-used") + + config = _reranker_runtime_config() + + assert config == { + "provider": "litellm-sdk", + "model": "vendor/rerank-v2", + "endpoint_fingerprint": _endpoint_fingerprint("https://gateway.example/v1"), + } + + +def test_model_config_uses_retain_provider_default_when_only_provider_is_overridden(monkeypatch): + monkeypatch.setenv("HMS_API_LLM_PROVIDER", "groq") + monkeypatch.setenv("HMS_API_LLM_MODEL", "custom-memory-model") + monkeypatch.setenv("HMS_API_RETAIN_LLM_PROVIDER", "openai") + monkeypatch.delenv("HMS_API_RETAIN_LLM_MODEL", raising=False) + + config = get_model_config() + + assert config["hms"]["model"] == "custom-memory-model" + assert config["retain"]["provider"] == "openai" + assert config["retain"]["model"] == "gpt-4o-mini" + + +@pytest.mark.asyncio +async def test_fresh_parallel_run_processes_items_before_any_bank_exists(): + runner = object.__new__(BenchmarkRunner) + runner.dataset = _Dataset() + runner.template_path = None + runner._agent_has_data = AsyncMock(return_value=False) + + async def process_item(item, *args, **kwargs): + return { + "item_id": item["id"], + "metrics": {"correct": 1, "total": 1, "invalid": 0}, + "num_sessions": 1, + } + + runner.process_single_item = AsyncMock(side_effect=process_item) + runner._save_incremental_results = Mock() + + results = await runner._process_items_parallel( + items=[{"id": "a"}, {"id": "b"}], + agent_id="longmemeval", + thinking_budget=10, + max_tokens=100, + skip_ingestion=False, + max_questions_per_item=None, + question_semaphore=asyncio.Semaphore(2), + eval_semaphore=asyncio.Semaphore(2), + filln=False, + max_concurrent_items=2, + ) + + assert {result["item_id"] for result in results} == {"a", "b"} + assert runner.process_single_item.await_count == 2 + runner._agent_has_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_fresh_item_reingests_even_when_a_bank_already_exists(): + runner = object.__new__(BenchmarkRunner) + runner.dataset = _Dataset() + runner.template_path = None + runner.memory = Mock() + runner.memory.delete_bank = AsyncMock() + runner._agent_has_data = AsyncMock(return_value=True) + runner.ingest_conversation = AsyncMock(return_value=1) + runner._audit_durable_ingestion = AsyncMock( + return_value={ + "durable_documents": 1, + "expected_documents": 1, + "missing_documents": [], + "documents_without_chunks": [], + } + ) + runner.evaluate_qa_task = AsyncMock(return_value=[]) + runner.calculate_metrics = AsyncMock(return_value={"accuracy": 0.0, "correct": 0, "total": 0, "invalid": 0}) + + await runner.process_single_item( + {"id": "a"}, + "longmemeval_a", + 1, + 1, + 10, + 100, + None, + False, + asyncio.Semaphore(1), + asyncio.Semaphore(1), + clear_this_agent=True, + skip_if_already_ingested=False, + ) + + runner._agent_has_data.assert_not_awaited() + runner.memory.delete_bank.assert_awaited_once() + runner.ingest_conversation.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_ingestion_raises_after_the_last_retry(monkeypatch): + runner = object.__new__(BenchmarkRunner) + runner.dataset = _Dataset() + runner.memory = Mock() + runner.memory.retain_batch_async = AsyncMock(side_effect=RuntimeError("provider unavailable")) + monkeypatch.setattr(asyncio, "sleep", AsyncMock()) + + with pytest.raises(RuntimeError, match="provider unavailable"): + await runner.ingest_conversation({"id": "a"}, "longmemeval_a") + + assert runner.memory.retain_batch_async.await_count == 3 + + +@pytest.mark.asyncio +async def test_judge_provider_failure_is_recorded_as_invalid(): + evaluator = object.__new__(LLMAnswerEvaluator) + evaluator.llm_config = Mock() + evaluator.llm_config.call = AsyncMock(side_effect=RuntimeError("provider unavailable")) + runner = object.__new__(BenchmarkRunner) + runner.answer_evaluator = evaluator + + metrics = await runner.calculate_metrics( + [ + { + "question": "question", + "correct_answer": "answer", + "predicted_answer": "prediction", + "category": "multi-session", + "is_invalid": False, + "error": None, + } + ], + asyncio.Semaphore(1), + "a", + ) + + assert metrics["invalid"] == 1 + assert metrics["correct"] == 0 + assert metrics["detailed_results"][0]["is_invalid"] is True + + +@pytest.mark.asyncio +async def test_judge_concurrency_limit_is_shared_across_items(): + active = 0 + maximum_active = 0 + + async def fake_call(**kwargs): + nonlocal active, maximum_active + active += 1 + maximum_active = max(maximum_active, active) + await asyncio.sleep(0) + active -= 1 + return SimpleNamespace(correct=True, reasoning="ok") + + evaluator = object.__new__(LLMAnswerEvaluator) + evaluator.llm_config = Mock() + evaluator.llm_config.call = fake_call + runner = object.__new__(BenchmarkRunner) + runner.answer_evaluator = evaluator + shared_semaphore = asyncio.Semaphore(1) + result = { + "question": "question", + "correct_answer": "answer", + "predicted_answer": "answer", + "category": "multi-session", + "is_invalid": False, + "error": None, + } + + await asyncio.gather( + runner.calculate_metrics([dict(result)], shared_semaphore, "a"), + runner.calculate_metrics([dict(result)], shared_semaphore, "b"), + ) + + assert maximum_active == 1 + + +@pytest.mark.asyncio +async def test_resume_skips_valid_items_and_reruns_invalid_items(tmp_path: Path): + output_path = tmp_path / "results.json" + output_path.write_text( + json.dumps( + { + "item_results": [ + { + "item_id": "valid", + "metrics": { + "total": 1, + "correct": 1, + "invalid": 0, + "detailed_results": [ + { + "is_correct": True, + "is_invalid": False, + "error": None, + "predicted_answer": "answer", + "correctness_reasoning": "correct", + } + ], + }, + "num_sessions": 1, + }, + { + "item_id": "invalid", + "metrics": { + "total": 1, + "correct": 0, + "invalid": 1, + "detailed_results": [ + { + "is_correct": None, + "is_invalid": True, + "error": "provider unavailable", + } + ], + }, + "num_sessions": 1, + }, + ] + } + ), + encoding="utf-8", + ) + runner = object.__new__(BenchmarkRunner) + runner.dataset = _Dataset() + runner._agent_has_data = AsyncMock(return_value=False) + + async def process_item(item, *args, **kwargs): + return { + "item_id": item["id"], + "metrics": { + "total": 1, + "correct": 1, + "invalid": 0, + "detailed_results": [ + { + "is_correct": True, + "is_invalid": False, + "error": None, + "predicted_answer": "answer", + "correctness_reasoning": "correct", + } + ], + }, + "num_sessions": 1, + } + + runner.process_single_item = AsyncMock(side_effect=process_item) + runner._save_incremental_results = Mock() + + results = await runner._process_items_sequential( + items=[{"id": "valid"}, {"id": "invalid"}], + agent_id="longmemeval", + thinking_budget=10, + max_tokens=100, + skip_ingestion=False, + max_questions_per_item=None, + question_semaphore=asyncio.Semaphore(1), + eval_semaphore=asyncio.Semaphore(1), + clear_agent_per_item=True, + filln=True, + output_path=output_path, + merge_with_existing=True, + rerun_invalid_existing=True, + ) + + assert runner.process_single_item.await_count == 1 + assert runner.process_single_item.await_args.args[0]["id"] == "invalid" + assert {result["item_id"] for result in results} == {"valid", "invalid"} + + +@pytest.mark.asyncio +async def test_retrieval_only_fails_when_any_selected_bank_is_missing(): + runner = object.__new__(BenchmarkRunner) + runner.dataset = _Dataset() + runner._agent_has_data = AsyncMock(side_effect=[True, False]) + runner.process_single_item = AsyncMock() + + with pytest.raises(RuntimeError, match="missing 1 bank"): + await runner._process_items_parallel( + items=[{"id": "a"}, {"id": "b"}], + agent_id="longmemeval", + thinking_budget=10, + max_tokens=100, + skip_ingestion=True, + max_questions_per_item=None, + question_semaphore=asyncio.Semaphore(1), + eval_semaphore=asyncio.Semaphore(1), + filln=False, + max_concurrent_items=2, + ) + + runner.process_single_item.assert_not_awaited() + + +def test_atomic_json_write_replaces_the_complete_artifact(tmp_path: Path): + output_path = tmp_path / "result.json" + output_path.write_text('{"stale": true}\n', encoding="utf-8") + + _write_json_atomic({"item_results": [{"item_id": "a"}]}, output_path) + + assert json.loads(output_path.read_text(encoding="utf-8")) == {"item_results": [{"item_id": "a"}]} + assert not (tmp_path / ".result.json.tmp").exists() + + +def test_incremental_checkpoint_preserves_run_manifest(tmp_path: Path): + runner = object.__new__(BenchmarkRunner) + runner._run_manifest = {"artifact_schema_version": 1, "dataset": {"sha256": "abc"}} + output_path = tmp_path / "result.json" + + runner._save_incremental_results( + [ + { + "item_id": "a", + "metrics": {"correct": 1, "total": 1, "invalid": 0}, + "num_sessions": 1, + } + ], + output_path, + ) + + saved = json.loads(output_path.read_text(encoding="utf-8")) + assert saved["run_manifest"] == runner._run_manifest diff --git a/lab/evaluation/benchmarks/longmemeval/README.md b/lab/evaluation/benchmarks/longmemeval/README.md new file mode 100644 index 0000000..99154dc --- /dev/null +++ b/lab/evaluation/benchmarks/longmemeval/README.md @@ -0,0 +1,218 @@ +# LongMemEval reproduction + +This adapter runs one LongMemEval question per isolated HMS memory bank. Each +question follows four stages: + +1. Retain every haystack session. +2. Recall evidence for the question. +3. Generate an answer from the recalled evidence. +4. Judge the generated answer against the reference. + +The repository ships only the adapter. It downloads the canonical dataset on +first use and writes all runtime state to ignored local paths. + +## Requirements + +- Python 3.11 or newer +- [uv](https://docs.astral.sh/uv/) +- Docker, or another PostgreSQL 16 server with pgvector +- Model and embedding endpoints with enough quota for the selected run + +Install the evaluation environment from the repository root: + +```bash +uv sync --project lab/evaluation --extra test +``` + +## Start PostgreSQL + +The benchmark process runs on the host, so the database must listen on a +host-reachable address. The main Compose database uses the internal hostname +`postgres` and does not publish its port. Start an isolated benchmark database: + +```bash +docker run --name hms-longmemeval-db \ + --detach \ + --publish 127.0.0.1:5432:5432 \ + --env POSTGRES_USER=hms \ + --env POSTGRES_PASSWORD=hms_longmemeval_change_me \ + --env POSTGRES_DB=hms \ + --volume hms-longmemeval-postgres:/var/lib/postgresql/data \ + pgvector/pgvector@sha256:1d533553fefe4f12e5d80c7b80622ba0c382abb5758856f52983d8789179f0fb +``` + +Use `docker start hms-longmemeval-db` after the first run. + +## Configure model roles + +Create an ignored benchmark environment file: + +```bash +cp lab/evaluation/benchmarks/longmemeval/longmemeval.env.example .env.longmemeval +chmod 600 .env.longmemeval +``` + +Replace every `*_change_me` value. Keep separate settings for Retain, the core +memory model, Answer, Judge, and embeddings. A score is comparable only when +these roles use the same providers, model identifiers, endpoint behavior, and +prompt configuration. + +The sample uses `gpt-5-mini` for all language-model roles and +`text-embedding-3-small` for embeddings. These are examples, not a bundled +score claim. + +## Dataset pin + +The default run downloads this immutable artifact: + +- Dataset: `xiaowu0162/longmemeval-cleaned` +- File: `longmemeval_s_cleaned.json` +- Revision: `98d7416c24c778c2fee6e6f3006e7a073259d48f` +- Size: `277383467` bytes +- SHA-256: `d6f21ea9d60a0d56f34a05b609c79c88a451d2ae03597821ea3d5a9678c3a442` +- Source: + +The verified file is stored at +`.aaaDATA/longmemeval/longmemeval_s_cleaned.json`. The downloader writes a +temporary file, verifies its byte size and checksum, then installs it with an +atomic rename. The repository does not redistribute the dataset. Review the +upstream dataset card and terms before use. + +## Run a smoke test + +Start with one question and one request at a time: + +```bash +HMS_ENV_FILE=.env.longmemeval \ +HMS_MAX_INSTANCES=1 \ +HMS_RESULTS_FILENAME=longmemeval-smoke.json \ +bash .aaaSCRIPT/run_benchmark.sh +``` + +The command downloads and verifies the dataset, migrates the database, runs all +four stages, audits durable source chunks, and writes: + +- Results: `.aaaRESULT/longmemeval-smoke.json` +- Summary: `.aaaRESULT/longmemeval-smoke.md` +- Console log: `.aaaLOG/longmemeval_.log` +- Retain state: PostgreSQL banks named `longmemeval_` + +Fresh runs clear and rebuild every selected bank before evaluation. They also +refuse to overwrite an existing result filename. Delete an obsolete local +artifact intentionally or choose a new filename; use `--resume` only for a +compatible interrupted run. + +Retain state is not a standalone output file. Source documents live in +`documents`, source chunks live in `chunks`, and extracted facts live in +`memory_units`. The run fails when an expected document is missing or has no +durable chunk. A document with zero extracted facts remains valid because its +source chunk is still available to recall. + +## Run all 500 questions + +Choose positive concurrency values that stay within the provider and database +limits. Four parallel items provide a conservative starting point. The question +and Judge limits are shared across the entire run, not multiplied per item: + +```bash +HMS_ENV_FILE=.env.longmemeval \ +HMS_RESULTS_FILENAME=longmemeval-full.json \ +HMS_PARALLEL=4 \ +HMS_MAX_CONCURRENT_QUESTIONS=4 \ +HMS_EVAL_SEMAPHORE_SIZE=4 \ +bash .aaaSCRIPT/run_benchmark.sh +``` + +A complete artifact must report: + +- `num_items: 500` +- `total_questions: 500` +- `total_invalid: 0` +- the pinned dataset revision and checksum in `run_manifest` + +Provider failures count as incorrect in the primary accuracy. The artifact also +records valid-only accuracy as a diagnostic. A full run exits with an error +when it has missing questions or invalid judgments. + +## Resume an interrupted run + +Results are saved after each completed question with an atomic file replace. +Every checkpoint includes the dataset, pipeline, model-role, endpoint, and code +identity needed for compatibility checks. Resume with the same filename: + +```bash +HMS_ENV_FILE=.env.longmemeval \ +HMS_RESULTS_FILENAME=longmemeval-full.json \ +HMS_PARALLEL=4 \ +bash .aaaSCRIPT/run_benchmark.sh --resume +``` + +The runner skips valid completed question IDs. Invalid questions and items that +did not finish run again from Retain, including clearing a partially retained +bank. Resume requires an existing result file and never selects a new +timestamped artifact. + +Resume refuses to mix incompatible dataset hashes, pipeline settings, model +roles, endpoint fingerprints, database modes, Git revisions, or relevant dirty +source trees. Keep those settings unchanged for an interrupted run. Use a new +result filename when changing the experiment. + +## Reuse retained banks + +Use retrieval-only mode only after a successful Retain pass with the same +dataset, embedding fingerprint, database, and bank IDs: + +```bash +HMS_ENV_FILE=.env.longmemeval \ +HMS_RESULTS_FILENAME=longmemeval-recall-only.json \ +HMS_RETRIEVAL_ONLY=1 \ +bash .aaaSCRIPT/run_benchmark.sh +``` + +The runner requires durable chunks for every selected bank before retrieval-only +evaluation; missing banks fail the run instead of silently reducing the sample. +The strict embedding fingerprint policy rejects memories built in another +vector space. + +## Reproduction profiles + +`HMS_PIPELINE=ledger` enables a category-aware retrieval plan and the structured +evidence ledger. The plan reads LongMemEval `question_type` metadata. This is a +benchmark-conditioned profile, not label-free production recall. + +Use `HMS_PIPELINE=standard` to disable category-conditioned planning. The +`structured_source` context format remains independent and renders bounded +source bundles with retained source windows. + +`HMS_PIPELINE=self_evolution` is an optional diagnosis-derived experiment. Do +not compare it as an untuned main result unless the evaluation protocol +explicitly allows tuning from prior failed cases. + +## Local validation + +The focused test suite does not call external model providers: + +```bash +uv run --project lab/evaluation --extra test python -m pytest \ + lab/evaluation/benchmarks/common/test_benchmark_runner.py \ + lab/evaluation/benchmarks/longmemeval/test_evidence_bundles.py \ + lab/evaluation/benchmarks/longmemeval/test_source_backfill.py \ + lab/evaluation/benchmarks/longmemeval/test_source_context_integration.py \ + -q +``` + +Check the command line and launcher without starting a live run: + +```bash +bash -n .aaaSCRIPT/run_benchmark.sh +uv run --project lab/evaluation \ + python -m benchmarks.longmemeval.longmemeval_benchmark --help +``` + +## Privacy and cost + +Remote Retain and embedding endpoints receive the conversation haystacks. +Answer and Judge endpoints receive benchmark questions and generated content. +Use endpoints that satisfy your privacy requirements. A 500-question run can +consume substantial tokens and may take hours. Run the smoke test before +committing that cost. diff --git a/lab/evaluation/benchmarks/longmemeval/__init__.py b/lab/evaluation/benchmarks/longmemeval/__init__.py new file mode 100644 index 0000000..3bbe507 --- /dev/null +++ b/lab/evaluation/benchmarks/longmemeval/__init__.py @@ -0,0 +1 @@ +"""LongMemEval benchmark implementation.""" diff --git a/lab/evaluation/benchmarks/longmemeval/evidence_bundles.py b/lab/evaluation/benchmarks/longmemeval/evidence_bundles.py new file mode 100644 index 0000000..876d466 --- /dev/null +++ b/lab/evaluation/benchmarks/longmemeval/evidence_bundles.py @@ -0,0 +1,316 @@ +"""Source-centric evidence selection for LongMemEval answer prompts. + +The recall API returns facts, while several facts can point at the same raw +source chunk. Rendering that relationship as ``fact -> chunk`` for every fact +inflates the prompt and makes repeated wording look like independent events. +This module keeps the retrieval order and provenance, but renders each source +chunk once with a small, query-focused set of facts. +""" + +from __future__ import annotations + +import json +import re +from dataclasses import dataclass +from typing import Any, Mapping, Sequence + +_STOPWORDS = frozenset( + { + "about", + "after", + "again", + "before", + "between", + "current", + "currently", + "different", + "during", + "first", + "from", + "have", + "many", + "much", + "previous", + "recently", + "since", + "that", + "the", + "then", + "there", + "this", + "total", + "what", + "when", + "where", + "which", + "with", + } +) +_SIGNAL_RE = re.compile( + r"(?:\$\s?\d|\b\d+(?:[.,]\d+)?%?|\b(?:jan|feb|mar|apr|may|jun|jul|aug|sep|oct|nov|dec)\b|" + r"\b(?:before|after|earlier|later|first|last|current|latest|total|spent|cost|discount)\b)", + re.IGNORECASE, +) + + +@dataclass(frozen=True) +class EvidenceBundle: + """Facts sharing one source, in the order in which the source was found.""" + + source_key: str + document_id: str | None + chunk_id: str | None + first_rank: int + facts: tuple[dict[str, Any], ...] + chunk: Mapping[str, Any] | None = None + + +@dataclass(frozen=True) +class RenderedEvidence: + """Rendered prompt text and the source excerpts actually visible in it.""" + + text: str + covered_by_document: dict[str, tuple[str, ...]] + + +def _terms(query: str) -> frozenset[str]: + return frozenset( + token for token in re.findall(r"[A-Za-z][A-Za-z0-9_'-]{2,}", query.lower()) if token not in _STOPWORDS + ) + + +def _normalise_text(value: Any) -> str: + return re.sub(r"\W+", " ", str(value or "").lower()).strip() + + +def _compact_chunk_text(text: str, query_terms: frozenset[str], limit: int) -> str: + """Keep query-relevant turns when a raw chunk exceeds the display budget.""" + + text = " ".join(str(text or "").replace("<|endoftext|>", " ").split()) + if len(text) <= limit: + return text + + turns: list[str] = [] + try: + parsed = json.loads(text) + except (TypeError, ValueError): + parsed = None + if isinstance(parsed, dict): + parsed = parsed.get("messages") or parsed.get("turns") or [parsed] + if isinstance(parsed, list): + for turn in parsed: + if isinstance(turn, Mapping): + content = turn.get("content") + if content is None: + continue + role = str(turn.get("role") or "").strip().lower() + body = str(content) + turns.append(f"{role}: {body}" if role else body) + elif turn: + turns.append(str(turn)) + if not turns: + head = max(1, (limit - 9) // 2) + tail = max(1, limit - 9 - head) + return f"{text[:head].rstrip()} ... {text[-tail:].lstrip()}" + + scores = [sum(term in turn.lower() for term in query_terms) for turn in turns] + focus = max(range(len(turns)), key=lambda index: (scores[index], -index)) + focus_limit = min(limit, max(480, int(limit * 0.62))) + remaining = max(0, limit - focus_limit) + neighbours = [index for index in range(len(turns)) if index != focus] + neighbour_limit = remaining // len(neighbours) if neighbours else 0 + pieces: list[str] = [] + for index, turn in enumerate(turns): + per_turn = focus_limit if index == focus else neighbour_limit + if per_turn <= 0: + continue + clean = " ".join(turn.split()) + if len(clean) > per_turn: + original_length = len(clean) + lower = clean.lower() + positions = [ + match.start() for term in query_terms for match in re.finditer(rf"\b{re.escape(term)}\b", lower) + ] + # The last match often carries the updated state in a long + # assistant turn, while retaining the whole neighbouring turn + # keeps the user assertion visible. + center = max(positions, default=0) + window = max(1, per_turn - 6) + start = max(0, min(center - window // 2, original_length - window)) + clean = ("..." if start else "") + clean[start : start + window] + if start + window < original_length: + clean += "..." + pieces.append(clean) + return "\n".join(pieces)[:limit] + + +def _priority(fact: Mapping[str, Any], query_terms: frozenset[str], rank: int) -> tuple[int, int]: + text = str(fact.get("text") or "") + lower = text.lower() + overlap = sum(term in lower for term in query_terms) + signal = 1 if _SIGNAL_RE.search(text) else 0 + # Relevance wins within a source; the original rank is a stable tie-breaker. + return (overlap * 4 + signal * 2, -rank) + + +def build_evidence_bundles( + results: Sequence[Mapping[str, Any]], + chunks: Mapping[str, Mapping[str, Any]] | None, + query: str, + *, + max_bundles: int = 96, + max_facts_per_bundle: int = 2, +) -> list[EvidenceBundle]: + """Group retrieved facts by source chunk without changing candidate recall. + + ``results`` is assumed to be in retrieval order. Facts without a chunk + remain addressable by document (or their own id), so observations and + source-less rows are not silently discarded. The function is deterministic and + has no database or model dependency, which makes it suitable for both + benchmark prompts and unit tests. + """ + + if max_bundles <= 0 or max_facts_per_bundle <= 0: + return [] + + query_terms = _terms(query) + grouped: dict[str, list[tuple[int, Mapping[str, Any]]]] = {} + order: list[str] = [] + for rank, fact in enumerate(results, 1): + if not isinstance(fact, Mapping): + continue + chunk_id = str(fact.get("chunk_id") or "") or None + document_id = str(fact.get("document_id") or "") or None + source_key = chunk_id or document_id or f"fact:{fact.get('id') or rank}" + if source_key not in grouped: + grouped[source_key] = [] + order.append(source_key) + grouped[source_key].append((rank, fact)) + + bundles: list[EvidenceBundle] = [] + chunks = chunks or {} + for source_key in order[:max_bundles]: + rows = grouped[source_key] + first_rank = rows[0][0] + # Avoid repeated extraction of the same sentence within a chunk while + # keeping distinct numeric/event facts available to the answer model. + seen_text: set[str] = set() + ranked_rows = sorted( + rows, + key=lambda row: _priority(row[1], query_terms, row[0]), + reverse=True, + ) + selected: list[tuple[int, Mapping[str, Any]]] = [] + for rank, fact in ranked_rows: + text_key = _normalise_text(fact.get("text")) + if text_key and text_key in seen_text: + continue + if text_key: + seen_text.add(text_key) + selected.append((rank, fact)) + if len(selected) >= max_facts_per_bundle: + break + + if not selected: + continue + selected.sort(key=lambda row: row[0]) + first_fact = selected[0][1] + chunk_id = str(first_fact.get("chunk_id") or "") or None + document_id = str(first_fact.get("document_id") or "") or None + chunk = chunks.get(chunk_id) if chunk_id else None + bundles.append( + EvidenceBundle( + source_key=source_key, + document_id=document_id, + chunk_id=chunk_id, + first_rank=first_rank, + facts=tuple(dict(fact) for _, fact in selected), + chunk=chunk, + ) + ) + + bundles.sort(key=lambda bundle: bundle.first_rank) + return bundles + + +def render_evidence_with_coverage( + bundles: Sequence[EvidenceBundle], + *, + max_chunk_chars: int = 1400, + max_total_chars: int | None = None, + query: str = "", +) -> RenderedEvidence: + """Render bundles with one raw source chunk per bundle. + + ``recall`` can return hundreds of candidates. A source-centric layout + removes duplicate chunks, but it still needs an explicit prompt budget; + otherwise one chunk per candidate can exceed the answer model's useful + attention window. When ``max_total_chars`` is set, whole bundles are + retained in retrieval order until the bound is reached. + """ + + if not bundles: + return RenderedEvidence(text="", covered_by_document={}) + + header = "\n".join( + [ + "=== Source-Centric Evidence Bundles ===", + "Each source chunk is shown once. Facts under the same source are not independent events; deduplicate them before counting.", + ] + ) + if max_total_chars is not None and len(header) > max_total_chars: + return RenderedEvidence(text="", covered_by_document={}) + + rendered_text = header + covered_by_document: dict[str, list[str]] = {} + for index, bundle in enumerate(bundles, 1): + source = bundle.chunk_id or bundle.document_id or bundle.source_key + bundle_lines = [f"Bundle {index} (retrieval_rank={bundle.first_rank}, source={source}):"] + for fact in bundle.facts: + when_parts = [] + if fact.get("occurred_start"): + when_parts.append(f"occurred={fact['occurred_start']}") + if fact.get("occurred_end") and fact.get("occurred_end") != fact.get("occurred_start"): + when_parts.append(f"ended={fact['occurred_end']}") + if fact.get("mentioned_at"): + when_parts.append(f"mentioned={fact['mentioned_at']}") + when = ", ".join(when_parts) or "unknown time" + bundle_lines.append(f"- Fact ({fact.get('fact_type', 'unknown')}, {when}): {fact.get('text', '')}") + chunk_text = "" + if bundle.chunk: + chunk_text = _compact_chunk_text(str(bundle.chunk.get("chunk_text") or ""), _terms(query), max_chunk_chars) + if chunk_text: + bundle_lines.append(f'- Source chunk: "{chunk_text}"') + bundle_text = "\n".join(bundle_lines) + candidate_text = f"{rendered_text}\n\n{bundle_text}" + if max_total_chars is not None and len(candidate_text) > max_total_chars: + break + rendered_text = candidate_text + if bundle.document_id and chunk_text: + covered_by_document.setdefault(bundle.document_id, []).append(chunk_text) + + return RenderedEvidence( + text=rendered_text, + covered_by_document={document_id: tuple(excerpts) for document_id, excerpts in covered_by_document.items()}, + ) + + +def render_evidence_bundles( + bundles: Sequence[EvidenceBundle], + *, + max_chunk_chars: int = 1400, + max_total_chars: int | None = None, + query: str = "", +) -> str: + """Render evidence as text, preserving the original helper contract. + + Call :func:`render_evidence_with_coverage` when the caller also needs the + exact excerpts admitted by the prompt budget. + """ + + return render_evidence_with_coverage( + bundles, + max_chunk_chars=max_chunk_chars, + max_total_chars=max_total_chars, + query=query, + ).text diff --git a/lab/evaluation/benchmarks/longmemeval/longmemeval.env.example b/lab/evaluation/benchmarks/longmemeval/longmemeval.env.example new file mode 100644 index 0000000..243bc6c --- /dev/null +++ b/lab/evaluation/benchmarks/longmemeval/longmemeval.env.example @@ -0,0 +1,51 @@ +# Copy this file to .env.longmemeval and replace every *_change_me value. +# The launcher reads it when HMS_ENV_FILE=.env.longmemeval. + +# The benchmark process runs on the host. Use a host-reachable PostgreSQL URL. +HMS_BENCHMARK_DATABASE_URL=postgresql://hms:hms_longmemeval_change_me@127.0.0.1:5432/hms +HMS_API_DATABASE_BACKEND=postgresql +HMS_API_DATABASE_SCHEMA=public +HMS_API_VECTOR_EXTENSION=pgvector + +# Core model for recall organization and memory reasoning. +HMS_API_LLM_PROVIDER=openai +HMS_API_LLM_MODEL=gpt-5-mini +HMS_API_LLM_API_KEY=openai_key_change_me +HMS_API_LLM_BASE_URL=https://api.openai.com/v1 + +# Retain model for conversation ingestion and structured extraction. +HMS_API_RETAIN_LLM_PROVIDER=openai +HMS_API_RETAIN_LLM_MODEL=gpt-5-mini +HMS_API_RETAIN_LLM_API_KEY=openai_key_change_me +HMS_API_RETAIN_LLM_BASE_URL=https://api.openai.com/v1 +HMS_API_RETAIN_EMBEDDING_FAILURE_POLICY=raise + +# Answer model for the recalled evidence prompt. +HMS_API_ANSWER_LLM_PROVIDER=openai +HMS_API_ANSWER_LLM_MODEL=gpt-5-mini +HMS_API_ANSWER_LLM_API_KEY=openai_key_change_me +HMS_API_ANSWER_LLM_BASE_URL=https://api.openai.com/v1 + +# Judge model for generated answer scoring. +HMS_API_JUDGE_LLM_PROVIDER=openai +HMS_API_JUDGE_LLM_MODEL=gpt-5-mini +HMS_API_JUDGE_LLM_API_KEY=openai_key_change_me +HMS_API_JUDGE_LLM_BASE_URL=https://api.openai.com/v1 + +# Dense retrieval and deterministic reciprocal-rank fusion. +HMS_API_EMBEDDINGS_PROVIDER=openai +HMS_API_EMBEDDINGS_OPENAI_MODEL=text-embedding-3-small +HMS_API_EMBEDDINGS_OPENAI_API_KEY=openai_key_change_me +HMS_API_EMBEDDINGS_OPENAI_BASE_URL=https://api.openai.com/v1 +HMS_API_EMBEDDING_FINGERPRINT_POLICY=strict +HMS_API_RERANKER_PROVIDER=rrf + +# Public reproduction profile. +HMS_PIPELINE=ledger +HMS_CONTEXT_FORMAT=structured_source +HMS_RETRIEVAL_ONLY=0 + +# All concurrency values must be positive. Judge concurrency is global. +HMS_PARALLEL=1 +HMS_MAX_CONCURRENT_QUESTIONS=1 +HMS_EVAL_SEMAPHORE_SIZE=1 diff --git a/lab/evaluation/benchmarks/longmemeval/longmemeval_benchmark.py b/lab/evaluation/benchmarks/longmemeval/longmemeval_benchmark.py new file mode 100644 index 0000000..8de6212 --- /dev/null +++ b/lab/evaluation/benchmarks/longmemeval/longmemeval_benchmark.py @@ -0,0 +1,2639 @@ +""" +LongMemEval-specific benchmark implementations. + +Provides dataset, answer generator, and evaluator for the LongMemEval benchmark. +""" + +import asyncio +import hashlib +import json +import logging +import os +import platform +import re +import subprocess +import sys +import urllib.request +from datetime import date, datetime, timedelta, timezone +from pathlib import Path +from typing import Any, Dict, List, Mapping, Optional, Sequence, Tuple + +import pydantic +from hms_api.engine.llm_wrapper import LLMConfig +from hms_api.engine.schema import fq_table +from openai import AsyncOpenAI + +from benchmarks.common.benchmark_runner import ( + BenchmarkDataset, + BenchmarkRunner, + LLMAnswerEvaluator, + LLMAnswerGenerator, + RecallPlan, + get_model_config, +) +from benchmarks.longmemeval.evidence_bundles import ( + RenderedEvidence, + build_evidence_bundles, + render_evidence_with_coverage, +) +from benchmarks.longmemeval.source_backfill import ( + render_source_chunk_snippets, + render_source_snippets, + select_missing_chunk_snippets, + select_source_snippets, +) + +LONGMEMEVAL_DATASET_REVISION = "98d7416c24c778c2fee6e6f3006e7a073259d48f" +LONGMEMEVAL_DATASET_FILENAME = "longmemeval_s_cleaned.json" +LONGMEMEVAL_DATASET_URL = ( + "https://huggingface.co/datasets/xiaowu0162/longmemeval-cleaned/resolve/" + f"{LONGMEMEVAL_DATASET_REVISION}/{LONGMEMEVAL_DATASET_FILENAME}" +) +LONGMEMEVAL_DATASET_SHA256 = "d6f21ea9d60a0d56f34a05b609c79c88a451d2ae03597821ea3d5a9678c3a442" +LONGMEMEVAL_DATASET_SIZE = 277_383_467 +LONGMEMEVAL_EXPECTED_ITEMS = 500 +REPOSITORY_ROOT = Path(__file__).resolve().parents[4] +DEFAULT_DATASET_PATH = REPOSITORY_ROOT / ".aaaDATA" / "longmemeval" / LONGMEMEVAL_DATASET_FILENAME +REPRODUCIBILITY_SOURCE_PATHS = ( + ".aaaSCRIPT/run_benchmark.sh", + "core/dataplane/hms_api", + "core/dataplane/pyproject.toml", + "lab/evaluation/benchmarks", + "lab/evaluation/pyproject.toml", + "pyproject.toml", + "uv.lock", +) + + +def _git_value(*args: str) -> Optional[str]: + try: + completed = subprocess.run( + ["git", *args], + cwd=REPOSITORY_ROOT, + check=True, + capture_output=True, + text=True, + timeout=10, + ) + return completed.stdout.strip() + except (OSError, subprocess.SubprocessError): + return None + + +def _git_bytes(*args: str) -> Optional[bytes]: + """Return raw Git output without assuming that source diffs are text.""" + + try: + completed = subprocess.run( + ["git", *args], + cwd=REPOSITORY_ROOT, + check=True, + capture_output=True, + timeout=10, + ) + return completed.stdout + except (OSError, subprocess.SubprocessError): + return None + + +def _source_tree_fingerprint() -> Optional[str]: + """Identify relevant tracked and untracked source changes for safe resume.""" + + tracked_diff = _git_bytes("diff", "--binary", "HEAD", "--", *REPRODUCIBILITY_SOURCE_PATHS) + untracked_output = _git_bytes( + "ls-files", + "--others", + "--exclude-standard", + "-z", + "--", + *REPRODUCIBILITY_SOURCE_PATHS, + ) + if tracked_diff is None or untracked_output is None: + return None + + untracked_paths = sorted(path for path in untracked_output.split(b"\0") if path) + if not tracked_diff and not untracked_paths: + return None + + digest = hashlib.sha256() + digest.update(b"hms-longmemeval-source-tree-v1\0") + digest.update(tracked_diff) + for encoded_path in untracked_paths: + relative_path = encoded_path.decode("utf-8", errors="surrogateescape") + source_path = REPOSITORY_ROOT / relative_path + digest.update(b"\0path\0") + digest.update(encoded_path) + digest.update(b"\0content\0") + try: + digest.update(source_path.read_bytes()) + except OSError: + return None + return f"sha256:{digest.hexdigest()}" + + +def _manifest_dataset_reference(dataset_path: Path) -> str: + """Return a portable dataset reference without exposing checkout paths.""" + + resolved_path = dataset_path.expanduser().resolve() + try: + return resolved_path.relative_to(REPOSITORY_ROOT.resolve()).as_posix() + except ValueError: + return f"external:{resolved_path.name}" + + +def build_run_manifest( + *, + dataset_path: Path, + context_format: str, + max_instances: Optional[int], + max_instances_per_category: Optional[int], + max_questions_per_instance: Optional[int], + question_id: Optional[str], + index_range: Optional[str], + category: Optional[str], + max_concurrent_items: int, + max_concurrent_questions: int, + eval_semaphore_size: int, + thinking_budget: int, + max_tokens: int, + oracle_planner_v26: bool, + oracle_planner_v220: bool, + query_expansion_enabled: bool, + query_rewriting_strategy: str, + session_expansion_weight: float, + skip_ingestion: bool, + ingest_only: bool, + force_reingest: bool, +) -> Dict[str, Any]: + """Build a non-secret manifest sufficient to audit a benchmark artifact.""" + + canonical = dataset_path.resolve() == DEFAULT_DATASET_PATH.resolve() + dataset_sha256 = _sha256_file(dataset_path) + git_commit = _git_value("rev-parse", "HEAD") + git_status = _git_value("status", "--porcelain") + planner = "self_evolution" if oracle_planner_v220 else "ledger" if oracle_planner_v26 else "standard" + return { + "artifact_schema_version": 1, + "dataset": { + "path": _manifest_dataset_reference(dataset_path), + "sha256": dataset_sha256, + "revision": LONGMEMEVAL_DATASET_REVISION if canonical else None, + "expected_full_items": LONGMEMEVAL_EXPECTED_ITEMS if canonical else None, + }, + "pipeline": { + "stages": ["retain", "recall", "answer", "judge"], + "planner": planner, + "context_format": context_format, + "thinking_budget": thinking_budget, + "max_tokens": max_tokens, + "query_expansion_enabled": query_expansion_enabled, + "query_rewriting_strategy": query_rewriting_strategy if query_expansion_enabled else "noop", + "session_expansion_weight": session_expansion_weight, + }, + "concurrency": { + "items": max_concurrent_items, + "questions": max_concurrent_questions, + "judge": eval_semaphore_size, + }, + "selection": { + "max_instances": max_instances, + "max_instances_per_category": max_instances_per_category, + "max_questions_per_instance": max_questions_per_instance, + "question_id": question_id, + "index_range": index_range, + "category": category, + }, + "execution": { + "skip_ingestion": skip_ingestion, + "ingest_only": ingest_only, + "force_reingest": force_reingest, + }, + "database": { + "backend": os.getenv("HMS_API_DATABASE_BACKEND", "postgresql"), + "schema": os.getenv("HMS_API_DATABASE_SCHEMA", "public"), + "vector_extension": os.getenv("HMS_API_VECTOR_EXTENSION", "pgvector"), + }, + "runtime": { + "git_commit": git_commit, + "git_dirty": bool(git_status) if git_status is not None else None, + "source_tree_fingerprint": _source_tree_fingerprint(), + "python": sys.version.split()[0], + "platform": platform.platform(), + }, + } + + +def _artifact_compatibility_contract( + manifest: Mapping[str, Any], + model_config: Mapping[str, Any], +) -> Dict[str, Any]: + """Select fields that must match before item results can share one artifact.""" + + runtime = manifest.get("runtime", {}) + dataset = manifest.get("dataset", {}) + dataset_identity = ( + { + "sha256": dataset.get("sha256"), + } + if isinstance(dataset, Mapping) + else None + ) + return { + "artifact_schema_version": manifest.get("artifact_schema_version"), + "dataset": dataset_identity, + "pipeline": manifest.get("pipeline"), + "database": manifest.get("database"), + "git_commit": runtime.get("git_commit") if isinstance(runtime, Mapping) else None, + "source_tree_fingerprint": (runtime.get("source_tree_fingerprint") if isinstance(runtime, Mapping) else None), + "model_config": model_config, + } + + +def validate_artifact_compatibility( + output_path: Path, + *, + current_manifest: Mapping[str, Any], + current_model_config: Mapping[str, Any], +) -> Dict[str, Any]: + """Load an existing artifact and reject unsafe cross-run result mixing.""" + + try: + existing = json.loads(output_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise ValueError(f"Cannot resume from unreadable result artifact {output_path}: {exc}") from exc + if not isinstance(existing, dict): + raise ValueError(f"Cannot resume from non-object result artifact: {output_path}") + + existing_manifest = existing.get("run_manifest") + existing_model_config = existing.get("model_config") + if not isinstance(existing_manifest, Mapping) or not isinstance(existing_model_config, Mapping): + raise ValueError( + f"Cannot safely merge {output_path}: existing artifact has no compatible run_manifest/model_config" + ) + + existing_contract = _artifact_compatibility_contract(existing_manifest, existing_model_config) + current_contract = _artifact_compatibility_contract(current_manifest, current_model_config) + mismatches = [ + field_name + for field_name in current_contract + if existing_contract.get(field_name) != current_contract.get(field_name) + ] + if mismatches: + raise ValueError( + f"Cannot safely merge {output_path}: incompatible experiment fields: {', '.join(mismatches)}. " + "Use a new HMS_RESULTS_FILENAME for a different experiment." + ) + return existing + + +def validate_output_target( + output_path: Path, + *, + merge_with_existing: bool, + resume: bool, +) -> None: + """Protect fresh results and enforce an explicit existing target for resume.""" + + if resume and not output_path.exists(): + raise FileNotFoundError(f"--resume requires an existing results file. Not found: {output_path}") + if output_path.exists() and not merge_with_existing: + raise FileExistsError( + f"Refusing to overwrite existing result artifact {output_path}. " + "Choose a new HMS_RESULTS_FILENAME or use --resume for a compatible interrupted run." + ) + + +ORACLE_PLANNER_V1_WEIGHTS = { + "single-session-user": 0.25, + "single-session-assistant": 0.25, + "single-session-preference": 0.30, + "knowledge-update": 0.50, + "temporal-reasoning": 0.60, + "multi-session": 0.80, +} + + +SELF_EVOLUTION_PROFILES = { + "oracle_v220": { + "base": "oracle_v26", + "diagnosis_source": "v2.6 failed LongMemEval cases only", + "evolution_targets": [ + "count/total deduplication", + "relative-date lookup grounding", + "amount/difference missing-side calibration", + "current/previous state arbitration", + ], + "selection_rule": "Keep V2.6 retrieval and ledger as the base; add only diagnosis-derived pre-generation evidence controls.", + }, +} + + +def _v26_base_retrieval_plan( + question: str, + question_type: Optional[str], + question_date: Optional[datetime], +) -> RecallPlan: + """Internal V2.6 retrieval base: oracle weights plus multi-session query expansion and appendix.""" + del question, question_date + weight = ORACLE_PLANNER_V1_WEIGHTS.get(question_type or "", 0.30) + if question_type == "multi-session": + return RecallPlan( + name="longmemeval_v26_base_retrieval", + session_expansion_weight=weight, + query_rewriting_enabled=True, + query_rewriting_strategy_name="llm_driven", + evidence_appendix_mode="cross_session", + ) + return RecallPlan(name="longmemeval_v26_base_retrieval", session_expansion_weight=weight) + + +def longmemeval_oracle_planner_v26( + question: str, + question_type: Optional[str], + question_date: Optional[datetime], +) -> RecallPlan: + """V2.6: base retrieval plus a pre-generation Structured Evidence Ledger.""" + plan = _v26_base_retrieval_plan(question, question_type, question_date) + plan.name = "longmemeval_oracle_v26" + return plan + + +def longmemeval_oracle_planner_v220( + question: str, + question_type: Optional[str], + question_date: Optional[datetime], +) -> RecallPlan: + """V2.20: V2.6 plus diagnosis-driven self-evolution controls.""" + plan = _v26_base_retrieval_plan(question, question_type, question_date) + plan.name = "longmemeval_oracle_v220" + return plan + + +class LongMemEvalDataset(BenchmarkDataset): + """LongMemEval dataset implementation.""" + + def load(self, path: Path, max_items: Optional[int] = None) -> List[Dict[str, Any]]: + """Load LongMemEval dataset from JSON file.""" + with open(path, "r", encoding="utf-8") as f: + dataset = json.load(f) + + if not isinstance(dataset, list): + raise ValueError(f"LongMemEval dataset must be a JSON list: {path}") + + question_ids: set[str] = set() + for index, item in enumerate(dataset): + if not isinstance(item, dict): + raise ValueError(f"LongMemEval item {index} must be a JSON object") + question_id = item.get("question_id") + if not isinstance(question_id, str) or not question_id: + raise ValueError(f"LongMemEval item {index} has no valid question_id") + if question_id in question_ids: + raise ValueError(f"Duplicate LongMemEval question_id: {question_id}") + question_ids.add(question_id) + for field_name in ("haystack_sessions", "haystack_dates", "haystack_session_ids"): + if not isinstance(item.get(field_name), list): + raise ValueError(f"LongMemEval item {question_id!r} has no valid {field_name}") + lengths = { + len(item["haystack_sessions"]), + len(item["haystack_dates"]), + len(item["haystack_session_ids"]), + } + if len(lengths) != 1: + raise ValueError(f"LongMemEval item {question_id!r} has misaligned session arrays") + + if max_items is not None: + dataset = dataset[:max_items] + + return dataset + + def get_item_id(self, item: Dict) -> str: + """Get question ID from LongMemEval item.""" + return item.get("question_id", "unknown") + + def prepare_sessions_for_ingestion(self, item: Dict) -> List[Dict[str, Any]]: + """ + Prepare LongMemEval conversation sessions for batch ingestion. + + Returns: + List of session dicts with 'content', 'context', 'event_date' + """ + sessions = item.get("haystack_sessions", []) + dates = item.get("haystack_dates", []) + session_ids = item.get("haystack_session_ids", []) + + # Dataset loading validates that all three arrays are aligned. + if not (len(sessions) == len(dates) == len(session_ids)): + raise ValueError(f"LongMemEval item {item.get('question_id', 'unknown')!r} has misaligned session arrays") + + batch_contents = [] + seen_document_ids = {} + + # Process each session + for idx, (session_turns, date_str, session_id) in enumerate(zip(sessions, dates, session_ids)): + # Parse session date + session_date = self._parse_date(date_str) if date_str else datetime.now(timezone.utc) + + # Clean session turns - remove has_answer key if present + cleaned_turns = [] + for turn in session_turns: + if isinstance(turn, dict): + cleaned_turn = {k: v for k, v in turn.items() if k != "has_answer"} + cleaned_turns.append(cleaned_turn) + else: + cleaned_turns.append(turn) + + session_content = json.dumps(cleaned_turns) + question_id = item.get("question_id", "unknown") + base_document_id = f"{question_id}_{session_id}" + + unique_document_id = base_document_id + if base_document_id in seen_document_ids: + seen_document_ids[base_document_id] += 1 + unique_document_id = f"{base_document_id}_chunk{seen_document_ids[base_document_id]}" + else: + seen_document_ids[base_document_id] = 0 + + batch_contents.append( + { + "content": session_content, + "context": f"Session {unique_document_id} - you are the assistant in this conversation - happened on {session_date.strftime('%Y-%m-%d %H:%M:%S')} UTC.", + "event_date": session_date, + "document_id": unique_document_id, + } + ) + + return batch_contents + + def get_qa_pairs(self, item: Dict) -> List[Dict[str, Any]]: + """ + Extract QA pairs from LongMemEval item. + + For LongMemEval, each item has one question. + + Returns: + List with single QA dict with 'question', 'answer', 'category', 'question_date' + """ + # Parse question_date if available + question_date = None + if "question_date" in item: + question_date = self._parse_date(item["question_date"]) + + return [ + { + "question": item.get("question", ""), + "answer": item.get("answer", ""), + "category": item.get("question_type", "unknown"), + "question_date": question_date, + } + ] + + def _parse_date(self, date_str: str) -> datetime: + """Parse date string to datetime object.""" + try: + # LongMemEval format: "2023/05/20 (Sat) 02:21" + # Try to parse the main part before the day name + date_str_cleaned = date_str.split("(")[0].strip() if "(" in date_str else date_str + + # Try multiple formats + for fmt in ["%Y/%m/%d %H:%M", "%Y-%m-%d %H:%M:%S", "%Y-%m-%d", "%Y/%m/%d"]: + try: + dt = datetime.strptime(date_str_cleaned, fmt) + return dt.replace(tzinfo=timezone.utc) + except ValueError: + continue + + # Fallback: try ISO format + return datetime.fromisoformat(date_str.replace("Z", "+00:00")) + except Exception: + raise ValueError(f"Failed to parse date string: {date_str}") + + +class QuestionAnswer(pydantic.BaseModel): + answer: str + reasoning: Optional[str] = None + + +class LongMemEvalAnswerGenerator(LLMAnswerGenerator): + """LongMemEval-specific answer generator using configurable LLM provider.""" + + def __init__( + self, + context_format: str = "structured_source", + evidence_mode: Optional[str] = None, + ): + """Initialize with LLM configuration for answer generation. + + Args: + context_format: How to format the retrieved context. Options: + - "json": Raw JSON dump of recall_result (original behavior) + - "structured": Human-readable format with facts grouped with source chunks + - "structured_compact": Source-centric bundles with each source chunk rendered once + """ + # Uses HMS_API_ANSWER_LLM_* env vars with fallback to HMS_API_LLM_* for + # benchmark-specific LLM configuration (separate from the API config system). + self.llm_config = LLMConfig( + provider=os.getenv("HMS_API_ANSWER_LLM_PROVIDER", os.getenv("HMS_API_LLM_PROVIDER", "openai")), + api_key=os.getenv("HMS_API_ANSWER_LLM_API_KEY", os.getenv("HMS_API_LLM_API_KEY", "")), + base_url=os.getenv("HMS_API_ANSWER_LLM_BASE_URL", os.getenv("HMS_API_LLM_BASE_URL", "")), + model=os.getenv("HMS_API_ANSWER_LLM_MODEL", os.getenv("HMS_API_LLM_MODEL", "gpt-4o-mini")), + reasoning_effort="high", + ) + self.client = self.llm_config._client + self.model = self.llm_config.model + self.context_format = context_format + self.evidence_mode = evidence_mode + + def _format_context_json(self, recall_result: Dict[str, Any]) -> str: + """Original JSON dump format.""" + return json.dumps(recall_result) + + def _format_context_structured(self, recall_result: Dict[str, Any]) -> str: + """Human-readable format with facts grouped with their source chunks. + + Format: + Fact 1: [fact text] + When: [date] + Source: + "[chunk text]" + + --- + + Fact 2: ... + + === Entity Observations === + Entity: [name] + - [observation 1] + - [observation 2] + """ + results = recall_result.get("results", []) + chunks = recall_result.get("chunks", {}) + entities = recall_result.get("entities", {}) + + if not results and not entities: + return "No memories found." + + formatted_parts = [] + + for i, fact in enumerate(results, 1): + fact_text = fact.get("text", "") + fact_type = fact.get("fact_type", "unknown") + + # Extract temporal information + occurred_start = fact.get("occurred_start") + occurred_end = fact.get("occurred_end") + mentioned_at = fact.get("mentioned_at") + + # Build temporal string + when_parts = [] + if occurred_start: + when_parts.append(f"occurred: {occurred_start}") + if mentioned_at: + when_parts.append(f"mentioned: {mentioned_at}") + when_str = " | ".join(when_parts) if when_parts else "unknown" + + # Get the source chunk if available + chunk_id = fact.get("chunk_id") + chunk_text = None + if chunk_id and chunk_id in chunks: + chunk_info = chunks[chunk_id] + chunk_text = chunk_info.get("chunk_text", "") + + # Build the formatted fact entry + entry_parts = [f"Fact {i} ({fact_type}): {fact_text}", f"When: {when_str}"] + + # Add context field if present + context = fact.get("context") + if context: + entry_parts.append(f"Context: {context}") + + # Add source chunk + if chunk_text: + # Truncate very long chunks + if len(chunk_text) > 1000: + chunk_text = chunk_text[:1000] + "..." + entry_parts.append(f'Source chunk:\n "{chunk_text}"') + + formatted_parts.append("\n".join(entry_parts)) + + # Add entity observations section if present + if entities: + entity_parts = ["=== Entity Observations ==="] + for entity_name, entity_state in entities.items(): + observations = entity_state.get("observations", []) + if observations: + entity_parts.append(f"\nEntity: {entity_name}") + for obs in observations: + obs_text = obs.get("text", "") + entity_parts.append(f" - {obs_text}") + if len(entity_parts) > 1: # More than just the header + formatted_parts.append("\n".join(entity_parts)) + + return "\n\n---\n\n".join(formatted_parts) + + def _format_context_source_centric(self, recall_result: Dict[str, Any], query: str) -> RenderedEvidence: + """Render a bounded, source-centric view without changing retrieval. + + The per-fact structured formatter appends the same raw chunk to every + fact extracted from it. LongMemEval sessions commonly produce many + near-duplicate facts, so that layout spends context on repetition and + makes counts look larger than they are. Bundles preserve the selected + facts and provenance while making the unit of evidence the source + chunk. + """ + + results = recall_result.get("results", []) + chunks = recall_result.get("chunks", {}) or {} + entities = recall_result.get("entities", {}) or {} + # Keep the answer prompt bounded even when recall returns hundreds of + # facts. The retained-source block is rendered separately and placed + # near the answer instructions, so the model can inspect provenance + # without competing with an unbounded candidate dump. + bundles = build_evidence_bundles( + results, + chunks, + query, + max_bundles=96, + max_facts_per_bundle=2, + ) + rendered_bundles = render_evidence_with_coverage( + bundles, + max_chunk_chars=900, + max_total_chars=76_000, + query=query, + ) + parts = [rendered_bundles.text] if rendered_bundles.text else [] + + if entities: + entity_parts = ["=== Entity Observations ==="] + for entity_name, entity_state in entities.items(): + observations = entity_state.get("observations", []) + if observations: + entity_parts.append(f"\nEntity: {entity_name}") + for observation in observations: + text = observation.get("text", "") + if text: + entity_parts.append(f" - {text}") + if len(entity_parts) > 1: + parts.append("\n".join(entity_parts)) + + return RenderedEvidence( + text="\n\n---\n\n".join(part for part in parts if part), + covered_by_document=rendered_bundles.covered_by_document, + ) + + def _get_context_instructions(self) -> str: + """Get instructions for interpreting the context based on format.""" + if self.context_format in {"structured", "structured_compact", "structured_source"}: + context_guide = """**Understanding the Retrieved Context:** +The context contains memory facts extracted from previous conversations, each with its source chunk. + +1. **Fact**: A high-level summary/atomic fact (e.g., "User loves hiking in mountains") + - This is the searchable summary of what was stored + +2. **Source Chunk**: The actual raw conversation where the fact was extracted from + - **This is your primary source for detailed information** + - Look here for specifics, context, quotes, and evidence + - Prioritize information from chunks when facts seem ambiguous + +3. **Temporal Information**: + - "occurred": When the event actually happened + - "mentioned": When it was discussed in conversation + - Use this to understand the timeline and resolve conflicts (prefer more recent info) + +4. **Context**: Additional metadata about the conversation session + +5. **Retained Source-Document Evidence** (when present): Verbatim windows or + exact retained chunks loaded from the same bank. Use these as the primary + evidence for details missing from a summarized fact, especially explicit + user state updates, dates, amounts, and purchases. Do not dismiss a source + window merely because the extracted fact is incomplete. +""" + else: + context_guide = "" + + base_instructions = """ +**Date Calculations (CRITICAL - read carefully):** +- When calculating days between two dates: count the days from Date A to Date B as (B - A) +- Example: Jan 1 to Jan 8 = 7 days (not 8) +- "X days ago" from Question Date means: Question Date minus X days +- When a fact says "three weeks ago" on a certain mentioned date, that refers to 3 weeks before THAT mentioned date, NOT the question date +- Always convert relative times ("last Friday", "two weeks ago") to absolute dates BEFORE comparing +- Double-check your arithmetic - off-by-one errors are very common +- **Important**: Read questions carefully for time anchors. "How many days ago did X happen when Y happened?" asks for the time between X and Y, NOT between X and the question date + +**Handling Relative Times in Facts:** +- If a fact says "last Friday" or "two weeks ago", anchor it to the fact's "mentioned" date, NOT the question date +- First convert ALL relative references to absolute dates, then answer the question +- Show your date conversion work in your reasoning + +**Counting Questions (CRITICAL for "how many" questions):** +- **Scan ALL facts first** - go through every single fact before counting, don't stop early +- **List each item explicitly in your reasoning** before giving the count: "1. X, 2. Y, 3. Z = 3 total" +- **Check all facts and chunks** before giving your final count +- **Watch for duplicates**: The same item may appear in multiple facts. Deduplicate by checking if two facts refer to the same underlying item/event +- **Watch for different descriptions of same thing**: "Dr. Patel (ENT specialist)" and "the ENT specialist" might be the same doctor +- **Don't over-interpret**: A project you "completed" is different from a project you're "leading" +- **Don't double-count**: If the same charity event is mentioned in two conversations, it's still one event + +**Disambiguation Guidance (CRITICAL - many errors come from over-counting):** +- **Assume overlap by default**: If two facts describe similar events (same type, similar timeframe, similar details), assume they are the SAME event unless there's clear evidence they are different +- If a person has a name AND a role mentioned, check if they're the same person before counting separately +- If an amount is mentioned multiple times on different dates, check if it's the same event or different events +- When facts reference the same underlying event from different sessions, count it once +- **Check for aliases**: "my college roommate's wedding" and "Emily's wedding" might be the same event +- **Check for time period overlap**: Two "week-long breaks" mentioned in overlapping time periods are likely the same break +- **When in doubt, undercount**: It's better to miss a duplicate than to count the same thing twice + +**Question Interpretation (read carefully):** +- "How many X before Y?" - count only X that happened BEFORE Y, not Y itself +- "How many properties viewed before making an offer on Z?" - count OTHER properties, not Z +- "How many X in the last week/month?" - calculate the exact date range from the question date, then filter +- Pay attention to qualifiers like "before", "after", "initially", "currently", "in total" + +**When to Say "I Don't Know":** +- If the question asks about something not in the retrieved context, say "I don't have information about X" +- If comparing two things (e.g., "which happened first, X or Y?") but only one is mentioned, explicitly say the other is missing +- Don't guess or infer dates that aren't explicitly stated in the facts or chunks +- If you cannot find a specific piece of information after checking all facts and chunks, admit it +- **Partial knowledge is OK**: If asked about two things and you only have info on one, provide what you know and note what's missing (don't just say "I don't know") + +**For Recommendation/Preference Questions (tips, suggestions, advice):** +- **DO NOT invent specific recommendations** (no made-up product names, course names, paper titles, channel names, etc.) +- **DO mention specific brands/products the user ALREADY uses** from the context +- Describe WHAT KIND of recommendation the user would prefer, referencing their existing tools/brands +- Keep answers concise - focus on key preferences (brand, quality level, specific interests) not exhaustive category lists +- First scan ALL facts for user's existing tools, brands, stated preferences + +**How to Answer:** +1. Scan ALL facts to find relevant memories - don't stop after finding a few +2. **Read the source chunks carefully** - they contain the actual details you need +3. Convert all relative times to absolute dates +4. Use temporal information to understand when things happened +5. Synthesize information from multiple facts if needed +6. If facts conflict, prefer more recent information +7. Double-check any date calculations before answering +8. **For counting questions ("how many")**: First list each unique item in your reasoning (1. X, 2. Y, 3. Z...), then count them +9. **For recommendations**: Reference the user's existing tools, experiences, or preferences explicitly +""" + return context_guide + base_instructions + + async def _format_source_document_backfill( + self, + question: str, + recall_result: Dict[str, Any], + bank_id: Optional[str], + question_type: Optional[str] = None, + rendered_coverage: Mapping[str, Sequence[str]] | None = None, + ) -> str: + """Recover bounded source windows for retrieved documents. + + Retain stores the original transcript separately from extracted facts. + Reading only documents referenced by the current candidate set keeps + this a scoped provenance lookup, rather than a bank-wide search. + """ + + if not bank_id: + return "" + database_url = os.environ.get("HMS_API_DATABASE_URL") + if not database_url: + return "" + + document_order: list[str] = [] + seen: set[str] = set() + for fact in recall_result.get("results", []): + document_id = str(fact.get("document_id") or "") + if document_id and document_id not in seen: + seen.add(document_id) + document_order.append(document_id) + if not document_order: + return "" + + returned_chunks = recall_result.get("chunks", {}) or {} + missing_chunk_ids: list[str] = [] + chunk_retrieval_rank: dict[str, int] = {} + for rank, fact in enumerate(recall_result.get("results", []), 1): + chunk_id = str(fact.get("chunk_id") or "") + if chunk_id and chunk_id not in returned_chunks and chunk_id not in missing_chunk_ids and rank <= 64: + missing_chunk_ids.append(chunk_id) + chunk_retrieval_rank[chunk_id] = rank + + try: + import asyncpg + + conn = await asyncpg.connect(database_url) + try: + rows = await conn.fetch( + f""" + SELECT id, original_text + FROM {fq_table("documents")} + WHERE bank_id = $1 AND id = ANY($2::text[]) + """, + bank_id, + document_order[:128], + ) + missing_chunk_rows = [] + if missing_chunk_ids: + missing_chunk_rows = await conn.fetch( + f""" + SELECT chunk_id, document_id, chunk_index, chunk_text + FROM {fq_table("chunks")} + WHERE bank_id = $1 AND chunk_id = ANY($2::text[]) + ORDER BY array_position($2::text[], chunk_id) + """, + bank_id, + missing_chunk_ids[:128], + ) + finally: + await conn.close() + except Exception as exc: + # Source backfill is an optional evidence aid. A database or JSON + # failure must leave ordinary recall and judging valid. + logging.warning("Source-document evidence backfill unavailable for %s: %s", bank_id, exc) + return "" + + documents = {str(row["id"]): row["original_text"] for row in rows if row["original_text"]} + exact_chunk_records = [] + for row in missing_chunk_rows: + record = dict(row) + record["_retrieval_rank"] = chunk_retrieval_rank.get(str(row["chunk_id"]), 10_000) + exact_chunk_records.append(record) + exact_chunks = select_missing_chunk_snippets( + exact_chunk_records, + question, + ) + base_coverage = { + str(document_id): list(excerpts) + for document_id, excerpts in (rendered_coverage or {}).items() + if document_id + } + admitted_exact_chunks = list(exact_chunks) + while True: + covered_chunks_by_document = { + document_id: list(excerpts) for document_id, excerpts in base_coverage.items() + } + for snippet in admitted_exact_chunks: + if snippet.document_id and snippet.text: + covered_chunks_by_document.setdefault(snippet.document_id, []).append(snippet.text) + + snippets = select_source_snippets( + documents, + question, + document_order, + max_documents=16, + max_snippets=12, + max_snippets_per_document=2, + max_chars_per_snippet=900, + max_total_chars=max(0, 9_000 - sum(len(snippet.text) for snippet in admitted_exact_chunks)), + covered_chunks=covered_chunks_by_document, + ) + while True: + blocks = [ + render_source_chunk_snippets(admitted_exact_chunks), + render_source_snippets(snippets), + ] + source_block = "\n\n".join(block for block in blocks if block) + if len(source_block) <= 14_000: + return source_block + if not snippets: + break + snippets.pop() + + if not admitted_exact_chunks: + return "" + # Re-run source selection when an exact entry cannot fit. Its + # original document turn is no longer considered covered. + admitted_exact_chunks.pop() + + def _needs_structured_evidence_ledger(self, question: str, question_type: Optional[str]) -> bool: + if self.evidence_mode not in {"oracle_v26", "oracle_v220"}: + return False + eligible_types = {"multi-session", "temporal-reasoning", "knowledge-update"} + if question_type not in eligible_types: + return False + + question_lower = question.lower() + markers = ( + "after", + "ago", + "amount", + "before", + "between", + "cashback", + "compared", + "cost", + "current", + "currently", + "date", + "days", + "difference", + "earliest", + "first", + "higher", + "hours", + "how long", + "how many", + "how much", + "in total", + "initially", + "latest", + "less", + "lower", + "months", + "more", + "most", + "order", + "percentage", + "previous", + "recently", + "since", + "spent", + "total", + "weeks", + "years", + ) + return any(marker in question_lower for marker in markers) + + @staticmethod + def _needs_v26_self_evolution_controller(question: str, question_type: Optional[str]) -> bool: + if question_type not in {"multi-session", "temporal-reasoning", "knowledge-update"}: + return False + question_lower = question.lower() + markers = ( + "ago", + "before", + "after", + "current", + "currently", + "difference", + "first", + "how many", + "how much", + "in total", + "initially", + "latest", + "previous", + "save", + "spent", + "total", + ) + return any(marker in question_lower for marker in markers) + + @staticmethod + def _v26_self_evolution_controller() -> str: + return """ +**V2.20 V2.6 Self-Evolution Controller:** +This controller was derived from V2.6 failure analysis. It does not replace the V2.6 evidence ledger; it tells you how to use that ledger more carefully. +- Count/total questions: enumerate unique real user events/items before giving the count. Do not count recommendations, options, generic background facts, plans, or duplicate extractions of the same event. If one required category is missing, answer with the known side plus "not enough information"; do not collapse missing evidence to 0. +- Amount/difference questions: compute only from amounts that are explicitly present for both sides requested by the question. If one side's amount is missing, say which side is missing instead of using generic ranges. +- Relative-date lookup questions: resolve the relative date from the question date, then prefer facts and source snippets whose event date or source text matches that resolved day. If the answer is described rather than named, return the full descriptive phrase. +- Current/previous/update questions: prefer the latest explicit state for "current" questions and the older explicit state for "previous/before" questions. Do not add old and new state values unless the question asks for a cumulative lifetime total. +- Final answer contract: start with the direct value/name/date/insufficient-information statement. Put caveats in reasoning after the direct answer. +""" + + @staticmethod + def _compact_text(text: Any, max_chars: int) -> str: + compact = " ".join(str(text or "").replace("<|endoftext|>", " ").split()) + if len(compact) > max_chars: + compact = compact[: max_chars - 3].rstrip() + "..." + return compact + + @staticmethod + def _bound_context_block(text: str, max_chars: int) -> str: + """Bound one evidence block while retaining both its head and tail. + + Evidence sections often contain a high-ranked header at the front and + exact source rows at the end. Keeping both sides is more useful than + a left-only cut, which can silently remove the answer-bearing span. + """ + + if len(text) <= max_chars: + return text + if max_chars <= 80: + return text[:max_chars] + head = max_chars // 2 + tail = max_chars - head - 40 + return f"{text[:head].rstrip()}\n... [evidence block bounded] ...\n{text[-tail:].lstrip()}" + + @staticmethod + def _content_terms(question: str) -> set[str]: + stopwords = { + "about", + "after", + "again", + "before", + "between", + "current", + "currently", + "different", + "during", + "first", + "from", + "have", + "many", + "much", + "previous", + "recently", + "since", + "that", + "the", + "then", + "there", + "this", + "total", + "what", + "when", + "where", + "which", + "with", + } + return { + token for token in re.findall(r"[A-Za-z][A-Za-z0-9_'-]{2,}", question.lower()) if token not in stopwords + } + + def _format_structured_evidence_ledger(self, question: str, recall_result: Dict[str, Any]) -> str: + results = recall_result.get("results", []) + chunks = recall_result.get("chunks", {}) + question_terms = self._content_terms(question) + signal_re = re.compile( + r"(\$?\d+(?:[.,]\d+)?%?|jan|feb|mar|apr|may|jun|jul|aug|sep|oct|nov|dec|" + r"monday|tuesday|wednesday|thursday|friday|saturday|sunday|" + r"today|yesterday|tomorrow|last|next|ago|week|month|year|day|hour|" + r"before|after|first|earlier|later|previous|current|latest|total|spent|cost|discount|cashback)", + re.IGNORECASE, + ) + + ledger_rows = [] + seen = set() + for fact in results[:180]: + text = self._compact_text(fact.get("text", ""), 360) + if not text: + continue + text_lower = text.lower() + term_overlap = sum(1 for term in question_terms if term in text_lower) + has_signal = bool(signal_re.search(text)) + if not has_signal and term_overlap < 2: + continue + + dedupe_key = re.sub(r"\W+", " ", text_lower).strip()[:180] + doc_id = str(fact.get("document_id") or "") + dedupe_key = f"{doc_id}:{dedupe_key}" + if dedupe_key in seen: + continue + seen.add(dedupe_key) + + ledger_rows.append( + { + "score": (3 if has_signal else 0) + term_overlap, + "doc": doc_id, + "type": fact.get("fact_type"), + "occurred": fact.get("occurred_start") or fact.get("occurred_end") or "-", + "mentioned": fact.get("mentioned_at") or "-", + "chunk_id": fact.get("chunk_id"), + "text": text, + } + ) + if len(ledger_rows) >= 70: + break + + ledger_rows.sort(key=lambda row: row["score"], reverse=True) + ledger_rows = ledger_rows[:45] + if not ledger_rows: + return "" + + lines = [ + "=== V2.6 Structured Evidence Ledger ===", + "Use this as a checklist for count/sum/date/order/update questions. It is extracted from the retrieved context; do not use it as new evidence beyond the facts and source chunks. Deduplicate repeated mentions of the same event before counting.", + "", + "Candidate facts:", + ] + used_chunks = [] + seen_chunks = set() + for idx, row in enumerate(ledger_rows, 1): + lines.append( + f"{idx}. occurred={row['occurred']} | mentioned={row['mentioned']} | " + f"doc={row['doc']} | type={row['type']} | {row['text']}" + ) + chunk_id = row.get("chunk_id") + if chunk_id and chunk_id in chunks and chunk_id not in seen_chunks: + seen_chunks.add(chunk_id) + used_chunks.append(chunk_id) + + if used_chunks: + lines.extend(["", "Raw source snippets for ledger rows:"]) + for idx, chunk_id in enumerate(used_chunks[:18], 1): + chunk_info = chunks.get(chunk_id) or {} + chunk_text = self._compact_text(chunk_info.get("chunk_text", ""), 650) + if chunk_text: + lines.append(f"{idx}. chunk={chunk_id} | {chunk_text}") + + return "\n".join(lines) + + @staticmethod + def _sort_date_key(value: Any) -> Tuple[int, str]: + if not value or value == "-": + return (1, "") + return (0, str(value)) + + @staticmethod + def _date_from_value(value: Any) -> Optional[date]: + if not value or value == "-": + return None + text = str(value).strip() + if not text: + return None + try: + return datetime.fromisoformat(text.replace("Z", "+00:00")).date() + except ValueError: + match = re.search(r"\d{4}-\d{2}-\d{2}", text) + if not match: + return None + try: + return datetime.strptime(match.group(0), "%Y-%m-%d").date() + except ValueError: + return None + + @staticmethod + def _resolved_relative_dates(question: str, question_date: Optional[datetime]) -> List[Tuple[str, date]]: + if question_date is None: + return [] + question_lower = question.lower() + base_date = question_date.date() + resolved: List[Tuple[str, date]] = [] + + for match in re.finditer(r"\b(\d+)\s+days?\s+ago\b", question_lower): + days = int(match.group(1)) + resolved.append((match.group(0), base_date - timedelta(days=days))) + + for match in re.finditer(r"\b(\d+)\s+weeks?\s+ago\b", question_lower): + weeks = int(match.group(1)) + resolved.append((match.group(0), base_date - timedelta(days=7 * weeks))) + + word_numbers = { + "one": 1, + "two": 2, + "three": 3, + "four": 4, + "five": 5, + "six": 6, + "seven": 7, + "eight": 8, + "nine": 9, + "ten": 10, + } + for word, value in word_numbers.items(): + if re.search(rf"\b{word}\s+days?\s+ago\b", question_lower): + resolved.append((f"{word} days ago", base_date - timedelta(days=value))) + if re.search(rf"\b{word}\s+weeks?\s+ago\b", question_lower): + resolved.append((f"{word} weeks ago", base_date - timedelta(days=7 * value))) + + if re.search(r"\byesterday\b", question_lower): + resolved.append(("yesterday", base_date - timedelta(days=1))) + + weekdays = { + "monday": 0, + "tuesday": 1, + "wednesday": 2, + "thursday": 3, + "friday": 4, + "saturday": 5, + "sunday": 6, + } + for weekday, target_idx in weekdays.items(): + if re.search(rf"\blast\s+{weekday}\b", question_lower): + days_back = (base_date.weekday() - target_idx) % 7 + if days_back == 0: + days_back = 7 + resolved.append((f"last {weekday}", base_date - timedelta(days=days_back))) + + deduped: List[Tuple[str, date]] = [] + seen = set() + for label, date_value in resolved: + key = (label, date_value.isoformat()) + if key not in seen: + seen.add(key) + deduped.append((label, date_value)) + return deduped + + @staticmethod + def _is_relative_date_lookup_question(question: str, question_type: Optional[str]) -> bool: + if question_type != "temporal-reasoning": + return False + question_lower = question.lower() + has_relative_date = bool( + re.search( + r"\b(\d+|one|two|three|four|five|six|seven|eight|nine|ten)\s+(days?|weeks?)\s+ago\b", + question_lower, + ) + or re.search(r"\blast\s+(monday|tuesday|wednesday|thursday|friday|saturday|sunday)\b", question_lower) + or re.search(r"\byesterday\b", question_lower) + ) + if not has_relative_date: + return False + + comparison_markers = ( + "first", + "earliest", + "latest", + "before", + "after", + "between", + "compared", + "higher", + "lower", + "more", + "less", + "total", + "how many", + "how much", + "how long", + ) + if any(marker in question_lower for marker in comparison_markers): + return False + + lookup_markers = ( + "what ", + "which ", + "who ", + "where ", + "from whom", + "by whom", + "did i buy", + "did i get", + "did i receive", + "did i purchase", + "started to listen", + ) + return any(marker in question_lower for marker in lookup_markers) + + def _format_resolved_date_evidence_block( + self, + question: str, + question_date: Optional[datetime], + question_type: Optional[str], + recall_result: Dict[str, Any], + ) -> str: + if self.evidence_mode != "oracle_v220" or question_type != "temporal-reasoning": + return "" + if not self._is_relative_date_lookup_question(question, question_type): + return "" + + resolved_dates = self._resolved_relative_dates(question, question_date) + if not resolved_dates: + return "" + + results = recall_result.get("results", []) + chunks = recall_result.get("chunks", {}) + rows = [] + seen = set() + target_dates = {date_value for _, date_value in resolved_dates} + for rank, fact in enumerate(results[:220], 1): + fact_dates = { + self._date_from_value(fact.get("occurred_start")), + self._date_from_value(fact.get("occurred_end")), + self._date_from_value(fact.get("mentioned_at")), + } + fact_dates.discard(None) + matched_dates = sorted(date_value for date_value in fact_dates if date_value in target_dates) + if not matched_dates: + continue + + text = self._compact_text(fact.get("text", ""), 340) + if not text: + continue + doc_id = str(fact.get("document_id") or "") + dedupe_text = re.sub(r"\W+", " ", text.lower()).strip()[:160] + dedupe_key = f"{doc_id}:{dedupe_text}" + if dedupe_key in seen: + continue + seen.add(dedupe_key) + rows.append( + { + "rank": rank, + "matched": ", ".join(date_value.isoformat() for date_value in matched_dates), + "doc": doc_id, + "type": fact.get("fact_type"), + "occurred": fact.get("occurred_start") or fact.get("occurred_end") or "-", + "mentioned": fact.get("mentioned_at") or "-", + "chunk_id": fact.get("chunk_id"), + "text": text, + } + ) + if len(rows) >= 24: + break + + if not rows: + return "" + + title = "=== V2.20 V2.6 Self-Evolved Relative-Date Evidence Block ===" + lines = [ + title, + "Relative date targets resolved from the question:", + ] + for label, date_value in resolved_dates: + lines.append(f"- {label} => {date_value.isoformat()}") + lines.extend( + [ + "Facts retrieved for those exact dates. Use them as same-day evidence, even when the surface noun in the question differs from the extracted fact wording.", + "", + "Same-date candidate facts:", + ] + ) + + used_chunks = [] + seen_chunks = set() + for idx, row in enumerate(rows, 1): + lines.append( + f"{idx}. target_date={row['matched']} | occurred={row['occurred']} | " + f"mentioned={row['mentioned']} | doc={row['doc']} | type={row['type']} | {row['text']}" + ) + chunk_id = row.get("chunk_id") + if chunk_id and chunk_id in chunks and chunk_id not in seen_chunks: + seen_chunks.add(chunk_id) + used_chunks.append(chunk_id) + + if used_chunks: + lines.extend(["", "Raw source snippets for same-date facts:"]) + for idx, chunk_id in enumerate(used_chunks[:10], 1): + chunk_info = chunks.get(chunk_id) or {} + chunk_text = self._compact_text(chunk_info.get("chunk_text", ""), 560) + if chunk_text: + lines.append(f"{idx}. chunk={chunk_id} | {chunk_text}") + + return "\n".join(lines) + + async def _format_resolved_date_memory_backfill( + self, + question: str, + question_date: Optional[datetime], + question_type: Optional[str], + bank_id: Optional[str], + ) -> str: + if self.evidence_mode != "oracle_v220" or question_type != "temporal-reasoning" or not bank_id: + return "" + if not self._is_relative_date_lookup_question(question, question_type): + return "" + + resolved_dates = self._resolved_relative_dates(question, question_date) + if not resolved_dates: + return "" + + database_url = os.environ.get("HMS_API_DATABASE_URL") + if not database_url: + return "" + + try: + import asyncpg + + conn = await asyncpg.connect(database_url) + try: + target_dates = [date_value for _, date_value in resolved_dates] + rows = await conn.fetch( + f""" + SELECT id, document_id, chunk_id, text, fact_type, + occurred_start, occurred_end, mentioned_at + FROM {fq_table("memory_units")} + WHERE bank_id = $1 + AND ( + occurred_start::date = ANY($2::date[]) + OR occurred_end::date = ANY($2::date[]) + OR mentioned_at::date = ANY($2::date[]) + ) + ORDER BY mentioned_at, document_id + LIMIT 180 + """, + bank_id, + target_dates, + ) + if not rows: + return "" + + target_date_set = set(target_dates) + question_terms = self._content_terms(question) + acquisition_re = re.compile(r"\b(acquired|got|bought|purchased|received|picked up|ordered)\b", re.I) + + def row_score(row: Any) -> Tuple[int, str]: + text = str(row["text"] or "") + text_lower = text.lower() + fact_dates = { + self._date_from_value(row["occurred_start"]), + self._date_from_value(row["occurred_end"]), + self._date_from_value(row["mentioned_at"]), + } + fact_dates.discard(None) + exact_occurred = ( + self._date_from_value(row["occurred_start"]) in target_date_set + or self._date_from_value(row["occurred_end"]) in target_date_set + ) + mentioned_match = self._date_from_value(row["mentioned_at"]) in target_date_set + term_overlap = sum(1 for term in question_terms if term in text_lower) + acquisition = bool(acquisition_re.search(text)) + score = ( + (12 if exact_occurred else 0) + + (2 if mentioned_match else 0) + + 3 * term_overlap + + (6 if acquisition else 0) + ) + return (-score, str(row["mentioned_at"] or ""), str(row["document_id"] or "")) + + rows = sorted(rows, key=row_score)[:36] + + chunk_ids = [row["chunk_id"] for row in rows if row["chunk_id"]] + chunk_rows = [] + if chunk_ids: + chunk_rows = await conn.fetch( + f""" + SELECT chunk_id, chunk_text + FROM {fq_table("chunks")} + WHERE bank_id = $1 AND chunk_id = ANY($2::text[]) + LIMIT 16 + """, + bank_id, + chunk_ids[:16], + ) + chunk_text_by_id = {row["chunk_id"]: row["chunk_text"] for row in chunk_rows} + finally: + await conn.close() + except Exception: + return "" + + title = "=== V2.20 V2.6 Self-Evolved Exact-Date Memory Backfill ===" + lines = [ + title, + "Memory units directly loaded from the fixed memory bank for the resolved relative-date target. This is not new extraction; it is a date-constrained evidence backfill from stored memories.", + "For acquisition questions, stored wording such as got, acquired, received, or bought should be treated as candidate acquisition evidence for the item named in the memory.", + "Resolved targets:", + ] + for label, date_value in resolved_dates: + lines.append(f"- {label} => {date_value.isoformat()}") + lines.extend(["", "Exact-date stored memories:"]) + + used_chunks = [] + seen_chunks = set() + for idx, row in enumerate(rows, 1): + text = self._compact_text(row["text"], 340) + lines.append( + f"{idx}. occurred={row['occurred_start'] or row['occurred_end'] or '-'} | " + f"mentioned={row['mentioned_at'] or '-'} | doc={row['document_id']} | " + f"type={row['fact_type']} | {text}" + ) + chunk_id = row["chunk_id"] + if chunk_id and chunk_id in chunk_text_by_id and chunk_id not in seen_chunks: + seen_chunks.add(chunk_id) + used_chunks.append(chunk_id) + + if used_chunks: + lines.extend(["", "Raw source snippets for exact-date backfill:"]) + for idx, chunk_id in enumerate(used_chunks[:8], 1): + chunk_text = self._compact_text(chunk_text_by_id.get(chunk_id, ""), 560) + if chunk_text: + lines.append(f"{idx}. chunk={chunk_id} | {chunk_text}") + + return "\n".join(lines) + + async def generate_answer( + self, + question: str, + recall_result: Dict[str, Any], + question_date: Optional[datetime] = None, + question_type: Optional[str] = None, + bank_id: Optional[str] = None, + ) -> Tuple[str, str, Optional[List[Dict[str, Any]]]]: + """ + Generate answer from retrieved memories using Groq. + + Args: + question: The question text + recall_result: Full RecallResult dict containing results, entities, chunks, and trace + question_date: Date when the question was asked (for temporal context) + question_type: Question category (e.g., 'single-session-user', 'multi-session-assistant') + + Returns: + Tuple of (answer, reasoning, None) + - None indicates to use the memories from recall_result + """ + # Format context based on selected mode + rendered_coverage: Mapping[str, Sequence[str]] = {} + if self.context_format in {"structured_compact", "structured_source"}: + rendered_context = self._format_context_source_centric(recall_result, question) + context = rendered_context.text + rendered_coverage = rendered_context.covered_by_document + elif self.context_format == "structured": + context = self._format_context_structured(recall_result) + else: + context = self._format_context_json(recall_result) + + source_block = "" + if self.context_format == "structured_source": + source_block = await self._format_source_document_backfill( + question, + recall_result, + bank_id, + question_type, + rendered_coverage, + ) + + if self._needs_structured_evidence_ledger(question, question_type): + ledger = self._format_structured_evidence_ledger(question, recall_result) + if ledger: + context = f"{context}\n\n{self._bound_context_block(ledger, 20_000)}" + if self.evidence_mode == "oracle_v220": + backfill_block = await self._format_resolved_date_memory_backfill( + question, + question_date, + question_type, + bank_id, + ) + if backfill_block: + context = f"{backfill_block}\n\n{context}" + date_block = self._format_resolved_date_evidence_block( + question, + question_date, + question_type, + recall_result, + ) + if date_block: + context = f"{context}\n\n{date_block}" + + # Put retained-source evidence after the compact fact/ledger sections, + # immediately before the answer prompt. This keeps provenance visible + # when the upstream provider applies an input-token window and makes + # the source text the last evidence the model reads. + if source_block: + context = f"{context}\n\n{source_block}" + + context_instructions = self._get_context_instructions() + if ( + self.evidence_mode == "oracle_v220" + and self._needs_structured_evidence_ledger(question, question_type) + and self._needs_v26_self_evolution_controller(question, question_type) + ): + context_instructions = f"{context_instructions}{self._v26_self_evolution_controller()}" + + # Format question date if provided + formatted_question_date = question_date.strftime("%Y-%m-%d %H:%M:%S UTC") if question_date else "Not specified" + + # Use LLM to generate answer. Provider and schema failures must + # propagate to the benchmark runner so the question is recorded as + # invalid rather than silently scored as an ordinary answer. + answer_obj = await self.llm_config.call( + messages=[ + { + "role": "user", + "content": f"""You are a helpful assistant that must answer user questions based on the previous conversations. + +{context_instructions}**Answer Guidelines:** +1. Start by scanning retrieved context to understand the facts and events that happened and the timeline. +2. Reason about all the memories and find the right answer, considering the most recent memory as an update of the current facts. +3. If you have 2 possible answers, just say both. + +In general the answer must be comprehensive and plenty of details from the retrieved context. + +For quantitative/counting questions ("how many..."): First list each unique item in your reasoning (1. X, 2. Y, 3. Z...), scanning ALL facts, then count them for your answer. +If questions asks a location (where...?) make sure to include the location name. +For recommendation questions ("can you recommend...", "suggest...", "any tips..."): DO NOT give actual recommendations. Instead, describe what KIND the user would prefer based on their context. Example answer format: "The user would prefer recommendations for [category] that focus on [their interest]. They would not prefer [what to avoid based on context]." +For questions asking for help or instructions, consider the users' recent memories and previous interactions with the assistant to understand their current situation better (recent purchases, specific product models used..) +For specific number/value questions, use the context to understand what is the most up-to-date number based on recency, but also include the reasoning (in the answer) on previous possible values and why you think are less relevant. +For open questions, include as much details as possible from different sources that are relevant. +For questions where a specific entity/role is mentioned and it's different from your memory, just say the truth, don't make up anything just to fulfill the question. For example, if the question is about a specific sport, you should consider if the memories and the question are about the same sport. (e.g. american football vs soccer, shows vs podcasts) +For comparative questions, say you don't know the answer if you don't have information about both sides. (or more sides) +For questions related to time/date, carefully review the question date and the memories date to correctly answer the question. +For questions related to time/date calculation (e.g. How many days passed between X and Y?), carefully review the memories date to correctly answer the question and only provide an answer if you have information about both X and Y, otherwise say it's not possible to calculate and why. + +Consider assistant's previous actions (e.g., bookings, reminders) as impactful to the user experiences. + + +Question: {question} +Question Date: {formatted_question_date} + +Retrieved Context: +{context} + + +Answer: +""", + } + ], + response_format=QuestionAnswer, + scope="memory", + max_completion_tokens=32768, + ) + reasoning_text = answer_obj.reasoning or "" + if reasoning_text: + reasoning_text = reasoning_text + " " + reasoning_text += f"(question date: {formatted_question_date})" + return answer_obj.answer, reasoning_text, None + + +def validate_runtime_options( + *, + max_instances: Optional[int], + max_instances_per_category: Optional[int], + max_questions_per_instance: Optional[int], + max_concurrent_items: int, + max_concurrent_questions: int, + eval_semaphore_size: int, + thinking_budget: int, + max_tokens: int, +) -> None: + """Reject zero and negative limits before any provider or database work.""" + + optional_positive = { + "max_instances": max_instances, + "max_instances_per_category": max_instances_per_category, + "max_questions_per_instance": max_questions_per_instance, + } + required_positive = { + "max_concurrent_items": max_concurrent_items, + "max_concurrent_questions": max_concurrent_questions, + "eval_semaphore_size": eval_semaphore_size, + "thinking_budget": thinking_budget, + "max_tokens": max_tokens, + } + for name, value in optional_positive.items(): + if value is not None and value < 1: + raise ValueError(f"{name} must be a positive integer, got {value}") + for name, value in required_positive.items(): + if value < 1: + raise ValueError(f"{name} must be a positive integer, got {value}") + + +def resolve_dataset_path(dataset_path: Optional[str]) -> Path: + """Resolve the selected dataset and enforce the canonical artifact pin.""" + + if dataset_path is None: + resolved = DEFAULT_DATASET_PATH + if not resolved.exists(): + download_dataset(resolved) + else: + resolved = Path(dataset_path).expanduser() + if not resolved.is_absolute(): + resolved = (REPOSITORY_ROOT / resolved).resolve() + if not resolved.exists(): + raise FileNotFoundError(f"Custom dataset not found: {resolved}") + + if resolved.resolve() == DEFAULT_DATASET_PATH.resolve(): + validate_canonical_dataset(resolved) + return resolved + + +def resolve_output_path(results_dir: Optional[str], results_filename: str) -> Path: + """Resolve a benchmark result artifact path from CLI settings.""" + + if results_dir: + result_root = Path(results_dir).expanduser() + if not result_root.is_absolute(): + result_root = (REPOSITORY_ROOT / result_root).resolve() + return result_root / results_filename + return Path(__file__).parent / "results" / results_filename + + +async def run_benchmark( + max_instances: int = None, + max_instances_per_category: int = None, + max_questions_per_instance: int = None, + thinking_budget: int = 500, + max_tokens: int = 8192, + skip_ingestion: bool = False, + filln: bool = False, + question_id: str = None, + index_range: str = None, + only_failed: bool = False, + only_invalid: bool = False, + only_ingested: bool = False, + category: str = None, + max_concurrent_items: int = 1, + results_filename: str = "benchmark_results.json", + results_dir: str = None, + context_format: str = "structured_source", + source_results: str = None, + ingest_only: bool = False, + force_reingest: bool = False, + max_concurrent_questions: int = 10, + eval_semaphore_size: int = 10, + dataset_path: Optional[str] = None, + query_expansion_enabled: bool = False, + query_rewriting_strategy: str = "llm_based", + session_expansion_weight: float = 0.3, + oracle_planner_v26: bool = False, + oracle_planner_v220: bool = False, + resume: bool = False, +): + """ + Run the LongMemEval benchmark. + + Args: + max_instances: Maximum number of instances to evaluate (None for all). Mutually exclusive with max_instances_per_category and category. + max_instances_per_category: Maximum number of instances per category (None for all). Mutually exclusive with max_instances and category. + max_questions_per_instance: Maximum questions per instance (for testing) + thinking_budget: Thinking budget for spreading activation search + max_tokens: Maximum tokens to retrieve from memories + skip_ingestion: Whether to skip ingestion and use existing data + filln: If True, only process question IDs not already present in the result artifact + question_id: Optional question ID to filter (e.g., 'sample-question-a'). Useful with --skip-ingestion. + only_failed: If True, only run questions that were previously marked as incorrect (is_correct=False) + only_invalid: If True, only run questions that were previously marked as invalid (is_invalid=True) + only_ingested: If True, only run questions whose memory bank already exists (has been ingested) + category: Optional category to filter questions (e.g., 'single-session-user', 'multi-session', 'temporal-reasoning'). Mutually exclusive with max_instances and max_instances_per_category. + max_concurrent_items: Maximum number of instances to process in parallel (default: 1 for sequential) + results_filename: Filename for results (default: benchmark_results.json). + results_dir: Optional directory for results. If None, defaults to results/ relative to script location. + context_format: How to format context for answer generation. "json" (raw JSON) or "structured" (human-readable with facts+chunks). + source_results: Source results file to read failed/invalid questions from (for --only-failed/--only-invalid). Defaults to benchmark_results.json. + ingest_only: Only ingest, skip evaluation + force_reingest: If True, always re-ingest even if data already exists (for re-running after fixing ingestion issues) + dataset_path: Optional custom dataset path. If None, uses the default dataset. + oracle_planner_v26: If True, use the V2.6 Structured Evidence Ledger. + oracle_planner_v220: If True, use pure v2.6 plus diagnosis-driven self-evolution controls. + resume: If True, merge with existing results and skip already processed items (default: False) + """ + from rich.console import Console + + console = Console() + validate_runtime_options( + max_instances=max_instances, + max_instances_per_category=max_instances_per_category, + max_questions_per_instance=max_questions_per_instance, + max_concurrent_items=max_concurrent_items, + max_concurrent_questions=max_concurrent_questions, + eval_semaphore_size=eval_semaphore_size, + thinking_budget=thinking_budget, + max_tokens=max_tokens, + ) + + # Validate mutually exclusive arguments + # --max-instances-per-category can't be combined with --max-instances or --category + # But --category CAN be combined with --max-instances (to limit questions within a category) + if max_instances_per_category is not None and (max_instances is not None or category is not None): + raise ValueError("--max-questions-per-category cannot be combined with --max-instances or --category") + + # Validate --only-ingested can't be combined with other dataset filters + if only_ingested: + incompatible_flags = [] + if only_failed: + incompatible_flags.append("--only-failed") + if only_invalid: + incompatible_flags.append("--only-invalid") + if category is not None: + incompatible_flags.append("--category") + if question_id is not None: + incompatible_flags.append("--question-id") + if max_instances_per_category is not None: + incompatible_flags.append("--max-instances-per-category") + + if incompatible_flags: + raise ValueError(f"--only-ingested cannot be combined with: {', '.join(incompatible_flags)}") + + # Determine dataset path. The canonical path is validated even when it is + # supplied explicitly through --dataset-path. + explicit_dataset_path = dataset_path is not None + dataset_path = resolve_dataset_path(dataset_path) + if explicit_dataset_path: + console.print(f"[cyan]Using custom dataset: {dataset_path}[/cyan]") + else: + console.print( + f"[cyan]Using pinned LongMemEval dataset {LONGMEMEVAL_DATASET_REVISION[:12]} " + f"(sha256:{LONGMEMEVAL_DATASET_SHA256[:12]}…)[/cyan]" + ) + + # Initialize components + dataset = LongMemEvalDataset() + + # Start with all items or load from dataset + original_dataset_items = None + filtered_items = None + + # Handle max_instances_per_category (aka max_questions_per_category) + if max_instances_per_category: + console.print(f"[cyan]Limiting to {max_instances_per_category} questions per category[/cyan]") + if original_dataset_items is None: + original_dataset_items = dataset.load(dataset_path, max_items=None) + + # Group by category and take max_instances_per_category from each + from collections import defaultdict + + category_items = defaultdict(list) + for item in original_dataset_items: + cat = item.get("question_type", "unknown") + category_items[cat].append(item) + + # Take up to max_instances_per_category from each category + filtered_items = [] + for cat, items in sorted(category_items.items()): + limited = items[:max_instances_per_category] + filtered_items.extend(limited) + console.print(f" [green]{cat}:[/green] {len(limited)} questions (of {len(items)} available)") + + console.print(f"[green]Total: {len(filtered_items)} questions across {len(category_items)} categories[/green]") + + # Load previous results if filtering for failed/invalid questions + failed_question_ids = set() + invalid_question_ids = set() + if only_failed or only_invalid: + # Use source_results if specified, otherwise default to benchmark_results.json + source_file = source_results if source_results else "benchmark_results.json" + results_path = Path(source_file).expanduser() + if not results_path.is_absolute(): + source_root = Path(results_dir) if results_dir else Path(__file__).parent / "results" + if not source_root.is_absolute(): + source_root = (REPOSITORY_ROOT / source_root).resolve() + results_path = source_root / results_path + if not results_path.exists(): + raise FileNotFoundError( + f"Cannot use --only-failed or --only-invalid; results file not found: {results_path}" + ) + + console.print(f"[cyan]Reading failed/invalid questions from: {source_file}[/cyan]") + with open(results_path, "r", encoding="utf-8") as f: + previous_results = json.load(f) + + # Extract question IDs that failed or are invalid + for item_result in previous_results.get("item_results", []): + item_id = item_result["item_id"] + for detail in item_result["metrics"].get("detailed_results", []): + if only_failed and detail.get("is_correct") == False and not detail.get("is_invalid", False): + failed_question_ids.add(item_id) + if only_invalid and detail.get("is_invalid", False): + invalid_question_ids.add(item_id) + + if only_failed: + console.print( + f"[cyan]Filtering to {len(failed_question_ids)} questions that failed (is_correct=False)[/cyan]" + ) + if only_invalid: + console.print( + f"[cyan]Filtering to {len(invalid_question_ids)} questions that were invalid (is_invalid=True)[/cyan]" + ) + + # Filter dataset by category if specified + if category: + console.print(f"[cyan]Filtering questions by category: {category}[/cyan]") + if original_dataset_items is None: + # Load full dataset without max_instances limit for filtering + original_dataset_items = dataset.load(dataset_path, max_items=None) + + filtered_items = [item for item in original_dataset_items if item.get("question_type") == category] + + if not filtered_items: + available_categories = set(item.get("question_type", "unknown") for item in original_dataset_items) + raise ValueError( + f"No questions found for category {category!r}. Available categories: {sorted(available_categories)}" + ) + + total_found = len(filtered_items) + will_run = min(total_found, max_instances) if max_instances else total_found + if max_instances and total_found > max_instances: + console.print( + f"[green]Found {total_found} questions for category '{category}' (will run {will_run} due to --max-instances)[/green]" + ) + else: + console.print(f"[green]Found {total_found} questions for category '{category}'[/green]") + + # Filter dataset by question_id(s) if specified + if question_id is not None: + # Parse comma-separated question IDs + target_ids = set(q.strip() for q in question_id.split(",") if q.strip()) + if not target_ids: + raise ValueError("--question-id did not contain a valid question ID") + + console.print(f"[cyan]Filtering to {len(target_ids)} question ID(s): {sorted(target_ids)}[/cyan]") + + # Load original items if not already loaded + if original_dataset_items is None: + original_dataset_items = dataset.load(dataset_path, max_items=None) + + # If we already have filtered_items from category filtering, filter those + # Otherwise start with all items + items_to_filter = filtered_items if filtered_items is not None else original_dataset_items + filtered_items = [item for item in items_to_filter if dataset.get_item_id(item) in target_ids] + + total_found = len(filtered_items) + missing_ids = target_ids - {dataset.get_item_id(item) for item in filtered_items} + if missing_ids: + console.print( + f"[yellow]Warning: {len(missing_ids)} question ID(s) not found in dataset: {sorted(missing_ids)}[/yellow]" + ) + if total_found == 0: + raise ValueError(f"None of the requested question IDs were found: {sorted(target_ids)}") + will_run = min(total_found, max_instances) if max_instances else total_found + if max_instances and total_found > max_instances: + console.print( + f"[green]Found {total_found} items matching question ID(s) (will run {will_run} due to --max-instances)[/green]" + ) + else: + console.print(f"[green]Found {total_found} items matching question ID(s)[/green]") + + # Filter dataset by index range if specified + if index_range: + try: + start_idx, end_idx = map(int, index_range.split(",")) + start_idx = max(1, start_idx) # Ensure 1-indexed, min 1 + end_idx = max(start_idx, end_idx) + + console.print(f"[cyan]Filtering to item index range: {start_idx}-{end_idx} (1-indexed)[/cyan]") + + if original_dataset_items is None: + original_dataset_items = dataset.load(dataset_path, max_items=None) + + items_to_filter = filtered_items if filtered_items is not None else original_dataset_items + filtered_items = [item for i, item in enumerate(items_to_filter, 1) if start_idx <= i <= end_idx] + + total_found = len(filtered_items) + if total_found == 0: + raise ValueError(f"--index-range {index_range!r} selected no questions") + will_run = min(total_found, max_instances) if max_instances else total_found + if max_instances and total_found > max_instances: + console.print( + f"[green]Found {total_found} questions in range {start_idx}-{end_idx} (will run {will_run} due to --max-instances)[/green]" + ) + else: + console.print(f"[green]Found {total_found} questions in range {start_idx}-{end_idx}[/green]") + except (ValueError, AttributeError) as exc: + raise ValueError(f"Invalid --index-range {index_range!r}; use 'start,end' (for example, '75,412')") from exc + + # Filter dataset based on failed/invalid flags + if only_failed or only_invalid: + target_ids = failed_question_ids if only_failed else invalid_question_ids + if not target_ids: + filter_type = "failed" if only_failed else "invalid" + raise ValueError(f"No {filter_type} questions were found in the source results") + + # Load original items if not already loaded + if original_dataset_items is None: + # Load full dataset without max_instances limit for filtering + original_dataset_items = dataset.load(dataset_path, max_items=None) + + # If we already have filtered_items from category filtering, filter those + # Otherwise start with all items + items_to_filter = filtered_items if filtered_items is not None else original_dataset_items + filtered_items = [item for item in items_to_filter if dataset.get_item_id(item) in target_ids] + + filter_type = "failed" if only_failed else "invalid" + total_found = len(filtered_items) + will_run = min(total_found, max_instances) if max_instances else total_found + if max_instances and total_found > max_instances: + console.print( + f"[green]Found {total_found} {filter_type} items to re-evaluate (will run {will_run} due to --max-instances)[/green]" + ) + else: + console.print(f"[green]Found {total_found} {filter_type} items to re-evaluate[/green]") + + output_path = resolve_output_path(results_dir, results_filename) + merge_with_existing = ( + filln + or question_id is not None + or only_failed + or only_invalid + or only_ingested + or category is not None + or max_instances_per_category is not None + or resume + ) + current_manifest = build_run_manifest( + dataset_path=dataset_path, + context_format=context_format, + max_instances=max_instances, + max_instances_per_category=max_instances_per_category, + max_questions_per_instance=max_questions_per_instance, + question_id=question_id, + index_range=index_range, + category=category, + max_concurrent_items=max_concurrent_items, + max_concurrent_questions=max_concurrent_questions, + eval_semaphore_size=eval_semaphore_size, + thinking_budget=thinking_budget, + max_tokens=max_tokens, + oracle_planner_v26=oracle_planner_v26, + oracle_planner_v220=oracle_planner_v220, + query_expansion_enabled=query_expansion_enabled, + query_rewriting_strategy=query_rewriting_strategy, + session_expansion_weight=session_expansion_weight, + skip_ingestion=skip_ingestion or only_ingested, + ingest_only=ingest_only, + force_reingest=force_reingest, + ) + current_model_config = get_model_config() + validate_output_target( + output_path, + merge_with_existing=merge_with_existing, + resume=resume, + ) + if output_path.exists() and merge_with_existing: + validate_artifact_compatibility( + output_path, + current_manifest=current_manifest, + current_model_config=current_model_config, + ) + output_path.parent.mkdir(parents=True, exist_ok=True) + if resume: + console.print(f"[cyan]Resume mode enabled: using compatible results from {output_path}[/cyan]") + + # Create local memory engine only after all local artifact checks pass. + from hms_api.engine.memory_engine import Budget + from hms_api.models import RequestContext + + from benchmarks.common.benchmark_runner import create_memory_engine + + memory = await create_memory_engine() + + evidence_mode = None + if oracle_planner_v220: + evidence_mode = "oracle_v220" + elif oracle_planner_v26: + evidence_mode = "oracle_v26" + + # Create answer generator + answer_generator = LongMemEvalAnswerGenerator( + context_format=context_format, + evidence_mode=evidence_mode, + ) + # Log context format being used + console.print(f"[blue]Context format: {context_format}[/blue]") + + answer_evaluator = LLMAnswerEvaluator() + + # Filter by only_ingested: only run items whose memory bank already exists + if only_ingested: + console.print("[cyan]Filtering to only items with existing memory banks...[/cyan]") + + # Load all items if not already loaded + if original_dataset_items is None: + original_dataset_items = dataset.load(dataset_path, max_items=None) + + items_to_check = filtered_items if filtered_items is not None else original_dataset_items + + # Check which items have existing banks + ingested_items = [] + pool = await memory._get_pool() + + for item in items_to_check: + item_id = dataset.get_item_id(item) + agent_id = f"longmemeval_{item_id}" + + # A retained bank is reusable when its source chunks are durable; + # fact extraction may legitimately produce zero facts. + async with pool.acquire() as conn: + has_chunks = await conn.fetchval( + f"SELECT EXISTS(SELECT 1 FROM {fq_table('chunks')} WHERE bank_id = $1)", + agent_id, + ) + if has_chunks: + ingested_items.append(item) + + filtered_items = ingested_items + console.print(f"[green]Found {len(filtered_items)} items with existing memory banks[/green]") + + if not filtered_items: + raise RuntimeError("No items with durable retained source chunks were found") + + # Determine query rewriting strategy + if query_expansion_enabled: + strategy_name = query_rewriting_strategy + else: + strategy_name = "noop" + + # Create benchmark runner + retrieval_planner = None + if oracle_planner_v220: + retrieval_planner = longmemeval_oracle_planner_v220 + elif oracle_planner_v26: + retrieval_planner = longmemeval_oracle_planner_v26 + + runner = BenchmarkRunner( + dataset=dataset, + answer_generator=answer_generator, + answer_evaluator=answer_evaluator, + memory=memory, + query_rewriting_strategy_name=strategy_name, + query_rewriting_enabled=query_expansion_enabled, + session_expansion_weight=session_expansion_weight, + retrieval_planner=retrieval_planner, + ) + + if query_expansion_enabled: + console.print(f"[cyan]Query expansion enabled: using {strategy_name} strategy[/cyan]") + + console.print(f"[cyan]Session expansion weight: {session_expansion_weight}[/cyan]") + if oracle_planner_v26: + console.print("[cyan]Oracle planner v2.6 enabled: base retrieval + structured evidence ledger[/cyan]") + for planner_category, planner_weight in sorted(ORACLE_PLANNER_V1_WEIGHTS.items()): + suffix = " + query expansion + evidence appendix" if planner_category == "multi-session" else "" + ledger = ( + " + high-risk ledger" + if planner_category in {"multi-session", "temporal-reasoning", "knowledge-update"} + else "" + ) + console.print(f" [cyan]{planner_category}:[/cyan] {planner_weight}{suffix}{ledger}") + if oracle_planner_v220: + profile = SELF_EVOLUTION_PROFILES["oracle_v220"] + console.print("[cyan]Oracle planner v2.20 enabled: pure v2.6 + diagnosis-driven self-evolution[/cyan]") + console.print(f" [cyan]base:[/cyan] {profile['base']}") + console.print(f" [cyan]diagnosis source:[/cyan] {profile['diagnosis_source']}") + console.print(f" [cyan]selection:[/cyan] {profile['selection_rule']}") + for planner_category, planner_weight in sorted(ORACLE_PLANNER_V1_WEIGHTS.items()): + suffix = " + query expansion + evidence appendix" if planner_category == "multi-session" else "" + ledger = ( + " + V2.6 ledger" + if planner_category in {"multi-session", "temporal-reasoning", "knowledge-update"} + else "" + ) + controller = ( + " + self-evolution controller" + if planner_category in {"multi-session", "temporal-reasoning", "knowledge-update"} + else "" + ) + date_block = " + self-evolved date evidence" if planner_category == "temporal-reasoning" else "" + console.print( + f" [cyan]{planner_category}:[/cyan] {planner_weight}{suffix}{ledger}{controller}{date_block}" + ) + + # If filtering by category, failed, invalid, only_ingested, or max_instances_per_category, we need to use a custom dataset that only returns those items + # We'll temporarily replace the dataset's load method + if filtered_items is not None: + original_load = dataset.load + + def filtered_load(path: Path, max_items: Optional[int] = None): + return filtered_items[:max_items] if max_items else filtered_items + + dataset.load = filtered_load + + # Run benchmark + # Single-phase approach: each question gets its own isolated agent_id + # This ensures each question only has access to its own context + + # Configuration for single-phase benchmark + separate_ingestion = False + clear_per_item = True # Use unique agent_id per question + + results = await runner.run( + dataset_path=dataset_path, + agent_id="longmemeval", # Will be suffixed with question_id per item + max_items=max_instances + if not max_instances_per_category + else None, # Don't apply max_items when using per-category limit + max_questions_per_item=max_questions_per_instance, + thinking_budget=thinking_budget, + max_tokens=max_tokens, + skip_ingestion=skip_ingestion or only_ingested, # Auto-skip ingestion when using --only-ingested + max_concurrent_questions=max_concurrent_questions, + eval_semaphore_size=eval_semaphore_size, + separate_ingestion_phase=separate_ingestion, + clear_agent_per_item=clear_per_item, + filln=filln or resume, # Resume skips item IDs already present in the result file. + specific_item=None, # Already filtered via filtered_items replacement + max_concurrent_items=max_concurrent_items, # Parallel instance processing + output_path=output_path, # Save results incrementally + merge_with_existing=merge_with_existing, # Merge when using --fill, --category, --only-failed, --only-invalid flags or specific question + ingest_only=ingest_only, # Only ingest, skip evaluation + force_reingest=force_reingest, # Force re-ingest even if data already exists + rerun_invalid_existing=resume, + run_manifest=current_manifest, + ) + results["run_manifest"] = current_manifest + runner.save_results(results, output_path) + + if ingest_only: + console.print("\n[green]✓[/green] Ingest-only mode completed. Data is ready for evaluation.") + console.print(" To run evaluation later with a different model:") + console.print(" 1. Update .env with your preferred model") + console.print(" 2. Run: HMS_BENCHMARK=longmemeval bash .aaaSCRIPT/run_benchmark.sh --only-ingested --fill") + return results + + full_run_requested = ( + dataset_path.resolve() == DEFAULT_DATASET_PATH.resolve() + and max_instances is None + and max_instances_per_category is None + and question_id is None + and index_range is None + and category is None + and not only_failed + and not only_invalid + and not only_ingested + ) + if full_run_requested: + if ( + results.get("num_items") != LONGMEMEVAL_EXPECTED_ITEMS + or results.get("total_questions") != LONGMEMEVAL_EXPECTED_ITEMS + ): + raise RuntimeError( + "Incomplete full LongMemEval run: " + f"items={results.get('num_items')}, questions={results.get('total_questions')}, " + f"expected={LONGMEMEVAL_EXPECTED_ITEMS}" + ) + if results.get("total_invalid", 0): + raise RuntimeError( + f"Full LongMemEval run contains {results['total_invalid']} invalid question(s); " + f"inspect {output_path} and resume after correcting the provider or runtime failure" + ) + + # Display results (final save already happened incrementally) + runner.display_results(results) + console.print(f"\n[green]✓[/green] Results saved incrementally to {output_path}") + + # Generate detailed report by question type + generate_type_report(results) + + # Generate markdown results table + generate_markdown_table(results, output_path) + + return results + + +def _sha256_file(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for block in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def validate_canonical_dataset(dataset_path: Path) -> None: + """Validate the immutable LongMemEval artifact used by the full benchmark.""" + + actual_size = dataset_path.stat().st_size + if actual_size != LONGMEMEVAL_DATASET_SIZE: + raise ValueError( + f"LongMemEval dataset size mismatch at {dataset_path}: " + f"expected {LONGMEMEVAL_DATASET_SIZE}, got {actual_size}" + ) + actual_sha256 = _sha256_file(dataset_path) + if actual_sha256 != LONGMEMEVAL_DATASET_SHA256: + raise ValueError( + f"LongMemEval dataset checksum mismatch at {dataset_path}: " + f"expected {LONGMEMEVAL_DATASET_SHA256}, got {actual_sha256}" + ) + + +def download_dataset(dataset_path: Path) -> None: + """Download, verify, and atomically install the pinned LongMemEval dataset.""" + + from rich.console import Console + + console = Console() + dataset_path.parent.mkdir(parents=True, exist_ok=True) + partial_path = dataset_path.with_name(f".{dataset_path.name}.part") + partial_path.unlink(missing_ok=True) + + console.print("[yellow]Dataset not found. Downloading the pinned LongMemEval artifact...[/yellow]") + console.print(f"[dim]URL: {LONGMEMEVAL_DATASET_URL}[/dim]") + console.print(f"[dim]Destination: {dataset_path}[/dim]") + + request = urllib.request.Request( + LONGMEMEVAL_DATASET_URL, + headers={"User-Agent": "HMS-LongMemEval-Reproduction/1.0"}, + ) + digest = hashlib.sha256() + total_bytes = 0 + try: + with urllib.request.urlopen(request, timeout=120) as response, partial_path.open("wb") as handle: + while True: + block = response.read(1024 * 1024) + if not block: + break + handle.write(block) + digest.update(block) + total_bytes += len(block) + handle.flush() + os.fsync(handle.fileno()) + + if total_bytes != LONGMEMEVAL_DATASET_SIZE: + raise ValueError( + f"Downloaded dataset size mismatch: expected {LONGMEMEVAL_DATASET_SIZE}, got {total_bytes}" + ) + actual_sha256 = digest.hexdigest() + if actual_sha256 != LONGMEMEVAL_DATASET_SHA256: + raise ValueError( + f"Downloaded dataset checksum mismatch: expected {LONGMEMEVAL_DATASET_SHA256}, got {actual_sha256}" + ) + os.replace(partial_path, dataset_path) + console.print("[green]✓ Pinned dataset downloaded and verified[/green]") + except Exception as exc: + partial_path.unlink(missing_ok=True) + raise RuntimeError(f"Failed to download the pinned LongMemEval dataset: {exc}") from exc + + +def generate_type_report(results: dict): + """Generate a detailed report by question type.""" + from rich.console import Console + from rich.table import Table + + console = Console() + + # Aggregate stats by question type + type_stats = {} + + for item_result in results["item_results"]: + metrics = item_result["metrics"] + by_category = metrics.get("category_stats", {}) + + for qtype, stats in by_category.items(): + if qtype not in type_stats: + type_stats[qtype] = {"total": 0, "correct": 0} + type_stats[qtype]["total"] += stats["total"] + type_stats[qtype]["correct"] += stats["correct"] + + # Display table + table = Table(title="Performance by Question Type") + table.add_column("Question Type", style="cyan") + table.add_column("Total", justify="right", style="yellow") + table.add_column("Correct", justify="right", style="green") + table.add_column("Accuracy", justify="right", style="magenta") + + for qtype, stats in sorted(type_stats.items()): + acc = (stats["correct"] / stats["total"] * 100) if stats["total"] > 0 else 0 + table.add_row(qtype, str(stats["total"]), str(stats["correct"]), f"{acc:.1f}%") + + console.print("\n") + console.print(table) + + +def generate_markdown_table(results: dict, json_output_path: Path): + """Generate a markdown results table with model configuration.""" + from rich.console import Console + + console = Console() + + # Aggregate stats by question type + type_stats = {} + + for item_result in results["item_results"]: + metrics = item_result["metrics"] + by_category = metrics.get("category_stats", {}) + + for qtype, stats in by_category.items(): + if qtype not in type_stats: + type_stats[qtype] = {"total": 0, "correct": 0, "invalid": 0} + type_stats[qtype]["total"] += stats["total"] + type_stats[qtype]["correct"] += stats["correct"] + type_stats[qtype]["invalid"] += stats.get("invalid", 0) + + # Build markdown content + lines = [] + lines.append("# LongMemEval Benchmark Results") + lines.append("") + + # Add model configuration + if "model_config" in results: + config = results["model_config"] + lines.append("## Model Configuration") + lines.append("") + lines.append(f"- **HMS**: {config['hms']['provider']}/{config['hms']['model']}") + if "retain" in config: + lines.append(f"- **Retain**: {config['retain']['provider']}/{config['retain']['model']}") + lines.append( + f"- **Answer Generation**: {config['answer_generation']['provider']}/{config['answer_generation']['model']}" + ) + lines.append(f"- **LLM Judge**: {config['judge']['provider']}/{config['judge']['model']}") + if "embeddings" in config: + lines.append(f"- **Embeddings**: {config['embeddings']['provider']}/{config['embeddings']['model']}") + lines.append("") + + lines.append( + f"**Overall Accuracy**: {results['overall_accuracy']:.2f}% ({results['total_correct']}/{results['total_questions']})" + ) + lines.append("") + + # Results by question type + lines.append("## Results by Question Type") + lines.append("") + lines.append("| Question Type | Total | Correct | Invalid | Accuracy |") + lines.append("|---------------|-------|---------|---------|----------|") + + for qtype in sorted(type_stats.keys()): + stats = type_stats[qtype] + acc = (stats["correct"] / stats["total"] * 100) if stats["total"] > 0 else 0 + invalid_str = str(stats["invalid"]) if stats["invalid"] > 0 else "-" + lines.append(f"| {qtype} | {stats['total']} | {stats['correct']} | {invalid_str} | {acc:.1f}% |") + + # Add overall row + total_invalid = results.get("total_invalid", 0) + invalid_str = str(total_invalid) if total_invalid > 0 else "-" + lines.append( + f"| **OVERALL** | **{results['total_questions']}** | **{results['total_correct']}** | **{invalid_str}** | **{results['overall_accuracy']:.1f}%** |" + ) + + # Write to file (same directory as JSON, but .md extension) + md_output_path = json_output_path.with_suffix(".md") + md_output_path.write_text("\n".join(lines) + "\n", encoding="utf-8") + console.print(f"\n[green]✓[/green] Results table saved to {md_output_path}") + + +if __name__ == "__main__": + import argparse + import logging + + parser = argparse.ArgumentParser(description="Run LongMemEval benchmark") + parser.add_argument( + "--max-instances", + type=int, + default=None, + help="Limit TOTAL number of questions to evaluate (default: all 500). For per-category limits, use --max-questions-per-category instead.", + ) + parser.add_argument( + "--max-instances-per-category", + "--max-questions-per-category", # Alias since each instance = 1 question in LongMemEval + type=int, + default=None, + dest="max_instances_per_category", + help="Limit number of questions per category (e.g., 20 = 20 questions from each of the 6 categories = 120 total). Cannot be combined with --max-instances or --category.", + ) + parser.add_argument( + "--max-questions", type=int, default=None, help="Limit number of questions per instance (for quick testing)" + ) + parser.add_argument( + "--thinking-budget", type=int, default=500, help="Thinking budget for spreading activation search" + ) + parser.add_argument("--max-tokens", type=int, default=8192, help="Maximum tokens to retrieve from memories") + parser.add_argument("--skip-ingestion", action="store_true", help="Skip ingestion and use existing data") + parser.add_argument( + "--fill", + action="store_true", + help="Only process questions not already in results file (for resuming interrupted runs)", + ) + parser.add_argument( + "--question-id", + type=str, + default=None, + help="Filter to specific question ID(s). Can be a single synthetic-style ID " + "(e.g., 'sample-question-a') or comma-separated IDs " + "(e.g., 'sample-question-a,sample-question-b'). Useful with --skip-ingestion " + "to test specific questions.", + ) + parser.add_argument( + "--index-range", + type=str, + default=None, + help="Filter to a range of item indices (e.g., '75,412'). Both start and end are inclusive, 1-indexed.", + ) + parser.add_argument( + "--only-failed", + action="store_true", + help="Only run questions that were previously marked as incorrect (is_correct=False). Requires existing results file.", + ) + parser.add_argument( + "--only-invalid", + action="store_true", + help="Only run questions that were previously marked as invalid (is_invalid=True). Requires existing results file.", + ) + parser.add_argument( + "--only-ingested", + action="store_true", + help="Only run questions whose memory bank already exists (has been ingested). Automatically skips ingestion. Cannot be combined with --only-failed, --only-invalid, --category, --question-id, or --max-instances-per-category.", + ) + parser.add_argument( + "--category", + type=str, + default=None, + help="Filter questions by category/question_type. Available categories: 'single-session-user', 'multi-session', 'single-session-preference', 'temporal-reasoning', 'knowledge-update', 'single-session-assistant'. Can be combined with --max-instances to limit questions within the category.", + ) + parser.add_argument( + "--parallel", + type=int, + default=1, + help="Number of instances to process in parallel (default: 1 for sequential). Higher values speed up evaluation but use more memory.", + ) + parser.add_argument( + "--results-filename", + type=str, + default="benchmark_results.json", + help="Filename for results output (default: benchmark_results.json).", + ) + parser.add_argument( + "--results-dir", + type=str, + default=None, + help="Optional directory for results. If not specified, uses results/ relative to script location.", + ) + parser.add_argument( + "--context-format", + type=str, + choices=["json", "structured", "structured_compact", "structured_source"], + default="structured_source", + help="How to format context: 'json' (raw), 'structured' (per-fact), 'structured_compact' (source bundles), or 'structured_source' (bundles plus retained source windows). Default: structured_source.", + ) + parser.add_argument( + "--source-results", + type=str, + default=None, + help="Source results file to read failed/invalid questions from (for --only-failed/--only-invalid). Defaults to benchmark_results.json if not specified.", + ) + parser.add_argument( + "--ingest-only", + action="store_true", + help="Only ingest conversation data (skip evaluation). Use with --fill to skip already ingested items. Use after ingest to do evaluation with different model.", + ) + parser.add_argument( + "--force-reingest", + action="store_true", + help="Force re-ingest even if data already exists. Use when you want to re-process ingestion for items that may have incomplete data.", + ) + parser.add_argument( + "--quiet", + action="store_true", + help="Suppress INFO level log messages (only show warnings and errors)", + ) + parser.add_argument( + "--max-concurrent-questions", + type=int, + default=10, + help="Maximum number of concurrent question processing (default: 10)", + ) + parser.add_argument( + "--eval-semaphore-size", + type=int, + default=10, + help="Maximum concurrent LLM judge requests (default: 10)", + ) + parser.add_argument( + "--dataset-path", + type=str, + default=None, + help="Optional custom dataset path. If not specified, uses the default dataset.", + ) + parser.add_argument( + "--enable-query-expansion", + action="store_true", + help="Enable query rewriting (default: disabled). When enabled, uses --query-rewriting-strategy to determine the strategy.", + ) + parser.add_argument( + "--query-rewriting-strategy", + type=str, + choices=["noop", "llm_based", "llm_driven"], + default="llm_based", + help="Query rewriting strategy to use when --enable-query-expansion is set. Options: 'noop' (no expansion), 'llm_based' (rule-based decision with LLM expansion), 'llm_driven' (LLM-driven analysis with entity expansion and time window calculation) (default: llm_based)", + ) + parser.add_argument( + "--session-expansion-weight", + type=float, + default=0.3, + help="Weight for session-based node expansion (default: 0.3). Set to 0 to disable.", + ) + parser.add_argument( + "--oracle-planner-v26", + action="store_true", + help="Use the V2.6 Structured Evidence Ledger for high-risk questions.", + ) + parser.add_argument( + "--oracle-planner-v220", + action="store_true", + help="Use pure v2.6 retrieval plus diagnosis-driven self-evolution controls.", + ) + parser.add_argument( + "--resume", + action="store_true", + help="Resume from a previous run by merging with existing results. Use with --results-filename to specify the same output file.", + ) + + args = parser.parse_args() + + log_level = logging.WARNING if args.quiet else logging.INFO + logging.basicConfig(level=log_level, format="%(asctime)s %(levelname)s %(message)s") + + # Validate that only one of --only-failed or --only-invalid is set + if args.only_failed and args.only_invalid: + parser.error("Cannot use both --only-failed and --only-invalid at the same time") + + planner_flags = [ + args.oracle_planner_v26, + args.oracle_planner_v220, + ] + if sum(1 for flag in planner_flags if flag) > 1: + parser.error("Cannot use more than one oracle planner flag at the same time") + + # Validate mutually exclusive arguments + # --max-instances-per-category can't be combined with --max-instances or --category + if args.max_instances_per_category is not None and (args.max_instances is not None or args.category is not None): + parser.error("--max-questions-per-category cannot be combined with --max-instances or --category") + + results = asyncio.run( + run_benchmark( + max_instances=args.max_instances, + max_instances_per_category=args.max_instances_per_category, + max_questions_per_instance=args.max_questions, + thinking_budget=args.thinking_budget, + max_tokens=args.max_tokens, + skip_ingestion=args.skip_ingestion, + filln=args.fill, + question_id=args.question_id, + index_range=args.index_range, + only_failed=args.only_failed, + only_invalid=args.only_invalid, + only_ingested=args.only_ingested, + category=args.category, + max_concurrent_items=args.parallel, + results_filename=args.results_filename, + results_dir=args.results_dir, + context_format=args.context_format, + source_results=args.source_results, + ingest_only=args.ingest_only, + force_reingest=args.force_reingest, + max_concurrent_questions=args.max_concurrent_questions, + eval_semaphore_size=args.eval_semaphore_size, + dataset_path=args.dataset_path, + query_expansion_enabled=args.enable_query_expansion, + query_rewriting_strategy=args.query_rewriting_strategy, + session_expansion_weight=args.session_expansion_weight, + oracle_planner_v26=args.oracle_planner_v26, + oracle_planner_v220=args.oracle_planner_v220, + resume=args.resume, + ) + ) diff --git a/lab/evaluation/benchmarks/longmemeval/source_backfill.py b/lab/evaluation/benchmarks/longmemeval/source_backfill.py new file mode 100644 index 0000000..6095c14 --- /dev/null +++ b/lab/evaluation/benchmarks/longmemeval/source_backfill.py @@ -0,0 +1,504 @@ +"""Query-time recovery of source spans retained in ``documents.original_text``. + +Fact extraction is intentionally lossy. When a retrieved fact points to a +document, the original retained transcript is still a trustworthy source for +details that were not materialized as a fact. This module selects small, +query-focused turn windows; it never invents text or broadens a bank scope. +""" + +from __future__ import annotations + +import json +import re +from collections import Counter +from dataclasses import dataclass +from typing import Any, Mapping, Sequence + +_STOPWORDS = frozenset( + { + "a", + "an", + "about", + "after", + "again", + "and", + "am", + "are", + "as", + "at", + "be", + "before", + "between", + "by", + "can", + "current", + "currently", + "could", + "different", + "did", + "do", + "does", + "during", + "for", + "first", + "from", + "have", + "has", + "had", + "how", + "i", + "in", + "interesting", + "is", + "it", + "me", + "many", + "much", + "might", + "my", + "of", + "on", + "or", + "previous", + "please", + "recommend", + "recommendation", + "recommendations", + "recent", + "recently", + "some", + "since", + "should", + "suggest", + "suggestion", + "suggestions", + "that", + "the", + "then", + "there", + "this", + "total", + "tell", + "to", + "find", + "looking", + "like", + "was", + "were", + "what", + "when", + "where", + "which", + "with", + "who", + "would", + "you", + "your", + } +) +_TOKEN_RE = re.compile(r"[A-Za-z][A-Za-z0-9_'-]{2,}") +_SIGNAL_RE = re.compile( + r"(?:\$\s?\d|\b\d+(?:[.,]\d+)?%?|\b(?:jan|feb|mar|apr|may|jun|jul|aug|sep|oct|nov|dec)\b|" + r"\b(?:before|after|earlier|later|first|last|current|latest|total|spent|cost|discount)\b)", + re.IGNORECASE, +) +_QUERY_SIGNAL_RE = re.compile( + r"\b(?:number|amount|cost|price|date|day|days|week|weeks|month|time|hour|hours|" + r"how many|how much|earliest|latest|first|last|before|after|current|total|difference|order)\b", + re.IGNORECASE, +) + + +@dataclass(frozen=True) +class SourceSnippet: + document_id: str + turn_start: int + turn_end: int + text: str + score: int + focus_turn: int | None = None + + +@dataclass(frozen=True) +class SourceChunkSnippet: + """A chunk referenced by a retrieved fact but omitted from the response.""" + + chunk_id: str + document_id: str | None + chunk_index: int | None + text: str + score: int + + +def _query_terms(query: str) -> frozenset[str]: + return frozenset(token for token in _TOKEN_RE.findall(query.lower()) if token not in _STOPWORDS) + + +def _turn_text(turn: Any) -> str: + if isinstance(turn, Mapping): + content = turn.get("content") + if isinstance(content, str): + text = content + elif content is not None: + text = json.dumps(content, ensure_ascii=False) + else: + text = "" + role = str(turn.get("role") or "").strip().lower() + return f"{role}: {text}" if role and text else text + return str(turn or "") + + +def _parse_turns(raw: Any) -> list[str]: + if isinstance(raw, str): + try: + raw = json.loads(raw) + except (TypeError, ValueError): + return [raw] if raw.strip() else [] + if isinstance(raw, Mapping): + raw = raw.get("messages") or raw.get("turns") or [raw] + if not isinstance(raw, Sequence) or isinstance(raw, (str, bytes, bytearray)): + return [] + return [_turn_text(turn) for turn in raw if _turn_text(turn).strip()] + + +def _compact(text: str, limit: int) -> str: + text = " ".join(text.replace("<|endoftext|>", " ").split()) + if len(text) <= limit: + return text + return text[: limit - 3].rstrip() + "..." + + +def _normalise_for_match(text: Any) -> str: + """Normalise text for source-coverage and duplicate checks.""" + + return re.sub(r"\W+", " ", str(text or "").lower()).strip() + + +def _content_without_role(text: str) -> str: + if re.match(r"^(?:user|assistant|system|tool):\s", text, re.IGNORECASE): + return text.split(":", 1)[1].lstrip() + return text + + +def _is_turn_covered(turn: str, covered_chunks: Sequence[str]) -> bool: + """Return whether a retained turn is already represented by a chunk. + + Recall chunks are JSON windows while source turns are rendered as plain + text. Suppress a backfill only when a substantial normalised substring is + visibly present in an existing chunk. + """ + + if not covered_chunks: + return False + normalised_turn = _normalise_for_match(_content_without_role(turn)) + if not normalised_turn: + return True + prefix = normalised_turn[:180] if len(normalised_turn) >= 180 else normalised_turn + for chunk in covered_chunks: + normalised_chunk = _normalise_for_match(chunk) + if not normalised_chunk: + continue + if normalised_turn in normalised_chunk: + return True + # Short turns are commonly embedded intact in a larger chunk. Do not + # treat a long-turn prefix as coverage: token-budget truncation can + # leave the opening sentence while dropping the answer-bearing tail. + if len(normalised_turn) <= 240 and len(prefix) >= 80 and prefix in normalised_chunk: + return True + return False + + +def _query_cluster_center(text: str, query_terms: frozenset[str], window: int) -> int | None: + """Find the densest deterministic cluster of query-term matches.""" + + matches = sorted( + (match.start(), match.end(), term) + for term in query_terms + for match in re.finditer(rf"\b{re.escape(term)}\b", text) + ) + if not matches: + return None + + left = 0 + counts: Counter[str] = Counter() + best_score: tuple[int, int, int, int] | None = None + best_bounds = (matches[0][0], matches[0][1]) + for right, (right_position, right_end, term) in enumerate(matches): + counts[term] += 1 + while left < right and right_end - matches[left][0] > window: + left_term = matches[left][2] + counts[left_term] -= 1 + if not counts[left_term]: + del counts[left_term] + left += 1 + left_position = matches[left][0] + span = right_end - left_position + score = (len(counts), right - left + 1, -span, right_position) + if best_score is None or score > best_score: + best_score = score + best_bounds = (left_position, right_end) + return sum(best_bounds) // 2 + + +def _compact_relevant(text: str, limit: int, query_terms: frozenset[str]) -> str: + """Compact one turn around a relevant match instead of from the left.""" + + text = " ".join(text.replace("<|endoftext|>", " ").split()) + if len(text) <= limit: + return text + role_match = re.match(r"^(?:user|assistant|system|tool):\s", text, re.IGNORECASE) + role_prefix = role_match.group(0) if role_match else "" + body = text[len(role_prefix) :] + if limit <= len(role_prefix) + 6: + return text[:limit] + + window = limit - len(role_prefix) - 6 + center = _query_cluster_center(body.lower(), query_terms, window) + if center is None: + signal_match = _SIGNAL_RE.search(body) + center = signal_match.start() if signal_match else 0 + start = max(0, center - window // 2) + if start + window > len(body): + start = max(0, len(body) - window) + excerpt = body[start : start + window].strip() + if start > 0: + excerpt = "..." + excerpt + if start + window < len(body): + excerpt = excerpt.rstrip() + "..." + return role_prefix + excerpt + + +def _render_turn_window( + turns: Sequence[str], + start: int, + end: int, + focus: int, + query_terms: frozenset[str], + limit: int, +) -> str: + """Render the focus turn first, then bounded neighbouring turns.""" + + indexes = list(range(start, end)) + if not indexes: + return "" + focus_limit = min(limit, max(640, int(limit * 0.62))) + remaining = max(0, limit - focus_limit) + neighbours = [index for index in indexes if index != focus] + neighbour_limit = remaining // len(neighbours) if neighbours else 0 + rendered: list[str] = [] + for index in indexes: + per_turn = focus_limit if index == focus else neighbour_limit + if per_turn <= 0: + continue + piece = _compact_relevant(turns[index], per_turn, query_terms) + if piece: + rendered.append(piece) + return "\n".join(rendered) + + +def select_source_snippets( + documents: Mapping[str, Any], + query: str, + document_order: Sequence[str], + *, + max_documents: int = 12, + max_snippets: int = 16, + max_snippets_per_document: int = 2, + max_chars_per_snippet: int = 1800, + max_total_chars: int = 18_000, + min_score: int = 2, + covered_chunks: Mapping[str, Sequence[str]] | None = None, +) -> list[SourceSnippet]: + """Select deterministic source windows from already-retrieved documents. + + Documents are considered in retrieval order. Within each document, a + scored turn and its immediate neighbours form one window, preserving the + conversational context around a matching answer. Existing returned chunks + can be supplied through ``covered_chunks``; a focus turn that is already + visible is omitted so this function only adds provenance recall lost. + Round-robin selection prevents one long transcript from consuming the + budget. + """ + + terms = _query_terms(query) + candidates_by_doc: dict[str, list[SourceSnippet]] = {} + seen_docs: set[str] = set() + for document_id in document_order[:max_documents]: + document_id = str(document_id) + if not document_id or document_id in seen_docs or document_id not in documents: + continue + seen_docs.add(document_id) + turns = _parse_turns(documents[document_id]) + if not turns: + continue + + scored_turns: list[tuple[int, int, int]] = [] + has_covered_overlap = False + for index, turn in enumerate(turns): + lower = turn.lower() + overlap = sum(term in lower for term in terms) + signal = 1 if _SIGNAL_RE.search(turn) else 0 + score = overlap * 5 + signal * 2 + if score < min_score: + continue + if covered_chunks and _is_turn_covered(turn, covered_chunks.get(document_id, ())): + has_covered_overlap = has_covered_overlap or overlap > 0 + continue + scored_turns.append((score, overlap, index)) + if any(overlap > 0 for _, overlap, _ in scored_turns): + scored_turns = [row for row in scored_turns if row[1] > 0] + elif has_covered_overlap or not _QUERY_SIGNAL_RE.search(query): + # Query-relevant evidence for this document is already visible; + # do not add an unrelated numeric/date-only fallback window. A + # signal-only fallback remains available for explicitly numeric or + # temporal questions such as "what was the amount?". + scored_turns = [] + scored_turns.sort(key=lambda row: (-row[0], -row[1], row[2])) + + snippets: list[SourceSnippet] = [] + seen_windows: set[str] = set() + for score, _, index in scored_turns: + start = max(0, index - 1) + end = min(len(turns), index + 2) + text = _render_turn_window(turns, start, end, index, terms, max_chars_per_snippet) + key = re.sub(r"\W+", " ", text.lower()).strip() + if not text or key in seen_windows: + continue + seen_windows.add(key) + snippets.append(SourceSnippet(document_id, start, end, text, score, index)) + if len(snippets) >= max_snippets_per_document: + break + if snippets: + candidates_by_doc[document_id] = snippets + + selected: list[SourceSnippet] = [] + total_chars = 0 + # Round-robin gives each high-ranked source a chance before a second window + # from any one document is added. + for offset in range(max_snippets_per_document): + for document_id in document_order[:max_documents]: + snippets = candidates_by_doc.get(str(document_id), []) + if offset >= len(snippets): + continue + snippet = snippets[offset] + if len(selected) >= max_snippets or total_chars + len(snippet.text) > max_total_chars: + return selected + selected.append(snippet) + total_chars += len(snippet.text) + return selected + + +def select_missing_chunk_snippets( + chunks: Sequence[Mapping[str, Any]], + query: str, + *, + max_chunks: int = 6, + max_chars_per_chunk: int = 1200, + max_total_chars: int = 7_000, +) -> list[SourceChunkSnippet]: + """Select exact chunk rows that recall referenced but did not return. + + This is the highest-fidelity recovery path: the chunk was already linked + to a retrieved fact, so no bank-wide search or new extraction is needed. + Rows are ranked by query overlap and retain input order as a stable tie + breaker (the caller supplies fact/retrieval order). + """ + + terms = _query_terms(query) + candidates: list[tuple[int, int, SourceChunkSnippet]] = [] + for position, row in enumerate(chunks): + chunk_id = str(row.get("chunk_id") or "") + if not chunk_id: + continue + raw_text = row.get("chunk_text") or "" + turns = _parse_turns(raw_text) + best_score = 0 + best_index = 0 + retrieval_rank = int(row.get("_retrieval_rank") or position + 1) + rank_bonus = max(0, 12 - min(retrieval_rank, 12)) + if turns: + turn_scores = [] + for index, turn in enumerate(turns): + overlap = sum(term in turn.lower() for term in terms) + signal = 1 if _SIGNAL_RE.search(turn) else 0 + turn_scores.append((overlap * 5 + signal * 2, overlap, index)) + best_score, _, best_index = max(turn_scores, key=lambda item: (item[0], item[1], -item[2])) + text = _render_turn_window( + turns, + max(0, best_index - 1), + min(len(turns), best_index + 2), + best_index, + terms, + max_chars_per_chunk, + ) + else: + text = _compact_relevant(str(raw_text), max_chars_per_chunk, terms) + best_score = sum(term in str(raw_text).lower() for term in terms) * 5 + if _SIGNAL_RE.search(str(raw_text)): + best_score += 2 + best_score += rank_bonus + if not text: + continue + candidates.append( + ( + -best_score, + position, + SourceChunkSnippet( + chunk_id=chunk_id, + document_id=str(row.get("document_id") or "") or None, + chunk_index=row.get("chunk_index"), + text=text, + score=best_score, + ), + ) + ) + + candidates.sort(key=lambda item: (item[0], item[1])) + selected: list[SourceChunkSnippet] = [] + seen_ids: set[str] = set() + total_chars = 0 + for _, _, candidate in candidates: + if candidate.chunk_id in seen_ids: + continue + if len(selected) >= max_chunks or total_chars + len(candidate.text) > max_total_chars: + break + seen_ids.add(candidate.chunk_id) + selected.append(candidate) + total_chars += len(candidate.text) + return selected + + +def render_source_chunk_snippets(snippets: Sequence[SourceChunkSnippet]) -> str: + """Render exact chunk recoveries with explicit chunk provenance.""" + + if not snippets: + return "" + lines = [ + "=== Retrieved-Chunk Provenance Recovery ===", + "These are exact retained chunks linked to retrieved facts but omitted from the normal response because of the chunk token budget.", + "", + ] + for index, snippet in enumerate(snippets, 1): + location = f"document={snippet.document_id or '-'} | chunk_index={snippet.chunk_index if snippet.chunk_index is not None else '-'}" + lines.append( + f"{index}. chunk={snippet.chunk_id} | {location} | relevance_score={snippet.score} | {snippet.text}" + ) + return "\n".join(lines) + + +def render_source_snippets(snippets: Sequence[SourceSnippet]) -> str: + """Render source snippets with explicit provenance and no synthetic facts.""" + + if not snippets: + return "" + lines = [ + "=== Retained Source-Document Evidence ===", + "These excerpts are verbatim windows from retained documents selected because their source document was retrieved. They are provenance evidence, not new inferred facts.", + "", + ] + for index, snippet in enumerate(snippets, 1): + lines.append( + f"{index}. document={snippet.document_id} | turns={snippet.turn_start}-{snippet.turn_end - 1} | " + f"relevance_score={snippet.score} | {snippet.text}" + ) + return "\n".join(lines) diff --git a/lab/evaluation/benchmarks/longmemeval/test_evidence_bundles.py b/lab/evaluation/benchmarks/longmemeval/test_evidence_bundles.py new file mode 100644 index 0000000..c6d5b8d --- /dev/null +++ b/lab/evaluation/benchmarks/longmemeval/test_evidence_bundles.py @@ -0,0 +1,178 @@ +import json +import os +import subprocess +import sys +from pathlib import Path + +from benchmarks.longmemeval.evidence_bundles import ( + build_evidence_bundles, + render_evidence_bundles, + render_evidence_with_coverage, +) + + +def _fact(text, *, chunk_id="chunk-1", document_id="doc-1", fact_type="world"): + return { + "id": text, + "text": text, + "fact_type": fact_type, + "document_id": document_id, + "chunk_id": chunk_id, + "occurred_start": None, + "occurred_end": None, + "mentioned_at": None, + } + + +def test_bundles_render_each_source_chunk_once_and_keep_distinct_facts(): + results = [ + _fact("A generic recommendation about travel."), + _fact("The train cost $50 and was booked on May 20."), + _fact("A second distinct event in the same source.", chunk_id="chunk-2"), + ] + chunks = { + "chunk-1": {"chunk_text": "raw chunk one", "chunk_index": 0}, + "chunk-2": {"chunk_text": "raw chunk two", "chunk_index": 1}, + } + + bundles = build_evidence_bundles(results, chunks, "How much did the train cost?") + rendered = render_evidence_with_coverage(bundles) + + assert len(bundles) == 2 + assert rendered.text.count("raw chunk one") == 1 + assert rendered.text.count("raw chunk two") == 1 + assert "The train cost $50" in rendered.text + assert rendered.covered_by_document == {"doc-1": ("raw chunk one", "raw chunk two")} + + +def test_bundle_selection_is_stable_for_missing_chunk_and_observation_rows(): + results = [ + _fact("First fact", chunk_id=None, document_id="doc-1"), + _fact("Second fact", chunk_id=None, document_id="doc-1"), + _fact("Third fact", chunk_id=None, document_id="doc-1"), + ] + bundles = build_evidence_bundles(results, {}, "What happened?", max_facts_per_bundle=2) + + assert len(bundles) == 1 + assert [fact["text"] for fact in bundles[0].facts] == ["First fact", "Second fact"] + + +def test_chunk_rendering_preserves_late_state_update_after_long_turn(): + long_assistant = "synthetic robotics workshop venue recommendations " * 180 + results = [_fact("User selected the Copper Finch Inn", chunk_id="chunk-1")] + chunks = { + "chunk-1": { + "chunk_text": json.dumps( + [ + {"role": "user", "content": "I am planning a fictional robotics workshop."}, + {"role": "assistant", "content": long_assistant}, + { + "role": "user", + "content": "I selected the Copper Finch Inn as the Northbridge lodging for the robotics workshop.", + }, + ] + ) + } + } + + query = "What lodging did I select for the Northbridge robotics workshop?" + bundles = build_evidence_bundles(results, chunks, query) + rendered = render_evidence_with_coverage( + bundles, + max_chunk_chars=700, + query=query, + ) + + assert "Copper Finch Inn" in rendered.text + + +def test_render_bound_keeps_whole_bundles_and_stops_before_budget(): + results = [_fact(f"fact {index}", chunk_id=f"chunk-{index}") for index in range(8)] + chunks = {f"chunk-{index}": {"chunk_text": "x" * 200} for index in range(8)} + + bundles = build_evidence_bundles(results, chunks, "What happened?", max_bundles=8) + rendered = render_evidence_with_coverage(bundles, max_chunk_chars=80, max_total_chars=500) + + assert len(rendered.text) <= 500 + assert "Bundle 1" in rendered.text + # A budget cut must not emit a partial bundle header without its fact. + assert rendered.text.count("Bundle ") == rendered.text.count("- Fact ") + + +def test_render_coverage_contains_only_admitted_compact_excerpts_at_exact_cap(): + results = [ + _fact("first fact", chunk_id="chunk-1", document_id="doc-1"), + _fact("second fact", chunk_id="chunk-2", document_id="doc-2"), + ] + raw_chunks = { + "chunk-1": {"chunk_text": json.dumps([{"role": "user", "content": "alpha " + "x " * 100}])}, + "chunk-2": {"chunk_text": json.dumps([{"role": "user", "content": "beta " + "y " * 100}])}, + } + bundles = build_evidence_bundles(results, raw_chunks, "alpha beta") + first_only = render_evidence_with_coverage(bundles[:1], max_chunk_chars=80, query="alpha beta") + + rendered = render_evidence_with_coverage( + bundles, + max_chunk_chars=80, + max_total_chars=len(first_only.text), + query="alpha beta", + ) + + assert rendered.text == first_only.text + assert len(rendered.text) == len(first_only.text) + assert set(rendered.covered_by_document) == {"doc-1"} + excerpt = rendered.covered_by_document["doc-1"][0] + assert excerpt in rendered.text + assert excerpt != raw_chunks["chunk-1"]["chunk_text"] + assert "doc-2" not in rendered.covered_by_document + + +def test_bundle_renderer_keeps_string_return_contract(): + bundles = build_evidence_bundles( + [_fact("visible fact", chunk_id="chunk-1", document_id="doc-1")], + {}, + "visible", + ) + + rendered = render_evidence_bundles(bundles) + + assert isinstance(rendered, str) + assert "visible fact" in rendered + + +def test_bundle_order_uses_source_first_rank_even_when_later_facts_are_selected(): + results = [ + _fact("generic", chunk_id="chunk-a", document_id="doc-a"), + _fact("beta match", chunk_id="chunk-b", document_id="doc-b"), + _fact("beta amount $10", chunk_id="chunk-a", document_id="doc-a"), + _fact("beta amount $20", chunk_id="chunk-a", document_id="doc-a"), + ] + + bundles = build_evidence_bundles(results, {}, "beta amount", max_facts_per_bundle=2) + + assert [bundle.source_key for bundle in bundles] == ["chunk-a", "chunk-b"] + assert bundles[0].first_rank == 1 + assert [fact["text"] for fact in bundles[0].facts] == ["beta amount $10", "beta amount $20"] + + +def test_long_turn_anchor_is_stable_across_hash_seeds_and_uses_latest_match(): + evaluation_root = Path(__file__).resolve().parents[2] + script = """ +import json +from benchmarks.longmemeval.evidence_bundles import _compact_chunk_text, _terms + +raw = json.dumps([{ + "role": "user", + "content": "alpha ANSWER_EARLY " + ("x " * 300) + " beta ANSWER_LATE", +}]) +print(_compact_chunk_text(raw, _terms("alpha beta"), 80)) +""" + outputs = set() + for seed in ("1", "2", "3", "4", "5", "6", "7", "8"): + env = os.environ.copy() + env["PYTHONHASHSEED"] = seed + env["PYTHONPATH"] = os.pathsep.join(path for path in (str(evaluation_root), env.get("PYTHONPATH", "")) if path) + outputs.add(subprocess.check_output([sys.executable, "-c", script], env=env, text=True).strip()) + + assert len(outputs) == 1 + assert "beta ANSWER_LATE" in outputs.pop() diff --git a/lab/evaluation/benchmarks/longmemeval/test_release_integrity.py b/lab/evaluation/benchmarks/longmemeval/test_release_integrity.py new file mode 100644 index 0000000..467e7ca --- /dev/null +++ b/lab/evaluation/benchmarks/longmemeval/test_release_integrity.py @@ -0,0 +1,308 @@ +import copy +import hashlib +import json +import os +import subprocess +from pathlib import Path +from unittest.mock import AsyncMock + +import pytest + +from benchmarks.longmemeval import longmemeval_benchmark as benchmark + + +@pytest.mark.asyncio +async def test_answer_provider_failure_propagates_to_the_runner(): + generator = object.__new__(benchmark.LongMemEvalAnswerGenerator) + generator.context_format = "json" + generator.evidence_mode = None + generator.llm_config = type("Config", (), {})() + generator.llm_config.call = AsyncMock(side_effect=RuntimeError("provider unavailable")) + + with pytest.raises(RuntimeError, match="provider unavailable"): + await generator.generate_answer( + "What happened?", + {"results": [{"text": "memory"}]}, + ) + + +def _manifest() -> dict: + return { + "artifact_schema_version": 1, + "dataset": { + "path": "datasets/dataset.json", + "sha256": "dataset-sha", + "revision": "revision", + "expected_full_items": 500, + }, + "pipeline": { + "planner": "ledger", + "context_format": "structured_source", + "thinking_budget": 500, + "max_tokens": 8192, + "query_expansion_enabled": False, + "query_rewriting_strategy": "noop", + "session_expansion_weight": 0.3, + }, + "database": { + "backend": "postgresql", + "schema": "public", + "vector_extension": "pgvector", + }, + "concurrency": {"items": 1, "questions": 1, "judge": 1}, + "runtime": { + "git_commit": "abc123", + "git_dirty": False, + "source_tree_fingerprint": None, + }, + } + + +def _model_config() -> dict: + role = { + "provider": "openai", + "model": "gpt-5-mini", + "endpoint_fingerprint": "sha256:endpoint", + } + return { + "hms": dict(role), + "retain": dict(role), + "answer_generation": dict(role), + "judge": dict(role), + "embeddings": { + "provider": "openai", + "model": "text-embedding-3-small", + "fingerprint_policy": "strict", + "endpoint_fingerprint": "sha256:embedding", + }, + "reranker": {"provider": "rrf", "model": ""}, + } + + +def _built_manifest(dataset_path: Path) -> dict: + return benchmark.build_run_manifest( + dataset_path=dataset_path, + context_format="structured_source", + max_instances=1, + max_instances_per_category=None, + max_questions_per_instance=1, + question_id=None, + index_range=None, + category=None, + max_concurrent_items=1, + max_concurrent_questions=1, + eval_semaphore_size=1, + thinking_budget=500, + max_tokens=8192, + oracle_planner_v26=False, + oracle_planner_v220=False, + query_expansion_enabled=False, + query_rewriting_strategy="noop", + session_expansion_weight=0.3, + skip_ingestion=False, + ingest_only=False, + force_reingest=False, + ) + + +def test_run_manifest_never_serializes_an_absolute_dataset_path(tmp_path: Path, monkeypatch): + repository_root = tmp_path / "checkout" + repository_dataset = repository_root / ".aaaDATA" / "longmemeval" / "dataset.json" + repository_dataset.parent.mkdir(parents=True) + repository_dataset.write_text("repository dataset", encoding="utf-8") + + external_dataset = tmp_path / "private" / "custom-dataset.json" + external_dataset.parent.mkdir() + external_dataset.write_text("external dataset", encoding="utf-8") + + monkeypatch.setattr(benchmark, "REPOSITORY_ROOT", repository_root) + monkeypatch.setattr(benchmark, "_git_value", lambda *args: "abc123" if args == ("rev-parse", "HEAD") else "") + monkeypatch.setattr(benchmark, "_source_tree_fingerprint", lambda: None) + + repository_manifest = _built_manifest(repository_dataset) + external_manifest = _built_manifest(external_dataset) + + assert repository_manifest["dataset"]["path"] == ".aaaDATA/longmemeval/dataset.json" + assert external_manifest["dataset"]["path"] == "external:custom-dataset.json" + assert not Path(repository_manifest["dataset"]["path"]).is_absolute() + assert not Path(external_manifest["dataset"]["path"]).is_absolute() + + +def test_resume_compatibility_ignores_concurrency_but_rejects_model_changes(tmp_path: Path): + output_path = tmp_path / "results.json" + manifest = _manifest() + model_config = _model_config() + output_path.write_text( + json.dumps( + { + "run_manifest": manifest, + "model_config": model_config, + "item_results": [], + } + ), + encoding="utf-8", + ) + + current_manifest = copy.deepcopy(manifest) + current_manifest["dataset"]["path"] = "external:dataset.json" + current_manifest["concurrency"]["items"] = 8 + benchmark.validate_artifact_compatibility( + output_path, + current_manifest=current_manifest, + current_model_config=model_config, + ) + + incompatible_models = copy.deepcopy(model_config) + incompatible_models["answer_generation"]["model"] = "different-model" + with pytest.raises(ValueError, match="model_config"): + benchmark.validate_artifact_compatibility( + output_path, + current_manifest=current_manifest, + current_model_config=incompatible_models, + ) + + incompatible_source = copy.deepcopy(current_manifest) + incompatible_source["runtime"]["source_tree_fingerprint"] = "sha256:different" + with pytest.raises(ValueError, match="source_tree_fingerprint"): + benchmark.validate_artifact_compatibility( + output_path, + current_manifest=incompatible_source, + current_model_config=model_config, + ) + + +def test_source_tree_fingerprint_tracks_relevant_dirty_content(monkeypatch): + def clean_git_bytes(*args: str) -> bytes: + return b"" + + monkeypatch.setattr(benchmark, "_git_bytes", clean_git_bytes) + assert benchmark._source_tree_fingerprint() is None + + def dirty_git_bytes(*args: str) -> bytes: + return b"diff --git a/source.py b/source.py\n+changed\n" if args[0] == "diff" else b"" + + monkeypatch.setattr(benchmark, "_git_bytes", dirty_git_bytes) + first = benchmark._source_tree_fingerprint() + assert first is not None + assert first.startswith("sha256:") + + def different_git_bytes(*args: str) -> bytes: + return b"diff --git a/source.py b/source.py\n+different\n" if args[0] == "diff" else b"" + + monkeypatch.setattr(benchmark, "_git_bytes", different_git_bytes) + assert benchmark._source_tree_fingerprint() != first + + +def test_explicit_canonical_dataset_path_is_checksum_validated(tmp_path: Path, monkeypatch): + canonical_path = tmp_path / "longmemeval_s_cleaned.json" + canonical_payload = b"good-data" + canonical_path.write_bytes(b"evil-data") + monkeypatch.setattr(benchmark, "DEFAULT_DATASET_PATH", canonical_path) + monkeypatch.setattr(benchmark, "LONGMEMEVAL_DATASET_SIZE", len(canonical_payload)) + monkeypatch.setattr( + benchmark, + "LONGMEMEVAL_DATASET_SHA256", + hashlib.sha256(canonical_payload).hexdigest(), + ) + + with pytest.raises(ValueError, match="checksum mismatch"): + benchmark.resolve_dataset_path(str(canonical_path)) + + canonical_path.write_bytes(canonical_payload) + assert benchmark.resolve_dataset_path(str(canonical_path)) == canonical_path + + +@pytest.mark.parametrize( + ("field_name", "value"), + [ + ("max_instances", 0), + ("max_concurrent_items", 0), + ("max_concurrent_questions", 0), + ("eval_semaphore_size", 0), + ("thinking_budget", 0), + ("max_tokens", 0), + ], +) +def test_runtime_limits_must_be_positive(field_name: str, value: int): + options = { + "max_instances": 1, + "max_instances_per_category": None, + "max_questions_per_instance": 1, + "max_concurrent_items": 1, + "max_concurrent_questions": 1, + "eval_semaphore_size": 1, + "thinking_budget": 1, + "max_tokens": 1, + } + options[field_name] = value + + with pytest.raises(ValueError, match=field_name): + benchmark.validate_runtime_options(**options) + + +def test_fresh_output_refuses_to_overwrite_existing_artifact(tmp_path: Path): + output_path = tmp_path / "results.json" + output_path.write_text('{"old": true}\n', encoding="utf-8") + + with pytest.raises(FileExistsError, match="Refusing to overwrite"): + benchmark.validate_output_target( + output_path, + merge_with_existing=False, + resume=False, + ) + + +def test_resume_requires_an_existing_artifact(tmp_path: Path): + with pytest.raises(FileNotFoundError, match="requires an existing"): + benchmark.validate_output_target( + tmp_path / "missing.json", + merge_with_existing=True, + resume=True, + ) + + +def _launcher_environment(tmp_path: Path) -> dict[str, str]: + environment_file = tmp_path / "empty.env" + environment_file.write_text("", encoding="utf-8") + return { + "PATH": os.environ["PATH"], + "HMS_ENV_FILE": str(environment_file), + "HMS_API_DATABASE_URL": "postgresql://hms:test@127.0.0.1:5432/hms", + "HMS_DATA_DIR": str(tmp_path / "data"), + "HMS_LOG_DIR": str(tmp_path / "logs"), + "HMS_RESULT_DIR": str(tmp_path / "results"), + } + + +def test_launcher_rejects_an_explicit_missing_environment_file(tmp_path: Path): + environment = _launcher_environment(tmp_path) + environment["HMS_ENV_FILE"] = str(tmp_path / "missing.env") + + completed = subprocess.run( + ["bash", str(benchmark.REPOSITORY_ROOT / ".aaaSCRIPT" / "run_benchmark.sh")], + cwd=benchmark.REPOSITORY_ROOT, + env=environment, + capture_output=True, + text=True, + timeout=10, + ) + + assert completed.returncode == 2 + assert "HMS_ENV_FILE does not exist" in completed.stderr + + +def test_launcher_rejects_zero_concurrency_before_starting_python(tmp_path: Path): + environment = _launcher_environment(tmp_path) + environment["HMS_PARALLEL"] = "0" + + completed = subprocess.run( + ["bash", str(benchmark.REPOSITORY_ROOT / ".aaaSCRIPT" / "run_benchmark.sh")], + cwd=benchmark.REPOSITORY_ROOT, + env=environment, + capture_output=True, + text=True, + timeout=10, + ) + + assert completed.returncode == 2 + assert "HMS_PARALLEL must be a positive integer" in completed.stderr diff --git a/lab/evaluation/benchmarks/longmemeval/test_source_backfill.py b/lab/evaluation/benchmarks/longmemeval/test_source_backfill.py new file mode 100644 index 0000000..a353b76 --- /dev/null +++ b/lab/evaluation/benchmarks/longmemeval/test_source_backfill.py @@ -0,0 +1,253 @@ +import json + +from benchmarks.longmemeval.source_backfill import ( + render_source_chunk_snippets, + render_source_snippets, + select_missing_chunk_snippets, + select_source_snippets, +) + + +def test_source_selection_keeps_answer_window_and_document_provenance(): + documents = { + "doc-a": json.dumps( + [ + {"role": "user", "content": "We discussed several fictional puzzle trials."}, + {"role": "assistant", "content": "Your best FableGrid total is 731 markers."}, + {"role": "user", "content": "Now let us discuss a different synthetic trial."}, + ] + ), + "doc-b": json.dumps([{"role": "user", "content": "A generic unrelated note."}]), + } + + snippets = select_source_snippets(documents, "What is my best FableGrid total?", ["doc-a", "doc-b"]) + rendered = render_source_snippets(snippets) + + assert snippets + assert snippets[0].document_id == "doc-a" + assert "731 markers" in rendered + assert "document=doc-a" in rendered + + +def test_source_selection_is_bounded_and_handles_non_json_documents(): + documents = {"doc-a": "plain retained text with 42", "doc-b": ""} + snippets = select_source_snippets( + documents, + "What number was recorded?", + ["doc-a", "doc-b"], + max_snippets=1, + max_chars_per_snippet=20, + ) + + assert len(snippets) == 1 + assert snippets[0].document_id == "doc-a" + assert len(snippets[0].text) <= 20 + + +def test_source_window_keeps_later_turn_after_long_assistant_turn(): + long_answer = "synthetic robotics workshop venue planning details " * 260 + documents = { + "doc-a": json.dumps( + [ + {"role": "user", "content": "I am planning a fictional robotics workshop."}, + {"role": "assistant", "content": long_answer}, + { + "role": "user", + "content": "I selected the Copper Finch Inn as the Northbridge lodging for the robotics workshop.", + }, + ] + ) + } + + snippets = select_source_snippets( + documents, + "What lodging did I select for the Northbridge robotics workshop?", + ["doc-a"], + max_snippets=1, + max_chars_per_snippet=900, + ) + + assert snippets + assert "Copper Finch Inn" in snippets[0].text + assert "assistant:" in snippets[0].text + + +def test_covered_query_turn_does_not_trigger_unrelated_fallback(): + documents = { + "doc-a": json.dumps( + [ + { + "role": "user", + "content": "I selected the Copper Finch Inn as the Northbridge lodging for the robotics workshop.", + }, + {"role": "assistant", "content": "The total cost was 42 dollars."}, + ] + ) + } + + snippets = select_source_snippets( + documents, + "What lodging did I select for the Northbridge robotics workshop?", + ["doc-a"], + covered_chunks={ + "doc-a": [ + json.dumps( + [ + { + "role": "user", + "content": "I selected the Copper Finch Inn as the Northbridge lodging for the robotics workshop.", + } + ] + ) + ] + }, + ) + + assert snippets == [] + + +def test_long_turn_compaction_prefers_dense_query_cluster_over_late_signal(): + long_turn = ( + "alpha appears once near the beginning. " + + ("unrelated context " * 80) + + "beta gamma beta gamma establish the relevant cluster. " + + ("archive detail " * 80) + + "The record was archived in 2025." + ) + documents = {"doc-a": json.dumps([{"role": "assistant", "content": long_turn}])} + + snippets = select_source_snippets( + documents, + "What did alpha beta gamma establish?", + ["doc-a"], + max_snippets=1, + max_chars_per_snippet=180, + ) + + assert snippets + assert "beta gamma" in snippets[0].text + assert "2025" not in snippets[0].text + + +def test_long_turn_compaction_prefers_query_anchor_over_unrelated_late_signal(): + long_turn = "alpha and omega started at $11. " + ("background detail " * 120) + "The resolved value was $73." + documents = {"doc-a": json.dumps([{"role": "assistant", "content": long_turn}])} + + first = select_source_snippets( + documents, + "What did alpha and omega establish?", + ["doc-a"], + max_snippets=1, + max_chars_per_snippet=180, + ) + second = select_source_snippets( + documents, + "What did alpha and omega establish?", + ["doc-a"], + max_snippets=1, + max_chars_per_snippet=180, + ) + + assert first and second + assert first[0].text == second[0].text + assert "$11" in first[0].text + assert "$73" not in first[0].text + + +def test_long_turn_compaction_keeps_query_amount_before_late_year(): + long_turn = ( + "The fence repair cost $42 and was completed promptly. " + + ("unrelated archive detail " * 120) + + "The archive was migrated in 2025." + ) + snippets = select_source_snippets( + {"doc-a": json.dumps([{"role": "assistant", "content": long_turn}])}, + "How much did the fence repair cost?", + ["doc-a"], + max_snippets=1, + max_chars_per_snippet=180, + ) + + assert snippets + assert "$42" in snippets[0].text + assert "2025" not in snippets[0].text + + +def test_long_turn_compaction_without_query_overlap_keeps_first_signal(): + long_turn = ( + "The recorded value was $42. " + ("unrelated archive detail " * 120) + "The archive was migrated in 2025." + ) + snippets = select_source_snippets( + {"doc-a": json.dumps([{"role": "assistant", "content": long_turn}])}, + "What was the amount?", + ["doc-a"], + max_snippets=1, + max_chars_per_snippet=180, + ) + + assert snippets + assert "$42" in snippets[0].text + assert "2025" not in snippets[0].text + + +def test_long_turn_compaction_keeps_long_query_term_inside_window_boundary(): + long_term = "supercalifragilistic" + long_turn = "alpha " + ("background " * 32) + long_term + snippets = select_source_snippets( + {"doc-a": json.dumps([{"role": "assistant", "content": long_turn}])}, + f"What is the {long_term} value?", + ["doc-a"], + max_snippets=1, + max_chars_per_snippet=100, + ) + + assert snippets + assert long_term in snippets[0].text + + +def test_long_turn_compaction_handles_budget_shorter_than_query_token(): + snippets = select_source_snippets( + {"doc-a": json.dumps([{"role": "assistant", "content": "supercalifragilistic"}])}, + "supercalifragilistic?", + ["doc-a"], + max_snippets=1, + max_chars_per_snippet=18, + ) + + assert snippets + assert len(snippets[0].text) <= 18 + + +def test_missing_chunk_recovery_keeps_exact_chunk_provenance(): + snippets = select_missing_chunk_snippets( + [ + { + "chunk_id": "chunk-7", + "document_id": "doc-a", + "chunk_index": 7, + "chunk_text": json.dumps([{"role": "user", "content": "The repaired fence cost 42 dollars."}]), + } + ], + "How much did the fence repair cost?", + ) + + rendered = render_source_chunk_snippets(snippets) + assert snippets and snippets[0].chunk_id == "chunk-7" + assert "42 dollars" in rendered + assert "chunk_index=7" in rendered + + +def test_long_turn_prefix_is_not_mistaken_for_complete_coverage(): + long_turn = ( + "The fictional exhibit is a clockwork kelpie with " + ("background detail " * 80) + "an amber mosaic tail." + ) + snippets = select_source_snippets( + {"doc-a": json.dumps([{"role": "assistant", "content": long_turn}])}, + "What description was given for the clockwork kelpie's tail?", + ["doc-a"], + covered_chunks={"doc-a": [json.dumps([{"role": "assistant", "content": long_turn[:180]}])]}, + max_snippets=1, + ) + + assert snippets + assert "amber mosaic tail" in snippets[0].text diff --git a/lab/evaluation/benchmarks/longmemeval/test_source_context_integration.py b/lab/evaluation/benchmarks/longmemeval/test_source_context_integration.py new file mode 100644 index 0000000..8d78a5c --- /dev/null +++ b/lab/evaluation/benchmarks/longmemeval/test_source_context_integration.py @@ -0,0 +1,161 @@ +import json +import sys +from types import SimpleNamespace + +import pytest + +from benchmarks.longmemeval.evidence_bundles import RenderedEvidence +from benchmarks.longmemeval.longmemeval_benchmark import LongMemEvalAnswerGenerator + + +def _generator(context_format: str = "structured_source") -> LongMemEvalAnswerGenerator: + """Create a formatter without constructing a real LLM client.""" + + generator = object.__new__(LongMemEvalAnswerGenerator) + generator.context_format = context_format + generator.evidence_mode = None + return generator + + +class _FakeConnection: + def __init__(self, documents, chunks): + self.documents = documents + self.chunks = chunks + self.closed = False + + async def fetch(self, query, *_args): + if "original_text" in query: + return self.documents + return self.chunks + + async def close(self): + self.closed = True + + +class _FakeAsyncpg: + def __init__(self, connection): + self.connection = connection + + async def connect(self, _database_url): + return self.connection + + +def test_source_centric_formatter_reports_only_visible_document_coverage(): + generator = _generator() + chunk_text = json.dumps([{"role": "user", "content": "The preferred color is green."}]) + rendered = generator._format_context_source_centric( + { + "results": [ + { + "id": "fact-1", + "text": "The preferred color is green.", + "fact_type": "preference", + "document_id": "doc-1", + "chunk_id": "chunk-1", + } + ], + "chunks": {"chunk-1": {"chunk_text": chunk_text}}, + }, + "What color do I prefer?", + ) + + assert isinstance(rendered, RenderedEvidence) + assert rendered.text + assert rendered.covered_by_document["doc-1"] + assert rendered.covered_by_document["doc-1"][0] in rendered.text + + +@pytest.mark.asyncio +async def test_backfill_recovers_original_text_when_raw_chunk_misses_bundle_budget(monkeypatch): + generator = _generator() + results = [] + chunks = {} + # The first fourteen bundles nearly fill the source-centric cap. The + # target chunk is still part of recall, but its bundle is not admitted. + for index in range(15): + chunk_id = f"chunk-{index}" + document_id = f"doc-{index}" + results.append( + { + "id": f"fact-{index}", + "text": f"noise {index} " + ("x" * 5000), + "fact_type": "note", + "document_id": document_id, + "chunk_id": chunk_id, + } + ) + chunks[chunk_id] = {"chunk_text": json.dumps([{"role": "user", "content": f"noise {index}"}])} + results.append( + { + "id": "target-fact", + "text": "The target code is blue.", + "fact_type": "preference", + "document_id": "target-doc", + "chunk_id": "target-chunk", + } + ) + chunks["target-chunk"] = {"chunk_text": json.dumps([{"role": "user", "content": "The target code is blue."}])} + recall_result = {"results": results, "chunks": chunks} + rendered = generator._format_context_source_centric(recall_result, "What is the target code?") + assert "target-doc" not in rendered.covered_by_document + + target_document = json.dumps( + [ + {"role": "user", "content": "The target code is blue and must be copied exactly."}, + ] + ) + connection = _FakeConnection( + documents=[{"id": "target-doc", "original_text": target_document}], + chunks=[], + ) + monkeypatch.setitem(sys.modules, "asyncpg", _FakeAsyncpg(connection)) + monkeypatch.setenv("HMS_API_DATABASE_URL", "postgresql://fake") + + source_block = await generator._format_source_document_backfill( + "What is the target code?", + recall_result, + "bank-1", + rendered_coverage=rendered.covered_by_document, + ) + + assert "The target code is blue and must be copied exactly." in source_block + assert "document=target-doc" in source_block + assert connection.closed + + +@pytest.mark.asyncio +async def test_generate_answer_passes_rendered_coverage_to_backfill(monkeypatch): + generator = _generator() + expected_coverage = {"doc-1": ("visible source excerpt",)} + generator._format_context_source_centric = lambda _result, _query: RenderedEvidence( + text="compact context", covered_by_document=expected_coverage + ) + captured = {} + + async def fake_backfill(question, recall_result, bank_id, question_type=None, rendered_coverage=None): + captured.update( + question=question, + recall_result=recall_result, + bank_id=bank_id, + question_type=question_type, + rendered_coverage=rendered_coverage, + ) + return "" + + generator._format_source_document_backfill = fake_backfill + + class _FakeLLM: + async def call(self, **_kwargs): + return SimpleNamespace(answer="ok", reasoning="") + + generator.llm_config = _FakeLLM() + answer, _reasoning, _memories = await generator.generate_answer( + "Which code?", + {"results": [], "chunks": {}}, + question_type="single-session-user", + bank_id="bank-1", + ) + + assert answer == "ok" + assert captured["rendered_coverage"] == expected_coverage + assert captured["bank_id"] == "bank-1" diff --git a/lab/evaluation/pyproject.toml b/lab/evaluation/pyproject.toml index 5112260..cf31e9a 100644 --- a/lab/evaluation/pyproject.toml +++ b/lab/evaluation/pyproject.toml @@ -20,12 +20,13 @@ dependencies = [ [project.optional-dependencies] test = [ "pytest>=8.0.0", + "pytest-asyncio>=0.21.0", "httpx>=0.27.0", "python-dotenv>=1.0.0", ] [tool.hatch.build.targets.wheel] -packages = ["hms_dev", "upgrade_tests"] +packages = ["hms_dev", "benchmarks", "upgrade_tests"] [tool.uv.sources] hms-api = { workspace = true } diff --git a/uv.lock b/uv.lock index b0ea4b9..564e25a 100644 --- a/uv.lock +++ b/uv.lock @@ -2108,6 +2108,7 @@ dependencies = [ test = [ { name = "httpx" }, { name = "pytest" }, + { name = "pytest-asyncio" }, { name = "python-dotenv" }, ] @@ -2125,6 +2126,7 @@ requires-dist = [ { name = "openai", specifier = ">=1.0.0" }, { name = "pydantic", specifier = ">=2.0.0" }, { name = "pytest", marker = "extra == 'test'", specifier = ">=8.0.0" }, + { name = "pytest-asyncio", marker = "extra == 'test'", specifier = ">=0.21.0" }, { name = "python-dotenv", marker = "extra == 'test'", specifier = ">=1.0.0" }, { name = "python-fasthtml", specifier = ">=0.12.33" }, { name = "rich", specifier = ">=13.0.0" }, From 2088295784ffa1ae82d90a9450e645c39106236e Mon Sep 17 00:00:00 2001 From: Dannong Xu Date: Sun, 26 Jul 2026 20:09:05 +0800 Subject: [PATCH 3/8] ci(retain): add database and release quality gates Run offline contracts, PostgreSQL and Oracle lifecycle checks, style validation, benchmark tests, and distribution notice verification with immutable action and service-image references. Refs #1 --- .github/workflows/retain-offline.yml | 236 +++++++++++++++++++++++++++ .github/workflows/retain-oracle.yml | 89 ++++++++++ 2 files changed, 325 insertions(+) create mode 100644 .github/workflows/retain-offline.yml create mode 100644 .github/workflows/retain-oracle.yml diff --git a/.github/workflows/retain-offline.yml b/.github/workflows/retain-offline.yml new file mode 100644 index 0000000..586eef6 --- /dev/null +++ b/.github/workflows/retain-offline.yml @@ -0,0 +1,236 @@ +name: Retain offline quality gate + +on: + pull_request: + push: + branches: + - main + workflow_dispatch: + +permissions: + contents: read + +concurrency: + group: retain-offline-${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +jobs: + offline-quality: + name: Python 3.11 offline checks + runs-on: ubuntu-latest + timeout-minutes: 20 + env: + CI: "true" + PYTHONHASHSEED: "0" + PYTHONPATH: ${{ github.workspace }}/lab/evaluation + UV_NO_PROGRESS: "1" + + steps: + - name: Check out the repository + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + + - name: Set up Python 3.11 + uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0 + with: + python-version: "3.11" + + - name: Set up uv + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 + with: + version: "0.9.3" + enable-cache: true + cache-dependency-glob: uv.lock + + - name: Install locked test dependencies + run: uv sync --locked --package hms-api-slim --extra test --group dev + + - name: Check Retain source style + run: | + uv run --no-sync --package hms-api-slim ruff check \ + --config core/dataplane/pyproject.toml \ + core/dataplane/hms_api/config.py \ + core/dataplane/hms_api/engine/db/ops_postgresql.py \ + core/dataplane/hms_api/engine/db/ops_oracle.py \ + core/dataplane/hms_api/engine/embedding_fingerprint.py \ + core/dataplane/hms_api/engine/entity_resolution_contracts.py \ + core/dataplane/hms_api/engine/entity_resolver.py \ + core/dataplane/hms_api/engine/ingestion \ + core/dataplane/hms_api/engine/memory_engine.py \ + core/dataplane/hms_api/engine/retain/chunk_storage.py \ + core/dataplane/hms_api/engine/retain/embedding_utils.py \ + core/dataplane/hms_api/engine/retain/entity_labels.py \ + core/dataplane/hms_api/engine/retain/entity_processing.py \ + core/dataplane/hms_api/engine/retain/fact_extraction.py \ + core/dataplane/hms_api/engine/retain/fact_storage.py \ + core/dataplane/hms_api/engine/retain/link_utils.py \ + core/dataplane/hms_api/engine/retain/types.py \ + core/dataplane/hms_api/worker/poller.py \ + core/dataplane/tests/test_async_batch_retain.py \ + core/dataplane/tests/test_ingestion_oracle_contracts.py \ + core/dataplane/tests/test_ingestion_oracle_live.py \ + core/dataplane/tests/test_ingestion_pipeline_contracts.py \ + core/dataplane/tests/test_ingestion_postgresql_live.py \ + core/dataplane/tests/test_op_cancellation.py + uv run --no-sync --package hms-api-slim ruff format --check \ + --config core/dataplane/pyproject.toml \ + core/dataplane/hms_api/config.py \ + core/dataplane/hms_api/engine/db/ops_postgresql.py \ + core/dataplane/hms_api/engine/db/ops_oracle.py \ + core/dataplane/hms_api/engine/embedding_fingerprint.py \ + core/dataplane/hms_api/engine/entity_resolution_contracts.py \ + core/dataplane/hms_api/engine/entity_resolver.py \ + core/dataplane/hms_api/engine/ingestion \ + core/dataplane/hms_api/engine/memory_engine.py \ + core/dataplane/hms_api/engine/retain/chunk_storage.py \ + core/dataplane/hms_api/engine/retain/embedding_utils.py \ + core/dataplane/hms_api/engine/retain/entity_labels.py \ + core/dataplane/hms_api/engine/retain/entity_processing.py \ + core/dataplane/hms_api/engine/retain/fact_extraction.py \ + core/dataplane/hms_api/engine/retain/fact_storage.py \ + core/dataplane/hms_api/engine/retain/link_utils.py \ + core/dataplane/hms_api/engine/retain/types.py \ + core/dataplane/hms_api/worker/poller.py \ + core/dataplane/tests/test_async_batch_retain.py \ + core/dataplane/tests/test_ingestion_oracle_contracts.py \ + core/dataplane/tests/test_ingestion_oracle_live.py \ + core/dataplane/tests/test_ingestion_pipeline_contracts.py \ + core/dataplane/tests/test_ingestion_postgresql_live.py \ + core/dataplane/tests/test_op_cancellation.py + + - name: Check LongMemEval source style + run: | + uv run --no-sync --package hms-api-slim ruff check \ + --config lab/evaluation/pyproject.toml \ + lab/evaluation/benchmarks + uv run --no-sync --package hms-api-slim ruff format --check \ + --config lab/evaluation/pyproject.toml \ + lab/evaluation/benchmarks + + - name: Compile and import the changed packages + run: | + uv run --no-sync --package hms-api-slim python -m compileall -q \ + core/dataplane/hms_api/engine/ingestion \ + lab/evaluation/benchmarks + uv run --no-sync --package hms-api-slim python -c \ + "from hms_api.engine.ingestion import RetainPipelineService; from benchmarks.longmemeval import evidence_bundles, source_backfill" + + - name: Validate the benchmark launcher + run: bash -n .aaaSCRIPT/run_benchmark.sh + + - name: Verify distributable license notices + run: | + uv build --package hms-api-slim --out-dir "$RUNNER_TEMP/hms-dist" + uv run --no-sync --package hms-api-slim python - "$RUNNER_TEMP/hms-dist" <<'PY' + import sys + import tarfile + import zipfile + from pathlib import Path + + dist_dir = Path(sys.argv[1]) + required = {"LICENSE", "THIRD_PARTY_NOTICES.md"} + artifacts = sorted( + path + for path in dist_dir.iterdir() + if path.suffix == ".whl" or path.name.endswith(".tar.gz") + ) + if len(artifacts) != 2: + raise SystemExit(f"expected one wheel and one sdist, found: {artifacts}") + for artifact in artifacts: + if artifact.suffix == ".whl": + with zipfile.ZipFile(artifact) as archive: + names = archive.namelist() + elif artifact.name.endswith(".tar.gz"): + with tarfile.open(artifact, "r:gz") as archive: + names = archive.getnames() + else: + raise SystemExit(f"unexpected distribution artifact: {artifact}") + basenames = {Path(name).name for name in names} + missing = required - basenames + if missing: + raise SystemExit(f"{artifact.name} is missing notices: {sorted(missing)}") + PY + grep -Fq "COPY core/dataplane/LICENSE ./api/" deploy/containers/standalone/Dockerfile + grep -Fq "COPY core/dataplane/THIRD_PARTY_NOTICES.md ./api/" deploy/containers/standalone/Dockerfile + + - name: Run Retain offline contract tests + run: | + uv run --no-sync --package hms-api-slim pytest \ + -o "addopts=" \ + -p no:rerunfailures \ + -q \ + core/dataplane/tests/test_ingestion_pipeline_contracts.py \ + core/dataplane/tests/test_ingestion_oracle_contracts.py \ + core/dataplane/tests/test_ingestion_oracle_live.py \ + core/dataplane/tests/test_embedding_fingerprint.py \ + core/dataplane/tests/test_config_validation.py \ + core/dataplane/tests/test_db_abstraction.py \ + core/dataplane/tests/test_fact_extraction_retry.py \ + core/dataplane/tests/test_link_utils.py \ + core/dataplane/tests/test_multimodal_engine_bridge.py \ + core/dataplane/tests/test_multimodal_security.py + + - name: Run LongMemEval offline tests + run: | + uv run --no-sync --package hms-api-slim pytest \ + -o "addopts=" \ + -p no:rerunfailures \ + -q \ + lab/evaluation/benchmarks/common/test_benchmark_runner.py \ + lab/evaluation/benchmarks/longmemeval + + postgresql-live: + name: PostgreSQL 16 Retain smoke + runs-on: ubuntu-latest + timeout-minutes: 20 + services: + postgres: + # Immutable multi-platform digest for pgvector/pgvector:pg16. + image: pgvector/pgvector@sha256:1d533553fefe4f12e5d80c7b80622ba0c382abb5758856f52983d8789179f0fb + env: + POSTGRES_USER: hms + POSTGRES_PASSWORD: hms_test_password + POSTGRES_DB: hms + ports: + - 5432:5432 + options: >- + --health-cmd "pg_isready -U hms -d hms" + --health-interval 5s + --health-timeout 5s + --health-retries 20 + env: + CI: "true" + PYTHONHASHSEED: "0" + UV_NO_PROGRESS: "1" + HMS_API_DATABASE_BACKEND: postgresql + HMS_API_DATABASE_URL: postgresql://hms:hms_test_password@127.0.0.1:5432/hms + HMS_API_LLM_PROVIDER: none + HMS_API_LLM_MODEL: none + + steps: + - name: Check out the repository + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + + - name: Set up Python 3.11 + uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0 + with: + python-version: "3.11" + + - name: Set up uv + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 + with: + version: "0.9.3" + enable-cache: true + cache-dependency-glob: uv.lock + + - name: Install locked PostgreSQL test dependencies + run: uv sync --locked --package hms-api-slim --extra test --group dev + + - name: Run the live PostgreSQL Retain smoke test + run: | + uv run --no-sync --package hms-api-slim pytest \ + -o "addopts=" \ + -p no:rerunfailures \ + -q \ + core/dataplane/tests/test_ingestion_postgresql_live.py \ + core/dataplane/tests/test_async_batch_retain.py \ + core/dataplane/tests/test_op_cancellation.py diff --git a/.github/workflows/retain-oracle.yml b/.github/workflows/retain-oracle.yml new file mode 100644 index 0000000..74e860e --- /dev/null +++ b/.github/workflows/retain-oracle.yml @@ -0,0 +1,89 @@ +name: Retain Oracle integration + +on: + pull_request: + paths: + - ".github/workflows/retain-oracle.yml" + - "pyproject.toml" + - "uv.lock" + - "core/dataplane/pyproject.toml" + - "core/dataplane/hms_api/**" + - "core/dataplane/tests/conftest.py" + - "core/dataplane/tests/test_ingestion_oracle_contracts.py" + - "core/dataplane/tests/test_ingestion_oracle_live.py" + push: + branches: + - main + paths: + - ".github/workflows/retain-oracle.yml" + - "pyproject.toml" + - "uv.lock" + - "core/dataplane/pyproject.toml" + - "core/dataplane/hms_api/**" + - "core/dataplane/tests/conftest.py" + - "core/dataplane/tests/test_ingestion_oracle_contracts.py" + - "core/dataplane/tests/test_ingestion_oracle_live.py" + workflow_dispatch: + +permissions: + contents: read + +concurrency: + group: retain-oracle-${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +jobs: + oracle-live: + name: Oracle 23ai Retain smoke + runs-on: ubuntu-latest + timeout-minutes: 30 + services: + oracle: + # Immutable digest for gvenzl/oracle-free:23-slim-faststart. + image: gvenzl/oracle-free@sha256:d8913e4e4769b6e60197949bef30a4391713afe662b4b4e71a2665c881bdac8b + env: + ORACLE_PASSWORD: OracleTest1 + ports: + - 1521:1521 + options: >- + --health-cmd healthcheck.sh + --health-interval 10s + --health-timeout 5s + --health-retries 30 + env: + CI: "true" + PYTHONHASHSEED: "0" + PYTHONPATH: ${{ github.workspace }}/core/dataplane + UV_NO_PROGRESS: "1" + ORACLE_TEST_DSN: oracle://SYSTEM:OracleTest1@127.0.0.1:1521/FREEPDB1 + HMS_API_DATABASE_BACKEND: oracle + HMS_API_LLM_PROVIDER: none + HMS_API_LLM_MODEL: none + + steps: + - name: Check out the repository + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + + - name: Set up Python 3.11 + uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0 + with: + python-version: "3.11" + + - name: Set up uv + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 + with: + version: "0.9.3" + enable-cache: true + cache-dependency-glob: uv.lock + + - name: Install locked Oracle test dependencies + run: uv sync --locked --package hms-api-slim --extra test --extra oracle --group dev + + - name: Run the live Oracle Retain smoke test + run: | + uv run --no-sync --package hms-api-slim pytest \ + -o "addopts=" \ + -p no:rerunfailures \ + -m oracle \ + -q \ + core/dataplane/tests/test_ingestion_oracle_live.py From df0ec7db6f6b30fb1e3934e20787f88252431138 Mon Sep 17 00:00:00 2001 From: Dannong Xu Date: Sun, 26 Jul 2026 20:16:06 +0800 Subject: [PATCH 4/8] ci(oracle): use image with Oracle Text Run the live Oracle gate on the regular fast-start image because the slim flavor removes the CTXSYS components required by the production migration. Refs #1 --- .github/workflows/retain-oracle.yml | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/.github/workflows/retain-oracle.yml b/.github/workflows/retain-oracle.yml index 74e860e..65f4a99 100644 --- a/.github/workflows/retain-oracle.yml +++ b/.github/workflows/retain-oracle.yml @@ -39,8 +39,9 @@ jobs: timeout-minutes: 30 services: oracle: - # Immutable digest for gvenzl/oracle-free:23-slim-faststart. - image: gvenzl/oracle-free@sha256:d8913e4e4769b6e60197949bef30a4391713afe662b4b4e71a2665c881bdac8b + # Immutable digest for gvenzl/oracle-free:23-faststart. The regular + # flavor is required because the slim flavor removes Oracle Text. + image: gvenzl/oracle-free@sha256:2dcb93d4b50c78a127c0fc40da4c6aafd1d719ff2e545dbb92304b3f40ce407c env: ORACLE_PASSWORD: OracleTest1 ports: From 59767dbd4d5b94e34336bd3dab9649e585fe33f0 Mon Sep 17 00:00:00 2001 From: Dannong Xu Date: Sun, 26 Jul 2026 20:20:52 +0800 Subject: [PATCH 5/8] fix(oracle): execute projection backfill as raw SQL Bypass SQLAlchemy text parsing for the Oracle projection JSON backfill so literal colon-number and colon-boolean tokens are not mistaken for bind parameters. Add an offline migration regression contract. Refs #1 --- ...p4q5r6s7t8u9_add_memory_unit_projection.py | 6 +++-- .../tests/test_ingestion_oracle_contracts.py | 22 +++++++++++++++++++ 2 files changed, 26 insertions(+), 2 deletions(-) diff --git a/core/dataplane/hms_api/alembic/versions/p4q5r6s7t8u9_add_memory_unit_projection.py b/core/dataplane/hms_api/alembic/versions/p4q5r6s7t8u9_add_memory_unit_projection.py index 98c29e7..89e190b 100644 --- a/core/dataplane/hms_api/alembic/versions/p4q5r6s7t8u9_add_memory_unit_projection.py +++ b/core/dataplane/hms_api/alembic/versions/p4q5r6s7t8u9_add_memory_unit_projection.py @@ -8,7 +8,6 @@ from collections.abc import Sequence from alembic import context, op - from hms_api.alembic._dialect import run_for_dialect revision: str = "p4q5r6s7t8u9" @@ -78,7 +77,10 @@ def _oracle_upgrade() -> None: END IF; END; """) - op.execute(f""" + # Bypass SQLAlchemy's text parser for this literal JSON. Otherwise tokens + # such as ``:1`` and ``:true`` inside the payload are treated as bind + # parameters before the statement reaches Oracle. + op.get_bind().exec_driver_sql(f""" UPDATE {schema}memory_units mu SET projection = '{{"embedding":{{"v":1,"ok":' || diff --git a/core/dataplane/tests/test_ingestion_oracle_contracts.py b/core/dataplane/tests/test_ingestion_oracle_contracts.py index 8ad8a8f..94f3bf5 100644 --- a/core/dataplane/tests/test_ingestion_oracle_contracts.py +++ b/core/dataplane/tests/test_ingestion_oracle_contracts.py @@ -47,6 +47,28 @@ from hms_api.engine.retain import chunk_storage +def test_projection_migration_executes_literal_json_as_raw_oracle_sql(monkeypatch) -> None: + """Literal JSON colons must not be parsed as SQLAlchemy bind parameters.""" + + from hms_api.alembic.versions import p4q5r6s7t8u9_add_memory_unit_projection as migration + + alembic_statements: list[str] = [] + driver_statements: list[str] = [] + bind = SimpleNamespace(exec_driver_sql=driver_statements.append) + + monkeypatch.setattr(migration, "_get_schema_prefix", lambda: '"TENANT".') + monkeypatch.setattr(migration.op, "execute", alembic_statements.append) + monkeypatch.setattr(migration.op, "get_bind", lambda: bind) + + migration._oracle_upgrade() + + assert len(alembic_statements) == 1 + assert 'ALTER TABLE "TENANT".memory_units ADD' in alembic_statements[0] + assert len(driver_statements) == 1 + assert 'UPDATE "TENANT".memory_units mu' in driver_statements[0] + assert '"embedding":{"v":1,"ok":' in driver_statements[0] + + class _Transaction: def __init__(self, events: list[str]) -> None: self._events = events From e9b1e91b42607e409d884596b88c16d5c28e98f7 Mon Sep 17 00:00:00 2001 From: Dannong Xu Date: Sun, 26 Jul 2026 20:26:17 +0800 Subject: [PATCH 6/8] fix(oracle): compare projection CLOB safely Use DBMS_LOB.COMPARE during projection backfill because Oracle does not permit direct CLOB equality comparisons. Add an offline regression contract for the migration SQL.\n\nRefs #1 --- .../versions/p4q5r6s7t8u9_add_memory_unit_projection.py | 3 ++- core/dataplane/tests/test_ingestion_oracle_contracts.py | 2 ++ 2 files changed, 4 insertions(+), 1 deletion(-) diff --git a/core/dataplane/hms_api/alembic/versions/p4q5r6s7t8u9_add_memory_unit_projection.py b/core/dataplane/hms_api/alembic/versions/p4q5r6s7t8u9_add_memory_unit_projection.py index 89e190b..c1cc47a 100644 --- a/core/dataplane/hms_api/alembic/versions/p4q5r6s7t8u9_add_memory_unit_projection.py +++ b/core/dataplane/hms_api/alembic/versions/p4q5r6s7t8u9_add_memory_unit_projection.py @@ -8,6 +8,7 @@ from collections.abc import Sequence from alembic import context, op + from hms_api.alembic._dialect import run_for_dialect revision: str = "p4q5r6s7t8u9" @@ -96,7 +97,7 @@ def _oracle_upgrade() -> None: ELSE 'false' END || '}},"extraction":{{"v":"legacy"}}}}' - WHERE projection = '{{}}' + WHERE DBMS_LOB.COMPARE(projection, TO_CLOB('{{}}')) = 0 """) diff --git a/core/dataplane/tests/test_ingestion_oracle_contracts.py b/core/dataplane/tests/test_ingestion_oracle_contracts.py index 94f3bf5..49c81a1 100644 --- a/core/dataplane/tests/test_ingestion_oracle_contracts.py +++ b/core/dataplane/tests/test_ingestion_oracle_contracts.py @@ -67,6 +67,8 @@ def test_projection_migration_executes_literal_json_as_raw_oracle_sql(monkeypatc assert len(driver_statements) == 1 assert 'UPDATE "TENANT".memory_units mu' in driver_statements[0] assert '"embedding":{"v":1,"ok":' in driver_statements[0] + assert "DBMS_LOB.COMPARE(projection, TO_CLOB('{}')) = 0" in driver_statements[0] + assert "WHERE projection = '{}'" not in driver_statements[0] class _Transaction: From 011a0f21baf2bea5cdef87fb1df58e44515af1df Mon Sep 17 00:00:00 2001 From: Dannong Xu Date: Wed, 29 Jul 2026 20:21:13 +0800 Subject: [PATCH 7/8] fix(retain): publish full windows atomically Keep existing documents and all FULL write windows inside one transaction, validate the complete publication before checkpointing, and roll Oracle cancellation back across savepoint and backend scopes. Add contract coverage for cancellation, ownership loss, invalid mappings, and successful publication. Refs #4 --- core/dataplane/hms_api/engine/db/oracle.py | 6 +- .../ingestion/persistence/unit_of_work.py | 123 ++++- .../hms_api/engine/ingestion/service.py | 256 +++++++--- core/dataplane/tests/test_db_abstraction.py | 59 ++- .../test_ingestion_pipeline_contracts.py | 481 +++++++++++++++++- 5 files changed, 841 insertions(+), 84 deletions(-) diff --git a/core/dataplane/hms_api/engine/db/oracle.py b/core/dataplane/hms_api/engine/db/oracle.py index 80f9957..8b7fb86 100644 --- a/core/dataplane/hms_api/engine/db/oracle.py +++ b/core/dataplane/hms_api/engine/db/oracle.py @@ -905,7 +905,7 @@ async def transaction(self) -> AsyncIterator["OracleConnection"]: cursor.close() try: yield self - except Exception: + except BaseException: cursor = self._conn.cursor() await cursor.execute(f"ROLLBACK TO SAVEPOINT {sp_name}") cursor.close() @@ -1319,7 +1319,7 @@ async def acquire(self) -> AsyncIterator[OracleConnection]: # Auto-commit on clean exit (asyncpg uses autocommit by default; # oracledb does not, so we must commit explicitly) await conn.commit() - except Exception: + except BaseException: await conn.rollback() raise finally: @@ -1333,7 +1333,7 @@ async def transaction(self) -> AsyncIterator[OracleConnection]: await self._set_session_schema(conn) yield OracleConnection(conn) await conn.commit() - except Exception: + except BaseException: await conn.rollback() raise finally: diff --git a/core/dataplane/hms_api/engine/ingestion/persistence/unit_of_work.py b/core/dataplane/hms_api/engine/ingestion/persistence/unit_of_work.py index 18255a6..97300b1 100644 --- a/core/dataplane/hms_api/engine/ingestion/persistence/unit_of_work.py +++ b/core/dataplane/hms_api/engine/ingestion/persistence/unit_of_work.py @@ -7,7 +7,7 @@ from __future__ import annotations -from collections.abc import Callable, Mapping, Sequence +from collections.abc import Awaitable, Callable, Mapping, Sequence from contextlib import AbstractAsyncContextManager from dataclasses import dataclass, field from enum import StrEnum @@ -562,6 +562,127 @@ async def write_display_entity_links(self, request: RetainWriteRequest, phase3_p ConnectionScope: TypeAlias = Callable[[], AbstractAsyncContextManager[Any]] +AtomicCommitCallback: TypeAlias = Callable[[Any, tuple[CoreWriteResult, ...]], Awaitable[None]] +AtomicValidationCallback: TypeAlias = Callable[[tuple[CoreWriteResult, ...]], None] + + +@dataclass(frozen=True, slots=True) +class AtomicWriteStep: + """One prepared core write and the adapter that owns its backend semantics.""" + + adapter: PersistenceAdapter + request: RetainWriteRequest + + +class AtomicWriteOwnershipLost(RuntimeError): + """A prepared write lost ownership before the atomic batch could publish.""" + + def __init__(self, window_index: int) -> None: + self.window_index = window_index + super().__init__(f"Atomic Retain write lost ownership at window {window_index}") + + +class AtomicRetainUnitOfWork: + """Publish prepared Retain windows in one transaction. + + Core writes are applied in order on one connection. Any exception, inactive + operation fence, ownership loss, checkpoint failure, outbox failure, or + commit failure rolls every window back. Resolver statistics are flushed + once after commit, followed by independent best-effort display-link work + for each fact-bearing window. + """ + + def __init__(self, *, connection_scope: ConnectionScope) -> None: + if not callable(connection_scope): + raise TypeError("connection_scope must be callable") + self._connection_scope = connection_scope + + async def execute( + self, + steps: Sequence[AtomicWriteStep], + *, + validation_callback: AtomicValidationCallback | None = None, + commit_callback: AtomicCommitCallback | None = None, + ) -> tuple[UnitOfWorkResult, ...]: + prepared = tuple(steps) + if not prepared: + raise ValueError("Atomic Retain execution requires at least one write step") + if any(not isinstance(step, AtomicWriteStep) for step in prepared): + raise TypeError("steps must contain only AtomicWriteStep values") + if validation_callback is not None and not callable(validation_callback): + raise TypeError("validation_callback must be callable or None") + if commit_callback is not None and not callable(commit_callback): + raise TypeError("commit_callback must be callable or None") + + cores: list[CoreWriteResult] = [] + async with self._connection_scope() as connection: + async with connection.transaction(): + for window_index, step in enumerate(prepared): + core = await step.adapter.write_core(connection, step.request) + if core.ownership is OwnershipDisposition.LOST: + raise AtomicWriteOwnershipLost(window_index) + cores.append(core) + immutable_cores = tuple(cores) + if validation_callback is not None: + validation_callback(immutable_cores) + if commit_callback is not None: + await commit_callback(connection, immutable_cores) + + return await self._post_commit(prepared, tuple(cores)) + + @staticmethod + async def _post_commit( + steps: tuple[AtomicWriteStep, ...], + cores: tuple[CoreWriteResult, ...], + ) -> tuple[UnitOfWorkResult, ...]: + fact_window_indices = tuple(index for index, core in enumerate(cores) if core.post_commit_required) + reports: list[PostCommitReport | None] = [ + ( + None + if core.post_commit_required + else PostCommitReport( + status=PostCommitStatus.SKIPPED, + skip_reason=PostCommitSkipReason.NO_FACTS, + ) + ) + for core in cores + ] + + if fact_window_indices: + try: + await steps[fact_window_indices[0]].adapter.flush_entity_stats() + except Exception as exc: + failure = PostCommitFailure(stage=PostCommitStage.ENTITY_STATS, exception=exc) + for index in fact_window_indices: + reports[index] = PostCommitReport( + status=PostCommitStatus.FAILED, + failure=failure, + ) + else: + for index in fact_window_indices: + try: + await steps[index].adapter.write_display_entity_links( + steps[index].request, + cores[index].phase3_payload, + ) + except Exception as exc: + reports[index] = PostCommitReport( + status=PostCommitStatus.FAILED, + failure=PostCommitFailure( + stage=PostCommitStage.DISPLAY_ENTITY_LINKS, + exception=exc, + ), + ) + else: + reports[index] = PostCommitReport(status=PostCommitStatus.COMPLETED) + + if any(report is None for report in reports): # pragma: no cover - exhaustive classification invariant + raise AssertionError("Atomic Retain post-commit report is incomplete") + return tuple( + UnitOfWorkResult(core=core, post_commit=report) + for core, report in zip(cores, reports, strict=True) + if report is not None + ) class RetainUnitOfWork: diff --git a/core/dataplane/hms_api/engine/ingestion/service.py b/core/dataplane/hms_api/engine/ingestion/service.py index 1f8e921..b0c59a9 100644 --- a/core/dataplane/hms_api/engine/ingestion/service.py +++ b/core/dataplane/hms_api/engine/ingestion/service.py @@ -60,8 +60,12 @@ from .persistence.backend import RetainBackendAdapters, retain_backend_adapters from .persistence.models import CommittedUnitBinding, ExistingDocument, OperationCheckpoint from .persistence.unit_of_work import ( + AtomicRetainUnitOfWork, + AtomicWriteOwnershipLost, + AtomicWriteStep, ChunkWrite, CoreGraphWrite, + CoreWriteResult, DeltaWriteRequest, ExistingChunkWrite, FactWrite, @@ -209,6 +213,73 @@ class _ProjectedChunkOutcome: embedding_seconds: float +@dataclass(frozen=True, slots=True) +class _AtomicFullPublication: + unit_ids_by_content: tuple[tuple[str, ...], ...] + committed_unit_ids: tuple[str, ...] + + +def _validate_atomic_full_publication( + document_source_indices: Sequence[int | None], + prepared_records: Sequence[Sequence[MemoryRecord]], + cores: Sequence[CoreWriteResult], +) -> _AtomicFullPublication: + """Validate every persisted FULL result before checkpointing or commit.""" + + record_windows = tuple(tuple(records) for records in prepared_records) + core_results = tuple(cores) + if len(core_results) != len(record_windows): + raise RetainResultMappingError( + "Atomic FULL publication returned a different number of write results than prepared windows" + ) + + window_results: list[tuple[Sequence[MemoryRecord], Sequence[tuple[str, str]]]] = [] + seen_bucket_unit_ids: set[str] = set() + for window_index, (records, core) in enumerate(zip(record_windows, core_results, strict=True)): + bindings = tuple(core.unit_ids_by_fact_key) + bucket_unit_ids = tuple(unit_id for bucket in core.unit_ids_by_content for unit_id in bucket) + binding_unit_ids = tuple(unit_id for _fact_key, unit_id in bindings) + if any(not isinstance(unit_id, str) or not unit_id for unit_id in binding_unit_ids): + raise RetainResultMappingError( + f"Atomic FULL window {window_index} returned an invalid fact-key binding unit ID" + ) + if any(not isinstance(unit_id, str) or not unit_id for unit_id in bucket_unit_ids): + raise RetainResultMappingError( + f"Atomic FULL window {window_index} returned an invalid content-bucket unit ID" + ) + if len(bucket_unit_ids) != len(set(bucket_unit_ids)): + raise RetainResultMappingError( + f"Atomic FULL window {window_index} returned duplicate content-bucket unit IDs" + ) + duplicate_across_windows = seen_bucket_unit_ids.intersection(bucket_unit_ids) + if duplicate_across_windows: + raise RetainResultMappingError("Atomic FULL publication returned duplicate unit IDs across windows") + seen_bucket_unit_ids.update(bucket_unit_ids) + if len(bucket_unit_ids) != len(binding_unit_ids) or set(bucket_unit_ids) != set(binding_unit_ids): + raise RetainResultMappingError( + f"Atomic FULL window {window_index} returned inconsistent content and fact-key unit IDs" + ) + window_results.append((records, bindings)) + + sources = tuple(document_source_indices) + try: + public_buckets = merge_window_unit_ids(sources, tuple(window_results)) + except (TypeError, ValueError) as exc: + raise RetainResultMappingError("Atomic FULL publication returned an invalid fact-key mapping") from exc + + committed_unit_ids: list[str] = [] + for records, bindings in window_results: + units_by_key = dict(bindings) + committed_unit_ids.extend(units_by_key[record.fact_key] for record in records) + + public_iterator = iter(public_buckets) + unit_ids_by_content = tuple(() if source_index is None else next(public_iterator) for source_index in sources) + return _AtomicFullPublication( + unit_ids_by_content=unit_ids_by_content, + committed_unit_ids=tuple(committed_unit_ids), + ) + + def _projection_pipeline_concurrency(config: Any) -> int: """Resolve a finite positive producer width without trusting bool-as-int.""" @@ -812,7 +883,7 @@ async def _execute_full_document_windows( agent_name: str, outbox_callback: Any, ) -> _DocumentOutcome: - """Execute ordered, memory-bounded FULL windows with one finalizer.""" + """Prepare ordered FULL windows, then publish them atomically.""" windows = plan_full_write_windows( plan.chunks, @@ -820,52 +891,25 @@ async def _execute_full_document_windows( ) final_content_hash = compute_document_hash(plan.combined_content) inflight_content_hash = f"{_INFLIGHT_CONTENT_HASH_PREFIX}{uuid.uuid4()}" if len(windows) > 1 else None - window_results: list[tuple[Sequence[MemoryRecord], Sequence[tuple[str, str]]]] = [] - committed_unit_ids: list[str] = [] + prepared_steps: list[AtomicWriteStep] = [] + prepared_records: list[tuple[MemoryRecord, ...]] = [] total_usage = TokenUsage() - for window in windows: - records, extraction_usage = await self._extract_and_project_selected_chunks( - invocation, - execution, - plan, - window.chunks, - agent_name=agent_name, - fact_position_offset=(window.global_indices[0] if window.global_indices else 0), - ) - total_usage = total_usage + extraction_usage - checkpoint_callback = None - if window.is_last: - checkpoint_callback = self._compose_checkpoint_callback( + # Phase 1 may keep task-local resolver statistics. Reset once before + # preparing the complete atomic publication, then let every window + # contribute to the same post-commit flush. + execution.entity_resolver.discard_pending_stats() + try: + for window in windows: + records, extraction_usage = await self._extract_and_project_selected_chunks( invocation, execution, plan, - expected_unit_ids_count=len(committed_unit_ids) + len(records), - prior_unit_ids=tuple(committed_unit_ids), + window.chunks, + agent_name=agent_name, + fact_position_offset=(window.global_indices[0] if window.global_indices else 0), ) - - ownership = _backend_adapters(execution).document_ownership( - schema=execution.schema, - fresh=window.is_first and plan.existing is None, - ) - unit_of_work = RetainUnitOfWork( - connection_scope=lambda: acquire_with_retry(execution.pool), - adapter=PersistenceWriter( - pool=execution.pool, - embeddings_model=execution.embeddings_model, - entity_resolver=execution.entity_resolver, - config=execution.resolved_config, - ownership=ownership, - operation_activity=_backend_adapters(execution).operation_activity_fence( - invocation.operation_id, - schema=execution.schema, - ), - schema=execution.schema, - sanitize_log_identifiers=invocation.sanitize_log_identifiers, - ), - ) - - try: + total_usage = total_usage + extraction_usage request = await self._build_full_window_request( invocation, execution, @@ -876,53 +920,106 @@ async def _execute_full_document_windows( is_last=window.is_last, inflight_content_hash=inflight_content_hash, final_content_hash=final_content_hash, - checkpoint_callback=checkpoint_callback, - outbox_callback=(outbox_callback if window.is_last else None), + reset_pending_stats=False, ) - async with _database_budget(execution.db_semaphore): - result = await unit_of_work.execute(request) - except FreshDocumentOwnershipConflict as exc: - execution.entity_resolver.discard_pending_stats() - raise RetainOwnershipLostError( - "Retain lost fresh-document ownership for " - f"{_log_identifier(invocation, plan.intent.document_id)!r}; retry" - ) from exc - except BaseException: - execution.entity_resolver.discard_pending_stats() - raise - if result.core.ownership is OwnershipDisposition.LOST: - execution.entity_resolver.discard_pending_stats() + ownership = _backend_adapters(execution).document_ownership( + schema=execution.schema, + fresh=window.is_first and plan.existing is None, + ) + prepared_steps.append( + AtomicWriteStep( + adapter=PersistenceWriter( + pool=execution.pool, + embeddings_model=execution.embeddings_model, + entity_resolver=execution.entity_resolver, + config=execution.resolved_config, + ownership=ownership, + operation_activity=_backend_adapters(execution).operation_activity_fence( + invocation.operation_id, + schema=execution.schema, + ), + schema=execution.schema, + sanitize_log_identifiers=invocation.sanitize_log_identifiers, + ), + request=request, + ) + ) + prepared_records.append(tuple(records)) + + checkpoint_callback = self._compose_checkpoint_callback( + invocation, + execution, + plan, + expected_unit_ids_count=sum(len(records) for records in prepared_records), + ) + validated_publication: _AtomicFullPublication | None = None + + def validate_publication(cores: tuple[CoreWriteResult, ...]) -> None: + nonlocal validated_publication + validated_publication = _validate_atomic_full_publication( + tuple(item.source_index for item in plan.intent.items), + prepared_records, + cores, + ) + + async def finalize_publication( + connection: Any, + _cores: tuple[CoreWriteResult, ...], + ) -> None: + if validated_publication is None: # pragma: no cover - UoW callback order invariant + raise AssertionError("Atomic FULL publication was not validated") + if checkpoint_callback is not None: + await checkpoint_callback( + connection, + (validated_publication.committed_unit_ids,), + ) + if outbox_callback is not None: + await outbox_callback(connection) + + atomic_unit_of_work = AtomicRetainUnitOfWork( + connection_scope=lambda: acquire_with_retry(execution.pool), + ) + async with _database_budget(execution.db_semaphore): + results = await atomic_unit_of_work.execute( + prepared_steps, + validation_callback=validate_publication, + commit_callback=finalize_publication, + ) + except FreshDocumentOwnershipConflict as exc: + execution.entity_resolver.discard_pending_stats() + raise RetainOwnershipLostError( + "Retain lost fresh-document ownership for " + f"{_log_identifier(invocation, plan.intent.document_id)!r}; retry" + ) from exc + except AtomicWriteOwnershipLost as exc: + execution.entity_resolver.discard_pending_stats() + raise RetainOwnershipLostError( + "Retain lost document ownership for " + f"{_log_identifier(invocation, plan.intent.document_id)!r} " + f"at FULL window {exc.window_index}; retry" + ) from exc + except BaseException: + execution.entity_resolver.discard_pending_stats() + raise + + if validated_publication is None: # pragma: no cover - UoW callback order invariant + raise AssertionError("Atomic FULL publication committed without validation") + for window, result in zip(windows, results, strict=True): + if result.core.ownership is OwnershipDisposition.LOST: # pragma: no cover - atomic UoW raises raise RetainOwnershipLostError( "Retain lost document ownership for " f"{_log_identifier(invocation, plan.intent.document_id)!r} " f"at FULL window {window.window_index}; retry" ) self._log_post_commit_failure(invocation, plan, result) - bindings = result.core.unit_ids_by_fact_key - window_results.append((records, bindings)) - units_by_key = dict(bindings) - if len(units_by_key) != len(bindings): - raise RetainResultMappingError("FULL window returned duplicate fact-key bindings") - try: - committed_unit_ids.extend(units_by_key[record.fact_key] for record in records) - except KeyError as exc: - raise RetainResultMappingError("FULL window returned an incomplete fact-key binding") from exc - public_buckets = merge_window_unit_ids( - tuple(item.source_index for item in plan.intent.items), - tuple(window_results), - ) - public_iterator = iter(public_buckets) - unit_ids_by_content = tuple( - () if item.source_index is None else next(public_iterator) for item in plan.intent.items - ) - if committed_unit_ids: + if validated_publication.committed_unit_ids: final_ann_completed = await self._run_full_semantic_ann_best_effort( invocation, execution, plan, - committed_unit_ids, + validated_publication.committed_unit_ids, ) if final_ann_completed: await self._record_final_ann_completed_best_effort( @@ -931,7 +1028,7 @@ async def _execute_full_document_windows( plan.intent.document_id, ) return _DocumentOutcome( - unit_ids_by_content=unit_ids_by_content, + unit_ids_by_content=validated_publication.unit_ids_by_content, usage=total_usage, processed_tokens=None, ) @@ -1383,6 +1480,7 @@ async def _build_full_window_request( final_content_hash: str, checkpoint_callback: Any = None, outbox_callback: Any = None, + reset_pending_stats: bool = True, ) -> WriteWindowRequest: retain_params, document_tags = retain_document_metadata(plan.intent.items) payload = await self._build_fact_payload( @@ -1391,6 +1489,7 @@ async def _build_full_window_request( plan, selected_chunks, records, + reset_pending_stats=reset_pending_stats, ) if is_first: document_window = FirstFullWriteWindow( @@ -1502,6 +1601,8 @@ async def _build_fact_payload( plan: _DocumentExecutionPlan, selected_chunks: Sequence[ChunkPlan], records: Sequence[MemoryRecord], + *, + reset_pending_stats: bool = True, ) -> _FactPayload: storage_contents = tuple(content_to_storage(item) for item in plan.intent.items) positions = content_positions(plan.intent.items) @@ -1534,7 +1635,8 @@ async def _build_fact_payload( phase1 = None if processed_facts: - execution.entity_resolver.discard_pending_stats() + if reset_pending_stats: + execution.entity_resolver.discard_pending_stats() phase1_kwargs = { "skip_semantic_ann": ( plan.change.kind is DocumentChangeKind.FULL diff --git a/core/dataplane/tests/test_db_abstraction.py b/core/dataplane/tests/test_db_abstraction.py index 6bd1a9d..677746d 100644 --- a/core/dataplane/tests/test_db_abstraction.py +++ b/core/dataplane/tests/test_db_abstraction.py @@ -4,13 +4,13 @@ without requiring a live database connection. """ +import asyncio import json import uuid from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest - from hms_api.engine.db import DatabaseBackend, DatabaseConnection, create_database_backend from hms_api.engine.db.postgresql import PostgreSQLBackend from hms_api.engine.db.result import DictResultRow as ResultRow @@ -842,3 +842,60 @@ def cursor(self): 'ALTER SESSION SET CURRENT_SCHEMA = "tenant""; DROP TABLE banks; --"', ] assert backend._session_user == "APP_OWNER" + + +class TestOracleCancellationRollback: + """Cancellation is a transaction failure and must never publish writes.""" + + @pytest.mark.asyncio + async def test_connection_savepoint_rolls_back_on_cancellation(self): + from hms_api.engine.db.oracle import OracleConnection + + queries: list[str] = [] + + class Cursor: + async def execute(self, query): + queries.append(query) + + def close(self): + return None + + raw_connection = SimpleNamespace(cursor=Cursor) + with pytest.raises(asyncio.CancelledError): + async with OracleConnection(raw_connection).transaction(): + raise asyncio.CancelledError + + savepoint = queries[0].removeprefix("SAVEPOINT ") + assert queries == [ + f"SAVEPOINT {savepoint}", + f"ROLLBACK TO SAVEPOINT {savepoint}", + ] + + @pytest.mark.asyncio + @pytest.mark.parametrize("scope_name", ["acquire", "transaction"]) + async def test_backend_scope_rolls_back_on_cancellation(self, monkeypatch, scope_name): + from hms_api.engine.db.oracle import OracleBackend + + physical_connection = SimpleNamespace( + commit=AsyncMock(), + rollback=AsyncMock(), + ) + pool = SimpleNamespace( + acquire=AsyncMock(return_value=physical_connection), + release=AsyncMock(), + ) + backend = OracleBackend() + backend._pool = pool + + async def skip_schema(_self, connection): + assert connection is physical_connection + + monkeypatch.setattr(OracleBackend, "_set_session_schema", skip_schema) + + with pytest.raises(asyncio.CancelledError): + async with getattr(backend, scope_name)(): + raise asyncio.CancelledError + + physical_connection.commit.assert_not_awaited() + physical_connection.rollback.assert_awaited_once_with() + pool.release.assert_awaited_once_with(physical_connection) diff --git a/core/dataplane/tests/test_ingestion_pipeline_contracts.py b/core/dataplane/tests/test_ingestion_pipeline_contracts.py index 80f06a5..ece0ae2 100644 --- a/core/dataplane/tests/test_ingestion_pipeline_contracts.py +++ b/core/dataplane/tests/test_ingestion_pipeline_contracts.py @@ -6,7 +6,7 @@ import uuid from contextlib import asynccontextmanager from types import SimpleNamespace -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, Mock import pytest from hms_api.engine import embedding_fingerprint as fingerprint_module @@ -24,13 +24,20 @@ RetainPublicationAborted, ) from hms_api.engine.ingestion import service as service_module -from hms_api.engine.ingestion.domain import DocumentChangeKind +from hms_api.engine.ingestion.domain import ChunkPlan, DocumentChangeKind from hms_api.engine.ingestion.persistence import writer as writer_module from hms_api.engine.ingestion.persistence.operation_fence import OperationActivityFence from hms_api.engine.ingestion.persistence.unit_of_work import ( + AtomicRetainUnitOfWork, + AtomicWriteOwnershipLost, + AtomicWriteStep, CoreGraphWrite, + CoreWriteResult, FirstFullWriteWindow, + LaterFullWriteWindow, MetadataOnlyWriteRequest, + OwnershipDisposition, + PostCommitStatus, RetainUnitOfWork, WriteWindowRequest, ) @@ -1091,6 +1098,476 @@ async def test_inactive_operation_rolls_back_before_fingerprint_or_any_core_writ callback.assert_not_awaited() +class _AtomicStateConnection(_Connection): + def __init__(self, events: list[str]) -> None: + super().__init__(events) + self.state = {"document": "old-version"} + + def transaction(self): + connection = self + + class _Transaction: + async def __aenter__(self): + connection.in_transaction = True + connection._events.append("begin") + self.snapshot = dict(connection.state) + return connection + + async def __aexit__(self, exc_type, _exc, _traceback): + if exc_type is not None: + connection.state = self.snapshot + connection._events.append("rollback") + else: + connection._events.append("commit") + connection.in_transaction = False + return False + + return _Transaction() + + +class _AtomicStateAdapter: + def __init__( + self, + events: list[str], + *, + name: str, + value: str, + core: CoreWriteResult, + failure: Exception | None = None, + ) -> None: + self._events = events + self._name = name + self._value = value + self._core = core + self._failure = failure + + async def write_core(self, connection, _request) -> CoreWriteResult: + assert connection.in_transaction + self._events.append(f"write:{self._name}") + connection.state["document"] = self._value + if self._failure is not None: + raise self._failure + return self._core + + async def flush_entity_stats(self) -> None: + self._events.append(f"flush:{self._name}") + + async def write_display_entity_links(self, _request, _phase3_payload) -> None: + self._events.append(f"display:{self._name}") + + +def _atomic_window_request(*, first: bool) -> WriteWindowRequest: + document_window = ( + FirstFullWriteWindow( + combined_content="replacement", + continuation_content_hash="retain-inflight:test", + ) + if first + else LaterFullWriteWindow( + expected_content_hash="retain-inflight:test", + completed_content_hash="replacement-hash", + ) + ) + return WriteWindowRequest( + bank_id="bank", + document_id="doc", + document_window=document_window, + contents=(RetainContent(content="replacement"),), + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "later_failure", + [ + RetainOperationInactiveError("operation cancelled"), + RuntimeError("later window failed"), + ], + ids=["cancellation", "error"], +) +async def test_atomic_full_windows_restore_old_version_on_later_failure(later_failure) -> None: + events: list[str] = [] + connection = _AtomicStateConnection(events) + first = _AtomicStateAdapter( + events, + name="first", + value="partial-replacement", + core=CoreWriteResult(ownership=OwnershipDisposition.OWNED), + ) + later = _AtomicStateAdapter( + events, + name="later", + value="incomplete-replacement", + core=CoreWriteResult(ownership=OwnershipDisposition.OWNED), + failure=later_failure, + ) + callback = AsyncMock() + + @asynccontextmanager + async def connection_scope(): + yield connection + + unit_of_work = AtomicRetainUnitOfWork(connection_scope=connection_scope) + with pytest.raises(type(later_failure), match=str(later_failure)): + await unit_of_work.execute( + ( + AtomicWriteStep(first, _atomic_window_request(first=True)), + AtomicWriteStep(later, _atomic_window_request(first=False)), + ), + commit_callback=callback, + ) + + assert connection.state == {"document": "old-version"} + assert events == ["begin", "write:first", "write:later", "rollback"] + callback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_existing_two_window_full_service_rolls_back_when_later_fence_observes_cancellation( + monkeypatch, +) -> None: + events: list[str] = [] + connection = _AtomicStateConnection(events) + prepared_indices: list[tuple[int, ...]] = [] + ownership_fresh_values: list[bool] = [] + checkpoint = AsyncMock() + outbox = AsyncMock() + + class CancellingFence: + calls = 0 + + async def assert_active(self, _connection, *, bank_id) -> None: + assert bank_id == "bank" + self.calls += 1 + events.append(f"fence:{self.calls}") + if self.calls == 2: + raise RetainOperationInactiveError("operation cancelled after first window") + + fence = CancellingFence() + + class ServiceWriter: + def __init__(self, *, operation_activity, **_kwargs) -> None: + self._operation_activity = operation_activity + + async def write_core(self, write_connection, request) -> CoreWriteResult: + await self._operation_activity.assert_active( + write_connection, + bank_id=request.bank_id, + ) + assert write_connection.in_transaction + events.append("write:first" if isinstance(request.document_window, FirstFullWriteWindow) else "write:later") + write_connection.state["document"] = ( + "partial-replacement" + if isinstance(request.document_window, FirstFullWriteWindow) + else "complete-replacement" + ) + return CoreWriteResult(ownership=OwnershipDisposition.OWNED) + + async def flush_entity_stats(self) -> None: + raise AssertionError("post-commit work must not run after rollback") + + async def write_display_entity_links(self, _request, _phase3_payload) -> None: + raise AssertionError("post-commit work must not run after rollback") + + backend_adapters = SimpleNamespace( + document_ownership=lambda *, schema, fresh: ownership_fresh_values.append(fresh) or object(), + operation_activity_fence=lambda _operation_id, *, schema: fence, + ) + + @asynccontextmanager + async def connection_scope(_pool): + yield connection + + pipeline = service_module.RetainPipelineService() + + async def extract( + _invocation, + _execution, + _plan, + selected_chunks, + **_kwargs, + ): + prepared_indices.append(tuple(chunk.global_index for chunk in selected_chunks)) + return (), TokenUsage() + + async def build_request( + _invocation, + _execution, + _plan, + _selected_chunks, + _records, + *, + is_first, + **_kwargs, + ): + return _atomic_window_request(first=is_first) + + monkeypatch.setattr(pipeline, "_extract_and_project_selected_chunks", extract) + monkeypatch.setattr(pipeline, "_build_full_window_request", build_request) + monkeypatch.setattr( + pipeline, + "_compose_checkpoint_callback", + lambda *_args, **_kwargs: checkpoint, + ) + monkeypatch.setattr(service_module, "_backend_adapters", lambda _execution: backend_adapters) + monkeypatch.setattr(service_module, "PersistenceWriter", ServiceWriter) + monkeypatch.setattr(service_module, "acquire_with_retry", connection_scope) + + chunks = ( + ChunkPlan( + chunk_key="chunk-0", + source_index=0, + global_index=0, + local_index=0, + text="first", + content_hash="hash-0", + ), + ChunkPlan( + chunk_key="chunk-1", + source_index=0, + global_index=1, + local_index=1, + text="second", + content_hash="hash-1", + ), + ) + plan = SimpleNamespace( + chunks=chunks, + combined_content="replacement", + existing=SimpleNamespace(content_hash="old-hash"), + intent=SimpleNamespace( + document_id="doc", + items=(SimpleNamespace(source_index=0),), + ), + ) + invocation = RetainInvocation( + bank_id="bank", + raw_contents=(), + request_context=object(), + operation_id=str(uuid.uuid4()), + ) + execution = RetainExecutionContext( + pool=SimpleNamespace(backend_type="postgresql"), + embeddings_model=_embedding_model(), + llm_config=object(), + entity_resolver=SimpleNamespace(discard_pending_stats=lambda: events.append("discard")), + format_date_fn=lambda *_args, **_kwargs: "", + resolved_config=_config(retain_chunk_batch_size=1), + ) + + with pytest.raises( + RetainOperationInactiveError, + match="cancelled after first window", + ): + await pipeline._execute_full_document_windows( + invocation, + execution, + plan, + agent_name="agent", + outbox_callback=outbox, + ) + + assert prepared_indices == [(0,), (1,)] + assert ownership_fresh_values == [False, False] + assert connection.state == {"document": "old-version"} + assert events == [ + "discard", + "begin", + "fence:1", + "write:first", + "fence:2", + "rollback", + "discard", + ] + checkpoint.assert_not_awaited() + outbox.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_atomic_full_windows_restore_old_version_on_later_ownership_loss() -> None: + events: list[str] = [] + connection = _AtomicStateConnection(events) + first = _AtomicStateAdapter( + events, + name="first", + value="partial-replacement", + core=CoreWriteResult(ownership=OwnershipDisposition.OWNED), + ) + later = _AtomicStateAdapter( + events, + name="later", + value="incomplete-replacement", + core=CoreWriteResult(ownership=OwnershipDisposition.LOST), + ) + validation = Mock() + callback = AsyncMock() + + @asynccontextmanager + async def connection_scope(): + yield connection + + with pytest.raises(AtomicWriteOwnershipLost, match="window 1"): + await AtomicRetainUnitOfWork(connection_scope=connection_scope).execute( + ( + AtomicWriteStep(first, _atomic_window_request(first=True)), + AtomicWriteStep(later, _atomic_window_request(first=False)), + ), + validation_callback=validation, + commit_callback=callback, + ) + + assert connection.state == {"document": "old-version"} + assert events == ["begin", "write:first", "write:later", "rollback"] + validation.assert_not_called() + callback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_atomic_full_windows_restore_old_version_on_invalid_global_result_mapping() -> None: + events: list[str] = [] + connection = _AtomicStateConnection(events) + first_core = CoreWriteResult( + ownership=OwnershipDisposition.OWNED, + unit_ids_by_content=(("duplicate-unit",),), + unit_ids_by_fact_key=(("fact-first", "duplicate-unit"),), + ) + later_core = CoreWriteResult( + ownership=OwnershipDisposition.OWNED, + unit_ids_by_content=(("duplicate-unit",),), + unit_ids_by_fact_key=(("fact-later", "duplicate-unit"),), + ) + first = _AtomicStateAdapter( + events, + name="first", + value="partial-replacement", + core=first_core, + ) + later = _AtomicStateAdapter( + events, + name="later", + value="incomplete-replacement", + core=later_core, + ) + callback = AsyncMock() + + @asynccontextmanager + async def connection_scope(): + yield connection + + def validate(cores) -> None: + assert connection.in_transaction + events.append("validate") + service_module._validate_atomic_full_publication( + (0, 1), + ( + (SimpleNamespace(fact_key="fact-first", source_index=0),), + (SimpleNamespace(fact_key="fact-later", source_index=1),), + ), + cores, + ) + + with pytest.raises( + service_module.RetainResultMappingError, + match="duplicate unit IDs across windows", + ): + await AtomicRetainUnitOfWork(connection_scope=connection_scope).execute( + ( + AtomicWriteStep(first, _atomic_window_request(first=True)), + AtomicWriteStep(later, _atomic_window_request(first=False)), + ), + validation_callback=validate, + commit_callback=callback, + ) + + assert connection.state == {"document": "old-version"} + assert events == ["begin", "write:first", "write:later", "validate", "rollback"] + callback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_atomic_full_windows_publish_together_and_finalize_once() -> None: + events: list[str] = [] + connection = _AtomicStateConnection(events) + first_core = CoreWriteResult( + ownership=OwnershipDisposition.OWNED, + unit_ids_by_content=(("unit-first",),), + unit_ids_by_fact_key=(("fact-first", "unit-first"),), + phase3_payload="phase-first", + post_commit_required=True, + ) + later_core = CoreWriteResult( + ownership=OwnershipDisposition.OWNED, + unit_ids_by_content=(("unit-later",),), + unit_ids_by_fact_key=(("fact-later", "unit-later"),), + phase3_payload="phase-later", + post_commit_required=True, + ) + first = _AtomicStateAdapter( + events, + name="first", + value="partial-replacement", + core=first_core, + ) + later = _AtomicStateAdapter( + events, + name="later", + value="complete-replacement", + core=later_core, + ) + + @asynccontextmanager + async def connection_scope(): + yield connection + + async def finalize(commit_connection, cores) -> None: + assert commit_connection is connection + assert commit_connection.in_transaction + assert commit_connection.state == {"document": "complete-replacement"} + assert cores == (first_core, later_core) + events.append("finalize") + + def validate(cores) -> None: + assert connection.in_transaction + assert cores == (first_core, later_core) + publication = service_module._validate_atomic_full_publication( + (0,), + ( + (SimpleNamespace(fact_key="fact-first", source_index=0),), + (SimpleNamespace(fact_key="fact-later", source_index=0),), + ), + cores, + ) + assert publication.unit_ids_by_content == (("unit-first", "unit-later"),) + assert publication.committed_unit_ids == ("unit-first", "unit-later") + events.append("validate") + + results = await AtomicRetainUnitOfWork(connection_scope=connection_scope).execute( + ( + AtomicWriteStep(first, _atomic_window_request(first=True)), + AtomicWriteStep(later, _atomic_window_request(first=False)), + ), + validation_callback=validate, + commit_callback=finalize, + ) + + assert connection.state == {"document": "complete-replacement"} + assert [result.post_commit.status for result in results] == [ + PostCommitStatus.COMPLETED, + PostCommitStatus.COMPLETED, + ] + assert events == [ + "begin", + "write:first", + "write:later", + "validate", + "finalize", + "commit", + "flush:first", + "display:first", + "display:later", + ] + + @pytest.mark.asyncio async def test_unhashed_existing_document_uses_locked_full_replacement(monkeypatch) -> None: events: list[str] = [] From 675f4d3b1cba42b1f447fa59533e4aac6aa802f6 Mon Sep 17 00:00:00 2001 From: Dannong Xu Date: Wed, 29 Jul 2026 20:21:56 +0800 Subject: [PATCH 8/8] fix(evaluation): verify reused bank provenance Audit selected LongMemEval banks against normalized document hashes and retained metadata before QA, mark reused or mixed Retain creator identity as unverifiable, and enable retrieval trace diagnostics without exposing trace to the answer generator. Cover sequential, parallel, shared-bank, only-ingested, and force-reingest paths. Refs #4 --- .../benchmarks/common/benchmark_runner.py | 469 +++++++++++++----- .../common/test_benchmark_runner.py | 441 +++++++++++++++- .../benchmarks/longmemeval/README.md | 24 +- .../longmemeval/longmemeval_benchmark.py | 104 +++- .../longmemeval/test_release_integrity.py | 129 ++++- 5 files changed, 1000 insertions(+), 167 deletions(-) diff --git a/lab/evaluation/benchmarks/common/benchmark_runner.py b/lab/evaluation/benchmarks/common/benchmark_runner.py index 79aebf2..b4cd4fb 100644 --- a/lab/evaluation/benchmarks/common/benchmark_runner.py +++ b/lab/evaluation/benchmarks/common/benchmark_runner.py @@ -34,6 +34,8 @@ # Configure logging from environment variable get_config().configure_logging() +from hms_api.engine.ingestion.adapters.storage_records import compute_document_hash, retain_document_metadata +from hms_api.engine.ingestion.normalization import normalize_content_item from hms_api.engine.memory_engine import Budget from hms_api.engine.schema import fq_table from hms_api.models import RequestContext @@ -152,12 +154,20 @@ class IngestionIntegrityError(RuntimeError): def __init__(self, report: Dict[str, Any]): self.report = report item_id = report.get("item_id", "unknown") - missing = report.get("missing_documents", []) - empty = report.get("documents_without_chunks", []) - super().__init__( - f"Durable ingestion audit failed for item {item_id!r}: " - f"missing_documents={missing}, documents_without_chunks={empty}" + failure_fields = ( + "missing_documents", + "documents_without_chunks", + "unexpected_documents", + "inflight_documents", + "content_hash_mismatches", + "context_mismatches", + "event_date_mismatches", + "unverifiable_documents", + "invalid_retain_params", + "failed_banks", ) + failures = {field: report.get(field) for field in failure_fields if report.get(field)} + super().__init__(f"Durable ingestion audit failed for item {item_id!r}: {failures}") def _write_json_atomic(payload: Dict[str, Any], output_path: Path) -> None: @@ -261,13 +271,54 @@ def get_model_config() -> Dict[str, Dict[str, str]]: } -def print_model_config(): - """Print the model configuration to console.""" +def get_artifact_model_config( + *, + retain_executed: bool = True, + retain_execution: Optional[str] = None, +) -> Dict[str, Dict[str, str]]: + """Return model identities that truthfully describe stages executed in this run. + + A reused memory bank predates the current benchmark process. Its Retain + model identity cannot be reconstructed from durable document rows, so the + current Retain environment must not be presented as the bank creator. + """ + + execution = retain_execution or ("executed" if retain_executed else "not_executed") + if execution not in {"executed", "not_executed", "partial_or_skipped"}: + raise ValueError(f"Unsupported Retain execution mode: {execution}") + config = get_model_config() + if execution == "not_executed": + config["retain"] = { + "execution": "not_executed", + "bank_creator_identity": "unverifiable", + } + elif execution == "partial_or_skipped": + config["retain"] = { + "execution": "partial_or_skipped", + "bank_creator_identity": "mixed_or_unverifiable", + } + return config + + +def format_retain_model_config(config: Mapping[str, str]) -> str: + """Render an executed, reused, or mixed Retain identity without overclaiming.""" + + if "provider" in config and "model" in config: + return f"{config['provider']}/{config['model']}" + if config.get("execution") == "partial_or_skipped": + return "partially executed or skipped; bank creator identity is mixed or unverifiable" + return "not executed; reused-bank creator identity is unverifiable" + + +def print_model_config(config: Optional[Mapping[str, Mapping[str, str]]] = None): + """Print the model configuration to console.""" + config = config or get_model_config() + retain_config = config.get("retain", {}) console.print("\n[bold cyan]Model Configuration:[/bold cyan]") console.print(f" HMS: {config['hms']['provider']}/{config['hms']['model']}") - console.print(f" Retain: {config['retain']['provider']}/{config['retain']['model']}") + console.print(f" Retain: {format_retain_model_config(retain_config)}") console.print( f" Answer Generation: {config['answer_generation']['provider']}/{config['answer_generation']['model']}" ) @@ -1057,6 +1108,7 @@ async def answer_question( query_rewriting_strategy_name=recall_query_rewriting_strategy, query_rewriting_enabled=recall_query_rewriting_enabled, session_expansion_weight=recall_session_expansion_weight, + enable_trace=True, ) recall_time = time.time() - recall_start_time @@ -1065,8 +1117,10 @@ async def answer_question( num_chunks = len(search_result.chunks) if search_result.chunks else 0 num_entities = len(search_result.entities) if search_result.entities else 0 - # Convert entire RecallResult to dictionary for answer generation - recall_result_dict = search_result.model_dump() + # Keep the detailed trace local to benchmark diagnostics. It can be + # large and contains retrieval internals that must not influence the + # answer generator. + recall_result_dict = search_result.model_dump(exclude={"trace"}) if plan.evidence_appendix_mode == "cross_session": recall_result_dict = add_cross_session_evidence_appendix(recall_result_dict) elif plan.evidence_appendix_mode == "cross_session_compact": @@ -1412,42 +1466,85 @@ async def judge_single(result): "detailed_results": judged_results, } - async def _agent_has_data(self, agent_id: str) -> bool: - """ - Check if an agent has any indexed memory units. - - Args: - agent_id: Agent ID to check - - Returns: - True if agent has at least one memory unit, False otherwise - """ - try: - # A bank is reusable only when it has durable source chunks. A - # document may legitimately produce no extracted facts. - pool = await self.memory._get_pool() - async with pool.acquire() as conn: - result = await conn.fetchval( - f"SELECT EXISTS(SELECT 1 FROM {fq_table('chunks')} WHERE bank_id = $1)", - agent_id, - ) - return bool(result) - except Exception as e: - console.print(f" [red]Warning: Error checking agent data: {e}[/red]") - return False + def _expected_document_snapshots( + self, + item: Dict[str, Any], + ) -> Tuple[Dict[str, Dict[str, Optional[str]]], List[str]]: + """Build the same content and metadata identities used by Retain.""" - async def _audit_durable_ingestion(self, item: Dict[str, Any], agent_id: str) -> Dict[str, Any]: - """Verify that every input document has at least one durable source chunk. + prepared = self.dataset.prepare_sessions_for_ingestion(item) + normalized = tuple( + normalize_content_item(content, source_index=source_index) for source_index, content in enumerate(prepared) + ) + explicit_ids = tuple( + dict.fromkeys(content.document_id for content in normalized if content.document_id is not None) + ) + grouped: Dict[str, List[Any]] = {} + unverifiable: List[str] = [] + + if len(explicit_ids) == 1: + # Retain assigns missing-ID items to the sole explicit document. + grouped[explicit_ids[0]] = list(normalized) + elif len(explicit_ids) > 1: + for content in normalized: + if content.document_id is None: + unverifiable.append(f"source_index:{content.source_index}") + continue + grouped.setdefault(content.document_id, []).append(content) + elif normalized: + # Retain generates a random document ID when the whole batch omits + # IDs. Such a bank cannot be safely matched to a later dataset item. + unverifiable.extend(f"source_index:{content.source_index}" for content in normalized) + + snapshots: Dict[str, Dict[str, Optional[str]]] = {} + for document_id, contents in grouped.items(): + if any(content.update_mode.value == "append" for content in contents): + # The final content hash also depends on pre-existing text, + # which is not part of the submitted benchmark item. + unverifiable.append(document_id) + continue + retain_params, _ = retain_document_metadata(tuple(contents)) + snapshots[document_id] = { + "content_hash": compute_document_hash("\n".join(content.content for content in contents)), + "context": retain_params.get("context"), + "event_date": retain_params.get("event_date"), + } + return snapshots, unverifiable + + @staticmethod + def _retain_params_mapping(value: Any) -> Dict[str, Any]: + """Normalize PostgreSQL/Oracle JSON representations to a mapping.""" + + if value is None: + return {} + if isinstance(value, Mapping): + return dict(value) + if isinstance(value, bytes): + value = value.decode("utf-8") + if isinstance(value, str): + parsed = json.loads(value) + if isinstance(parsed, Mapping): + return dict(parsed) + raise ValueError(f"retain_params must be a JSON object, got {type(value).__name__}") + + async def _audit_durable_ingestion( + self, + item: Dict[str, Any], + agent_id: str, + *, + reject_unexpected_documents: bool = True, + allowed_document_ids: Optional[Iterable[str]] = None, + ) -> Dict[str, Any]: + """Verify a bank is the exact durable Retain output expected for an item. - Fact extraction is lossy by design, so a document with zero facts is - reported but remains valid. Missing documents and zero-chunk documents - are integrity failures because recall cannot recover their source text. + Zero extracted facts remain valid because source chunks are recallable. + Missing/empty documents, stale content or Retain metadata, unexpected + documents, and in-flight writes make an item-scoped bank unsafe to + reuse. """ - prepared = self.dataset.prepare_sessions_for_ingestion(item) - expected_document_ids = list( - dict.fromkeys(str(content["document_id"]) for content in prepared if content.get("document_id") is not None) - ) + expected, unverifiable = self._expected_document_snapshots(item) + expected_document_ids = list(expected) report: Dict[str, Any] = { "item_id": self.dataset.get_item_id(item), "bank_id": agent_id, @@ -1456,47 +1553,168 @@ async def _audit_durable_ingestion(self, item: Dict[str, Any], agent_id: str) -> "missing_documents": [], "documents_without_chunks": [], "documents_without_facts": [], + "unexpected_documents": [], + "additional_documents": [], + "inflight_documents": [], + "content_hash_mismatches": [], + "context_mismatches": [], + "event_date_mismatches": [], + "unverifiable_documents": unverifiable, + "invalid_retain_params": [], } - if not expected_document_ids: - return report pool = await self.memory._get_pool() async with pool.acquire() as conn: rows = await conn.fetch( f""" SELECT d.id, - COUNT(DISTINCT c.chunk_id) AS chunk_count, - COUNT(DISTINCT m.id) AS fact_count + d.content_hash, + d.retain_params, + ( + SELECT COUNT(*) + FROM {fq_table("chunks")} AS c + WHERE c.bank_id = d.bank_id AND c.document_id = d.id + ) AS chunk_count, + ( + SELECT COUNT(*) + FROM {fq_table("memory_units")} AS m + WHERE m.bank_id = d.bank_id AND m.document_id = d.id + ) AS fact_count FROM {fq_table("documents")} AS d - LEFT JOIN {fq_table("chunks")} AS c - ON c.bank_id = d.bank_id AND c.document_id = d.id - LEFT JOIN {fq_table("memory_units")} AS m - ON m.bank_id = d.bank_id AND m.document_id = d.id - WHERE d.bank_id = $1 AND d.id = ANY($2::text[]) - GROUP BY d.id + WHERE d.bank_id = $1 """, agent_id, - expected_document_ids, ) by_id = {str(row["id"]): row for row in rows} - report["durable_documents"] = len(by_id) - report["missing_documents"] = [document_id for document_id in expected_document_ids if document_id not in by_id] - report["documents_without_chunks"] = [ - document_id - for document_id in expected_document_ids - if document_id in by_id and int(by_id[document_id]["chunk_count"]) == 0 - ] - report["documents_without_facts"] = [ - document_id - for document_id in expected_document_ids - if document_id in by_id and int(by_id[document_id]["fact_count"]) == 0 - ] + observed_ids = set(by_id) + expected_ids = set(expected_document_ids) + report["missing_documents"] = sorted(expected_ids - observed_ids) + allowed_ids = set(allowed_document_ids) if allowed_document_ids is not None else expected_ids + unexpected_documents = sorted(observed_ids - allowed_ids) + if reject_unexpected_documents or allowed_document_ids is not None: + report["unexpected_documents"] = unexpected_documents + else: + report["additional_documents"] = unexpected_documents + + for document_id, row in by_id.items(): + content_hash = str(row["content_hash"] or "") + if content_hash.startswith("retain-inflight:"): + report["inflight_documents"].append(document_id) + if document_id not in expected: + continue + + if int(row["chunk_count"] or 0) == 0: + report["documents_without_chunks"].append(document_id) + if int(row["fact_count"] or 0) == 0: + report["documents_without_facts"].append(document_id) + + snapshot = expected[document_id] + if content_hash != snapshot["content_hash"]: + report["content_hash_mismatches"].append( + { + "document_id": document_id, + "expected": snapshot["content_hash"], + "actual": content_hash, + } + ) + + try: + retain_params = self._retain_params_mapping(row["retain_params"]) + except (TypeError, ValueError, json.JSONDecodeError): + report["invalid_retain_params"].append(document_id) + continue + + if retain_params.get("context") != snapshot["context"]: + report["context_mismatches"].append( + { + "document_id": document_id, + "expected": snapshot["context"], + "actual": retain_params.get("context"), + } + ) + if retain_params.get("event_date") != snapshot["event_date"]: + report["event_date_mismatches"].append( + { + "document_id": document_id, + "expected": snapshot["event_date"], + "actual": retain_params.get("event_date"), + } + ) + + invalid_document_ids = { + *report["missing_documents"], + *report["documents_without_chunks"], + *report["inflight_documents"], + *report["invalid_retain_params"], + } + for mismatch_field in ("content_hash_mismatches", "context_mismatches", "event_date_mismatches"): + invalid_document_ids.update(mismatch["document_id"] for mismatch in report[mismatch_field]) + report["durable_documents"] = sum( + 1 for document_id in expected_document_ids if document_id not in invalid_document_ids + ) - if report["missing_documents"] or report["documents_without_chunks"]: + failure_fields = ( + "missing_documents", + "documents_without_chunks", + "unexpected_documents", + "inflight_documents", + "content_hash_mismatches", + "context_mismatches", + "event_date_mismatches", + "unverifiable_documents", + "invalid_retain_params", + ) + if any(report[field] for field in failure_fields): raise IngestionIntegrityError(report) return report + async def _preflight_reusable_items( + self, + items: Iterable[Dict[str, Any]], + agent_id: str, + *, + clear_agent_per_item: bool, + require_all: bool, + ) -> set[str]: + """Return item IDs with an exact reusable bank, before any QA starts.""" + + reusable_item_ids: set[str] = set() + failures: List[IngestionIntegrityError] = [] + item_list = list(items) + shared_allowed_document_ids: Optional[set[str]] = None + if not clear_agent_per_item: + shared_allowed_document_ids = set() + for item in item_list: + snapshots, _ = self._expected_document_snapshots(item) + shared_allowed_document_ids.update(snapshots) + for item in item_list: + item_id = self.dataset.get_item_id(item) + item_agent_id = f"{agent_id}_{item_id}" if clear_agent_per_item else agent_id + try: + await self._audit_durable_ingestion( + item, + item_agent_id, + reject_unexpected_documents=True, + allowed_document_ids=shared_allowed_document_ids, + ) + except IngestionIntegrityError as exc: + failures.append(exc) + else: + reusable_item_ids.add(item_id) + + if failures and require_all: + failed_banks = [str(exc.report.get("bank_id", "unknown")) for exc in failures] + raise IngestionIntegrityError( + { + "item_id": "reuse-preflight", + "bank_id": agent_id, + "failed_banks": failed_banks, + "bank_reports": [exc.report for exc in failures], + } + ) + return reusable_item_ids + async def process_single_item( self, item: Dict, @@ -1514,6 +1732,7 @@ async def process_single_item( ingest_only: bool = False, skip_if_already_ingested: bool = False, force_reingest: bool = False, + bank_is_item_scoped: bool = True, ) -> Dict: """ Process a single item (ingest + evaluate). @@ -1537,8 +1756,16 @@ async def process_single_item( # Check if already ingested (for smart resume) already_ingested = False if skip_if_already_ingested and not force_reingest: - already_ingested = await self._agent_has_data(agent_id) - if already_ingested: + try: + await self._audit_durable_ingestion( + item, + agent_id, + reject_unexpected_documents=bank_is_item_scoped, + ) + except IngestionIntegrityError: + already_ingested = False + else: + already_ingested = True console.print(f" [{step}] [yellow]⊘[/yellow] Skipping - already ingested") if not already_ingested or force_reingest: @@ -1563,7 +1790,11 @@ async def process_single_item( step += 1 console.print(f" [{step}] Auditing durable ingestion...") - ingestion_audit = await self._audit_durable_ingestion(item, agent_id) + ingestion_audit = await self._audit_durable_ingestion( + item, + agent_id, + reject_unexpected_documents=bank_is_item_scoped, + ) console.print( " [green]✓[/green] " f"{ingestion_audit['durable_documents']}/{ingestion_audit['expected_documents']} " @@ -1641,6 +1872,7 @@ async def run( force_reingest: bool = False, # If True, always re-ingest even if data already exists rerun_invalid_existing: bool = False, # Resume mode: rerun invalid existing item results run_manifest: Optional[Dict[str, Any]] = None, # Stable metadata included in every checkpoint + model_config: Optional[Mapping[str, Mapping[str, str]]] = None, # Executed-stage model identities ) -> Dict[str, Any]: """ Run the full benchmark evaluation. @@ -1676,10 +1908,12 @@ async def run( raise ValueError(f"{name} must be a positive integer, got {value}") self._run_manifest = run_manifest + selected_model_config = model_config or get_model_config() + self._model_config = {role: dict(settings) for role, settings in selected_model_config.items()} self._diagnostic_cache_dir = output_path.parent / "retrieval_cache" if output_path is not None else None # Print model configuration - print_model_config() + print_model_config(self._model_config) # Load dataset console.print(f"\n[1] Loading dataset from {dataset_path}...") @@ -1895,7 +2129,7 @@ async def _run_single_phase( "total_valid": total_valid, "valid_only_accuracy": valid_only_accuracy, "num_items": len(all_results), - "model_config": get_model_config(), + "model_config": getattr(self, "_model_config", None) or get_model_config(), "item_results": all_results, } @@ -1936,26 +2170,24 @@ async def _process_items_sequential( # Pre-load durable banks for reuse and ingest-only fill modes. ingested_item_ids = set() if skip_ingestion or ingest_only: - console.print("[cyan]Checking which items are already ingested...[/cyan]") - for item in items: - item_id = self.dataset.get_item_id(item) - item_agent_id = f"{agent_id}_{item_id}" if clear_agent_per_item else agent_id - if await self._agent_has_data(item_agent_id): - ingested_item_ids.add(item_id) + console.print("[cyan]Auditing retained banks against selected dataset items...[/cyan]") + ingested_item_ids = await self._preflight_reusable_items( + items, + agent_id, + clear_agent_per_item=clear_agent_per_item, + require_all=skip_ingestion, + ) if ingested_item_ids: - console.print(f"[cyan]Found {len(ingested_item_ids)} items already ingested[/cyan]") + console.print(f"[cyan]Found {len(ingested_item_ids)} exact reusable item banks[/cyan]") else: - console.print("[cyan]No items found with existing ingest data[/cyan]") - if skip_ingestion: + console.print("[cyan]No exact reusable item banks found[/cyan]") + if ingest_only and not skip_ingestion and not clear_agent_per_item: requested_item_ids = {self.dataset.get_item_id(item) for item in items} - missing_item_ids = sorted(requested_item_ids - ingested_item_ids) - if missing_item_ids: - preview = ", ".join(missing_item_ids[:10]) - suffix = " ..." if len(missing_item_ids) > 10 else "" - raise RuntimeError( - "Retrieval-only mode requires durable retained chunks for every selected item; " - f"missing {len(missing_item_ids)} bank(s): {preview}{suffix}" - ) + if ingested_item_ids != requested_item_ids: + # Repairing any item in a shared bank may clear or replace + # state needed by otherwise reusable items. Rebuild the + # selected union together instead of skipping a stale mix. + ingested_item_ids.clear() for i, item in enumerate(items, 1): item_id = self.dataset.get_item_id(item) @@ -1970,19 +2202,14 @@ async def _process_items_sequential( # Only clear on first item for shared agent_id clear_this_agent = i == 1 - # Skip items without existing ingest data (only applies when using --skip-ingestion) - # When not using --skip-ingestion, we should ingest the items - if skip_ingestion and item_id not in ingested_item_ids and not ingest_only: - raise RuntimeError(f"Retrieval-only bank disappeared before evaluation: {item_agent_id}") - # For ingest_only mode, skip already ingested items - if ingest_only and not skip_ingestion and item_id in ingested_item_ids: + if ingest_only and not skip_ingestion and not force_reingest and item_id in ingested_item_ids: console.print(f"\n[bold blue]Item {i}/{len(items)}[/bold blue] (ID: {item_id})") console.print(" [yellow]⊘[/yellow] Skipping - already ingested") continue # Then check fill status (results file) - if filln: + if filln and not force_reingest: skip_item_ids = resume_complete_item_ids if rerun_invalid_existing else existing_item_ids if item_id in skip_item_ids: console.print(f"\n[bold blue]Item {i}/{len(items)}[/bold blue] (ID: {item_id})") @@ -2005,6 +2232,7 @@ async def _process_items_sequential( ingest_only, skip_if_already_ingested=False, force_reingest=force_reingest, + bank_is_item_scoped=clear_agent_per_item, ) # Replace existing result or append new one @@ -2061,26 +2289,17 @@ async def _process_items_parallel( # ingest-only run. ingested_item_ids = set() if skip_ingestion or ingest_only: - console.print("[cyan]Checking which items are already ingested...[/cyan]") - for item in items: - item_id = self.dataset.get_item_id(item) - item_agent_id = f"{agent_id}_{item_id}" - if await self._agent_has_data(item_agent_id): - ingested_item_ids.add(item_id) + console.print("[cyan]Auditing retained banks against selected dataset items...[/cyan]") + ingested_item_ids = await self._preflight_reusable_items( + items, + agent_id, + clear_agent_per_item=True, + require_all=skip_ingestion, + ) if ingested_item_ids: - console.print(f"[cyan]Found {len(ingested_item_ids)} items already ingested[/cyan]") + console.print(f"[cyan]Found {len(ingested_item_ids)} exact reusable item banks[/cyan]") else: - console.print("[cyan]No items found with existing ingest data[/cyan]") - if skip_ingestion: - requested_item_ids = {self.dataset.get_item_id(item) for item in items} - missing_item_ids = sorted(requested_item_ids - ingested_item_ids) - if missing_item_ids: - preview = ", ".join(missing_item_ids[:10]) - suffix = " ..." if len(missing_item_ids) > 10 else "" - raise RuntimeError( - "Retrieval-only mode requires durable retained chunks for every selected item; " - f"missing {len(missing_item_ids)} bank(s): {preview}{suffix}" - ) + console.print("[cyan]No exact reusable item banks found[/cyan]") # Create semaphore for item-level parallelism item_semaphore = asyncio.Semaphore(max_concurrent_items) @@ -2091,19 +2310,14 @@ async def process_item_wrapper(i: int, item: Dict) -> Optional[Dict]: item_id = self.dataset.get_item_id(item) item_agent_id = f"{agent_id}_{item_id}" - # Only reuse mode requires data to exist before this item runs. - # A fresh parallel run must proceed to ingestion. - if skip_ingestion and item_id not in ingested_item_ids and not ingest_only: - raise RuntimeError(f"Retrieval-only bank disappeared before evaluation: {item_agent_id}") - # For ingest_only mode, skip already ingested items - if ingest_only and not skip_ingestion and item_id in ingested_item_ids: + if ingest_only and not skip_ingestion and not force_reingest and item_id in ingested_item_ids: console.print(f"\n[bold blue]Item {i}/{len(items)}[/bold blue] (ID: {item_id})") console.print(" [yellow]⊘[/yellow] Skipping - already ingested") return None # Then check fill status (results file) - if filln: + if filln and not force_reingest: skip_item_ids = resume_complete_item_ids if rerun_invalid_existing else existing_item_ids if item_id in skip_item_ids: console.print(f"\n[bold blue]Item {i}/{len(items)}[/bold blue] (ID: {item_id})") @@ -2127,6 +2341,7 @@ async def process_item_wrapper(i: int, item: Dict) -> Optional[Dict]: ingest_only=ingest_only, skip_if_already_ingested=False, force_reingest=force_reingest, + bank_is_item_scoped=True, ) return result @@ -2254,6 +2469,14 @@ async def _run_two_phase( else: console.print("\n[3] Skipping ingestion (using existing data)") + console.print(" [cyan]Auditing retained documents against all selected items...[/cyan]") + await self._preflight_reusable_items( + items, + agent_id, + clear_agent_per_item=False, + require_all=True, + ) + # Phase 2: Evaluation console.print("\n[5] Phase 2: Evaluating all questions...") @@ -2310,6 +2533,7 @@ async def _run_two_phase( "total_valid": total_valid, "valid_only_accuracy": valid_only_accuracy, "num_items": len(all_results), + "model_config": getattr(self, "_model_config", None) or get_model_config(), "item_results": all_results, } @@ -2323,7 +2547,8 @@ def display_results(self, results: Dict[str, Any]): console.print("[bold cyan]Model Configuration:[/bold cyan]") console.print(f" HMS: {config['hms']['provider']}/{config['hms']['model']}") if "retain" in config: - console.print(f" Retain: {config['retain']['provider']}/{config['retain']['model']}") + retain_config = config["retain"] + console.print(f" Retain: {format_retain_model_config(retain_config)}") console.print( f" Answer Generation: {config['answer_generation']['provider']}/{config['answer_generation']['model']}" ) @@ -2450,7 +2675,7 @@ def _save_incremental_results(self, all_results: List[Dict], output_path: Path): "total_valid": total_valid, "valid_only_accuracy": valid_only_accuracy, "num_items": len(ordered_results), - "model_config": get_model_config(), + "model_config": getattr(self, "_model_config", None) or get_model_config(), "item_results": ordered_results, } run_manifest = getattr(self, "_run_manifest", None) diff --git a/lab/evaluation/benchmarks/common/test_benchmark_runner.py b/lab/evaluation/benchmarks/common/test_benchmark_runner.py index 19c8d9a..f110309 100644 --- a/lab/evaluation/benchmarks/common/test_benchmark_runner.py +++ b/lab/evaluation/benchmarks/common/test_benchmark_runner.py @@ -1,33 +1,97 @@ import asyncio import json +from datetime import datetime, timezone from pathlib import Path from types import SimpleNamespace from unittest.mock import AsyncMock, Mock import pytest +from hms_api.engine.ingestion.adapters.storage_records import compute_document_hash from benchmarks.common.benchmark_runner import ( BenchmarkRunner, + IngestionIntegrityError, LLMAnswerEvaluator, _embedding_runtime_config, _endpoint_fingerprint, _reranker_runtime_config, _write_json_atomic, + get_artifact_model_config, get_model_config, ) +_EVENT_DATE = datetime(2024, 1, 2, 3, 4, 5, tzinfo=timezone.utc) + class _Dataset: def get_item_id(self, item): return item["id"] def prepare_sessions_for_ingestion(self, item): - return [{"content": item["id"], "document_id": f"document-{item['id']}"}] + return [ + { + "content": item["id"], + "context": f"context-{item['id']}", + "event_date": _EVENT_DATE, + "document_id": f"document-{item['id']}", + } + ] def get_qa_pairs(self, item): return [{"question": item["id"], "answer": "answer", "category": "test"}] +class _Acquire: + def __init__(self, connection): + self.connection = connection + + async def __aenter__(self): + return self.connection + + async def __aexit__(self, exc_type, exc, traceback): + return False + + +class _Pool: + def __init__(self, rows): + self.connection = SimpleNamespace(fetch=AsyncMock(return_value=rows)) + + def acquire(self): + return _Acquire(self.connection) + + +def _runner_with_document_rows(rows): + runner = object.__new__(BenchmarkRunner) + runner.dataset = _Dataset() + pool = _Pool(rows) + runner.memory = SimpleNamespace(_get_pool=AsyncMock(return_value=pool)) + return runner + + +def _document_row( + item_id: str, + *, + content_hash: str | None = None, + retain_params=None, + chunks: int = 1, + facts: int = 0, +): + return { + "id": f"document-{item_id}", + "content_hash": content_hash if content_hash is not None else compute_document_hash(item_id), + "retain_params": ( + retain_params + if retain_params is not None + else { + "context": f"context-{item_id}", + "event_date": _EVENT_DATE.isoformat(), + } + ), + "chunk_count": chunks, + "fact_count": facts, + } + + def test_embedding_runtime_config_uses_provider_specific_identity(monkeypatch): monkeypatch.setenv("HMS_API_EMBEDDINGS_PROVIDER", "cohere") monkeypatch.setenv("HMS_API_EMBEDDINGS_COHERE_MODEL", "embed-v4.0") @@ -73,12 +137,87 @@ def test_model_config_uses_retain_provider_default_when_only_provider_is_overrid assert config["retain"]["model"] == "gpt-4o-mini" +def test_reused_bank_model_config_does_not_claim_current_retain_identity(monkeypatch): + monkeypatch.setenv("HMS_API_RETAIN_LLM_PROVIDER", "openai") + monkeypatch.setenv("HMS_API_RETAIN_LLM_MODEL", "current-but-not-bank-creator") + + config = get_artifact_model_config(retain_executed=False) + + assert config["retain"] == { + "execution": "not_executed", + "bank_creator_identity": "unverifiable", + } + + +def test_mixed_ingest_only_model_config_does_not_claim_one_global_retain_identity(): + config = get_artifact_model_config(retain_execution="partial_or_skipped") + + assert config["retain"] == { + "execution": "partial_or_skipped", + "bank_creator_identity": "mixed_or_unverifiable", + } + + +@pytest.mark.asyncio +async def test_ingestion_audit_accepts_exact_document_with_zero_facts(): + runner = _runner_with_document_rows([_document_row("a", facts=0)]) + + report = await runner._audit_durable_ingestion({"id": "a"}, "longmemeval_a") + + assert report["durable_documents"] == 1 + assert report["documents_without_facts"] == ["document-a"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("row", "failure_field"), + [ + (_document_row("a", content_hash="stale"), "content_hash_mismatches"), + ( + _document_row( + "a", + retain_params={"context": "stale-context", "event_date": _EVENT_DATE.isoformat()}, + ), + "context_mismatches", + ), + ( + _document_row( + "a", + retain_params={"context": "context-a", "event_date": "2020-01-01T00:00:00+00:00"}, + ), + "event_date_mismatches", + ), + (_document_row("a", chunks=0), "documents_without_chunks"), + ], +) +async def test_ingestion_audit_rejects_stale_or_unqueryable_document(row, failure_field): + runner = _runner_with_document_rows([row]) + + with pytest.raises(IngestionIntegrityError) as exc_info: + await runner._audit_durable_ingestion({"id": "a"}, "longmemeval_a") + + assert exc_info.value.report[failure_field] + + +@pytest.mark.asyncio +async def test_ingestion_audit_rejects_unexpected_and_inflight_documents(): + inflight = _document_row("a", content_hash="retain-inflight:operation") + unexpected = _document_row("other") + runner = _runner_with_document_rows([inflight, unexpected]) + + with pytest.raises(IngestionIntegrityError) as exc_info: + await runner._audit_durable_ingestion({"id": "a"}, "longmemeval_a") + + assert exc_info.value.report["inflight_documents"] == ["document-a"] + assert exc_info.value.report["unexpected_documents"] == ["document-other"] + + @pytest.mark.asyncio async def test_fresh_parallel_run_processes_items_before_any_bank_exists(): runner = object.__new__(BenchmarkRunner) runner.dataset = _Dataset() runner.template_path = None - runner._agent_has_data = AsyncMock(return_value=False) + runner._preflight_reusable_items = AsyncMock() async def process_item(item, *args, **kwargs): return { @@ -105,7 +244,7 @@ async def process_item(item, *args, **kwargs): assert {result["item_id"] for result in results} == {"a", "b"} assert runner.process_single_item.await_count == 2 - runner._agent_has_data.assert_not_awaited() + runner._preflight_reusable_items.assert_not_awaited() @pytest.mark.asyncio @@ -115,7 +254,6 @@ async def test_fresh_item_reingests_even_when_a_bank_already_exists(): runner.template_path = None runner.memory = Mock() runner.memory.delete_bank = AsyncMock() - runner._agent_has_data = AsyncMock(return_value=True) runner.ingest_conversation = AsyncMock(return_value=1) runner._audit_durable_ingestion = AsyncMock( return_value={ @@ -143,7 +281,6 @@ async def test_fresh_item_reingests_even_when_a_bank_already_exists(): skip_if_already_ingested=False, ) - runner._agent_has_data.assert_not_awaited() runner.memory.delete_bank.assert_awaited_once() runner.ingest_conversation.assert_awaited_once() @@ -274,7 +411,6 @@ async def test_resume_skips_valid_items_and_reruns_invalid_items(tmp_path: Path) ) runner = object.__new__(BenchmarkRunner) runner.dataset = _Dataset() - runner._agent_has_data = AsyncMock(return_value=False) async def process_item(item, *args, **kwargs): return { @@ -324,10 +460,21 @@ async def process_item(item, *args, **kwargs): async def test_retrieval_only_fails_when_any_selected_bank_is_missing(): runner = object.__new__(BenchmarkRunner) runner.dataset = _Dataset() - runner._agent_has_data = AsyncMock(side_effect=[True, False]) + runner._audit_durable_ingestion = AsyncMock( + side_effect=[ + {"durable_documents": 1}, + IngestionIntegrityError( + { + "item_id": "b", + "bank_id": "longmemeval_b", + "missing_documents": ["document-b"], + } + ), + ] + ) runner.process_single_item = AsyncMock() - with pytest.raises(RuntimeError, match="missing 1 bank"): + with pytest.raises(IngestionIntegrityError, match="reuse-preflight"): await runner._process_items_parallel( items=[{"id": "a"}, {"id": "b"}], agent_id="longmemeval", @@ -341,9 +488,283 @@ async def test_retrieval_only_fails_when_any_selected_bank_is_missing(): max_concurrent_items=2, ) + assert runner._audit_durable_ingestion.await_count == 2 + runner.process_single_item.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_sequential_retrieval_only_preflights_every_bank_before_qa(): + runner = object.__new__(BenchmarkRunner) + runner.dataset = _Dataset() + runner._audit_durable_ingestion = AsyncMock( + side_effect=[ + {"durable_documents": 1}, + IngestionIntegrityError( + { + "item_id": "b", + "bank_id": "longmemeval_b", + "content_hash_mismatches": [{"document_id": "document-b"}], + } + ), + ] + ) + runner.process_single_item = AsyncMock() + + with pytest.raises(IngestionIntegrityError, match="reuse-preflight"): + await runner._process_items_sequential( + items=[{"id": "a"}, {"id": "b"}], + agent_id="longmemeval", + thinking_budget=10, + max_tokens=100, + skip_ingestion=True, + max_questions_per_item=None, + question_semaphore=asyncio.Semaphore(1), + eval_semaphore=asyncio.Semaphore(1), + clear_agent_per_item=True, + filln=False, + ) + + assert runner._audit_durable_ingestion.await_count == 2 runner.process_single_item.assert_not_awaited() +@pytest.mark.asyncio +async def test_ingest_only_reingests_nonmatching_bank_and_skips_exact_bank(): + runner = object.__new__(BenchmarkRunner) + runner.dataset = _Dataset() + runner._preflight_reusable_items = AsyncMock(return_value={"a"}) + runner.process_single_item = AsyncMock( + return_value={ + "item_id": "b", + "metrics": {"correct": 0, "total": 0, "invalid": 0}, + "num_sessions": 1, + } + ) + + results = await runner._process_items_sequential( + items=[{"id": "a"}, {"id": "b"}], + agent_id="longmemeval", + thinking_budget=10, + max_tokens=100, + skip_ingestion=False, + max_questions_per_item=None, + question_semaphore=asyncio.Semaphore(1), + eval_semaphore=asyncio.Semaphore(1), + clear_agent_per_item=True, + filln=False, + ingest_only=True, + ) + + assert [result["item_id"] for result in results] == ["b"] + assert runner.process_single_item.await_count == 1 + assert runner.process_single_item.await_args.args[0] == {"id": "b"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("ingest_only", [False, True]) +async def test_force_reingest_overrides_fill_skip_sequentially(tmp_path: Path, ingest_only: bool): + output_path = tmp_path / "results.json" + output_path.write_text( + json.dumps( + { + "item_results": [ + { + "item_id": "a", + "metrics": {"correct": 0, "total": 0, "invalid": 0}, + "num_sessions": 0, + } + ] + } + ), + encoding="utf-8", + ) + runner = object.__new__(BenchmarkRunner) + runner.dataset = _Dataset() + runner._preflight_reusable_items = AsyncMock(return_value={"a"}) + runner._save_incremental_results = Mock() + runner.process_single_item = AsyncMock( + return_value={ + "item_id": "a", + "metrics": {"correct": 0, "total": 0, "invalid": 0}, + "num_sessions": 1, + } + ) + + results = await runner._process_items_sequential( + items=[{"id": "a"}], + agent_id="longmemeval", + thinking_budget=10, + max_tokens=100, + skip_ingestion=False, + max_questions_per_item=None, + question_semaphore=asyncio.Semaphore(1), + eval_semaphore=asyncio.Semaphore(1), + clear_agent_per_item=True, + filln=True, + output_path=output_path, + merge_with_existing=True, + ingest_only=ingest_only, + force_reingest=True, + ) + + assert [result["item_id"] for result in results] == ["a"] + assert runner.process_single_item.await_args.kwargs["force_reingest"] is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("ingest_only", [False, True]) +async def test_force_reingest_overrides_fill_skip_in_parallel(tmp_path: Path, ingest_only: bool): + output_path = tmp_path / "results.json" + output_path.write_text( + json.dumps( + { + "item_results": [ + { + "item_id": "a", + "metrics": {"correct": 0, "total": 0, "invalid": 0}, + "num_sessions": 0, + } + ] + } + ), + encoding="utf-8", + ) + runner = object.__new__(BenchmarkRunner) + runner.dataset = _Dataset() + runner._preflight_reusable_items = AsyncMock(return_value={"a"}) + runner._save_incremental_results = Mock() + runner.process_single_item = AsyncMock( + return_value={ + "item_id": "a", + "metrics": {"correct": 0, "total": 0, "invalid": 0}, + "num_sessions": 1, + } + ) + + results = await runner._process_items_parallel( + items=[{"id": "a"}], + agent_id="longmemeval", + thinking_budget=10, + max_tokens=100, + skip_ingestion=False, + max_questions_per_item=None, + question_semaphore=asyncio.Semaphore(1), + eval_semaphore=asyncio.Semaphore(1), + filln=True, + max_concurrent_items=1, + output_path=output_path, + merge_with_existing=True, + ingest_only=ingest_only, + force_reingest=True, + ) + + assert [result["item_id"] for result in results] == ["a"] + assert runner.process_single_item.await_args.kwargs["force_reingest"] is True + + +@pytest.mark.asyncio +async def test_two_phase_shared_bank_rejects_documents_outside_selected_union_before_qa(): + runner = _runner_with_document_rows( + [ + _document_row("a"), + _document_row("b"), + _document_row("stale"), + ] + ) + runner.evaluate_qa_task = AsyncMock() + + with pytest.raises(IngestionIntegrityError) as exc_info: + await runner._run_two_phase( + items=[{"id": "a"}, {"id": "b"}], + agent_id="shared", + thinking_budget=10, + max_tokens=100, + skip_ingestion=True, + max_questions_per_item=None, + max_concurrent_questions=1, + eval_semaphore_size=1, + ) + + assert "document-stale" in str(exc_info.value.report["bank_reports"]) + runner.evaluate_qa_task.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_recall_trace_populates_diagnostics_but_is_not_sent_to_answer_generator(): + fact = SimpleNamespace( + id="fact-1", + document_id="document-a", + context="context-a", + occurred_start=_EVENT_DATE.isoformat(), + fact_type="world", + metadata={"proof_count": 2}, + model_dump=lambda: { + "id": "fact-1", + "document_id": "document-a", + "text": "remembered fact", + }, + ) + trace = { + "rrf_merged": [ + { + "node_id": "fact-1", + "text": "remembered fact", + "rrf_score": 0.5, + "final_rrf_rank": 1, + } + ], + "reranked": [ + { + "node_id": "fact-1", + "text": "remembered fact", + "rerank_score": 0.9, + "rerank_rank": 1, + "rrf_rank": 1, + "score_components": {"combined_score": 0.8}, + } + ], + } + + class _RecallResult: + results = [fact] + chunks = None + entities = None + + def __init__(self): + self.trace = trace + + def model_dump(self, *, exclude=None): + payload = {"results": [fact.model_dump()], "trace": trace} + for field in exclude or set(): + payload.pop(field, None) + return payload + + generator = Mock() + generator.needs_external_search.return_value = True + generator.generate_answer = AsyncMock(return_value=("answer", "reasoning", None)) + runner = object.__new__(BenchmarkRunner) + runner.answer_generator = generator + runner.retrieval_planner = None + runner.query_rewriting_enabled = False + runner.query_rewriting_strategy_name = "noop" + runner.session_expansion_weight = 0.3 + runner.memory = SimpleNamespace( + recall_async=AsyncMock(return_value=_RecallResult()), + _cross_encoder=SimpleNamespace(model_name="reranker", provider_name="test"), + ) + + _, _, _, _, retrieval_details = await runner.answer_question( + "longmemeval_a", + "What happened?", + ) + + assert runner.memory.recall_async.await_args.kwargs["enable_trace"] is True + generator_payload = generator.generate_answer.await_args.args[1] + assert "trace" not in generator_payload + assert retrieval_details["coarse_search_results"].total_candidates == 1 + assert len(retrieval_details["reranked_results"].reranked_candidates) == 1 + + def test_atomic_json_write_replaces_the_complete_artifact(tmp_path: Path): output_path = tmp_path / "result.json" output_path.write_text('{"stale": true}\n', encoding="utf-8") @@ -356,7 +777,8 @@ def test_atomic_json_write_replaces_the_complete_artifact(tmp_path: Path): def test_incremental_checkpoint_preserves_run_manifest(tmp_path: Path): runner = object.__new__(BenchmarkRunner) - runner._run_manifest = {"artifact_schema_version": 1, "dataset": {"sha256": "abc"}} + runner._run_manifest = {"artifact_schema_version": 2, "dataset": {"sha256": "abc"}} + runner._model_config = get_artifact_model_config(retain_executed=False) output_path = tmp_path / "result.json" runner._save_incremental_results( @@ -372,3 +794,4 @@ def test_incremental_checkpoint_preserves_run_manifest(tmp_path: Path): saved = json.loads(output_path.read_text(encoding="utf-8")) assert saved["run_manifest"] == runner._run_manifest + assert saved["model_config"]["retain"]["bank_creator_identity"] == "unverifiable" diff --git a/lab/evaluation/benchmarks/longmemeval/README.md b/lab/evaluation/benchmarks/longmemeval/README.md index 99154dc..d6e76dc 100644 --- a/lab/evaluation/benchmarks/longmemeval/README.md +++ b/lab/evaluation/benchmarks/longmemeval/README.md @@ -169,10 +169,26 @@ HMS_RETRIEVAL_ONLY=1 \ bash .aaaSCRIPT/run_benchmark.sh ``` -The runner requires durable chunks for every selected bank before retrieval-only -evaluation; missing banks fail the run instead of silently reducing the sample. -The strict embedding fingerprint policy rejects memories built in another -vector space. +Before any question is evaluated, the runner matches every selected bank to the +current dataset item. It verifies the complete document-ID set, normalized +content hashes, retained `context` and `event_date`, durable chunks, and that no +Retain write is still in flight. Missing, stale, empty, or unexpected documents +fail retrieval-only mode instead of silently reusing the wrong state. Documents +with zero extracted facts remain valid when their source chunks are durable. +The strict embedding fingerprint policy separately rejects memories built in +another vector space. + +Durable rows do not record enough information to reconstruct the Retain +pipeline, model, or source revision that originally created an older bank. +Retrieval-only artifacts therefore omit `retain` from +`run_manifest.pipeline.stages`, mark the reused-bank Retain provenance as +`unverifiable`, and do not present the current Retain model configuration as the +bank creator. The recorded Git/source identity applies only to stages executed +by the current benchmark process. Fresh runs record `current_run` Retain +provenance. Because a non-forced ingest-only run can skip exact existing banks +and ingest only the missing or stale subset, its global Retain provenance is +marked `mixed_or_reused_bank` and unverifiable. Use `--force-reingest` when the +artifact must attest that every selected bank was created by the current run. ## Reproduction profiles diff --git a/lab/evaluation/benchmarks/longmemeval/longmemeval_benchmark.py b/lab/evaluation/benchmarks/longmemeval/longmemeval_benchmark.py index 8de6212..e22ef64 100644 --- a/lab/evaluation/benchmarks/longmemeval/longmemeval_benchmark.py +++ b/lab/evaluation/benchmarks/longmemeval/longmemeval_benchmark.py @@ -29,7 +29,8 @@ LLMAnswerEvaluator, LLMAnswerGenerator, RecallPlan, - get_model_config, + format_retain_model_config, + get_artifact_model_config, ) from benchmarks.longmemeval.evidence_bundles import ( RenderedEvidence, @@ -172,8 +173,34 @@ def build_run_manifest( git_commit = _git_value("rev-parse", "HEAD") git_status = _git_value("status", "--porcelain") planner = "self_evolution" if oracle_planner_v220 else "ledger" if oracle_planner_v26 else "standard" + if skip_ingestion: + executed_stages = [] if ingest_only else ["recall", "answer", "judge"] + ingestion_provenance = { + "mode": "reused_bank", + "status": "unverifiable", + "retain_execution": "not_executed", + "content_identity": "verified_at_runtime", + "unverifiable_fields": ["retain_pipeline", "retain_model", "retain_code"], + } + elif ingest_only and not force_reingest: + executed_stages = ["retain"] + ingestion_provenance = { + "mode": "mixed_or_reused_bank", + "status": "unverifiable", + "retain_execution": "partial_or_skipped", + "content_identity": "verified_at_runtime", + "unverifiable_fields": ["reused_retain_pipeline", "reused_retain_model", "reused_retain_code"], + } + else: + executed_stages = ["retain"] if ingest_only else ["retain", "recall", "answer", "judge"] + ingestion_provenance = { + "mode": "current_run", + "status": "current_run", + "retain_execution": "executed", + "content_identity": "verified_after_ingestion", + } return { - "artifact_schema_version": 1, + "artifact_schema_version": 2, "dataset": { "path": _manifest_dataset_reference(dataset_path), "sha256": dataset_sha256, @@ -181,7 +208,7 @@ def build_run_manifest( "expected_full_items": LONGMEMEVAL_EXPECTED_ITEMS if canonical else None, }, "pipeline": { - "stages": ["retain", "recall", "answer", "judge"], + "stages": executed_stages, "planner": planner, "context_format": context_format, "thinking_budget": thinking_budget, @@ -208,12 +235,14 @@ def build_run_manifest( "ingest_only": ingest_only, "force_reingest": force_reingest, }, + "ingestion_provenance": ingestion_provenance, "database": { "backend": os.getenv("HMS_API_DATABASE_BACKEND", "postgresql"), "schema": os.getenv("HMS_API_DATABASE_SCHEMA", "public"), "vector_extension": os.getenv("HMS_API_VECTOR_EXTENSION", "pgvector"), }, "runtime": { + "identity_scope": "executed_stages", "git_commit": git_commit, "git_dirty": bool(git_status) if git_status is not None else None, "source_tree_fingerprint": _source_tree_fingerprint(), @@ -242,6 +271,7 @@ def _artifact_compatibility_contract( "artifact_schema_version": manifest.get("artifact_schema_version"), "dataset": dataset_identity, "pipeline": manifest.get("pipeline"), + "ingestion_provenance": manifest.get("ingestion_provenance"), "database": manifest.get("database"), "git_commit": runtime.get("git_commit") if isinstance(runtime, Mapping) else None, "source_tree_fingerprint": (runtime.get("source_tree_fingerprint") if isinstance(runtime, Mapping) else None), @@ -303,6 +333,24 @@ def validate_output_target( ) +async def _filter_items_with_reusable_banks( + items: Sequence[Dict[str, Any]], + *, + dataset: BenchmarkDataset, + runner: BenchmarkRunner, + agent_id: str = "longmemeval", +) -> List[Dict[str, Any]]: + """Filter ``--only-ingested`` items through the shared exact-bank preflight.""" + + reusable_item_ids = await runner._preflight_reusable_items( + items, + agent_id, + clear_agent_per_item=True, + require_all=False, + ) + return [item for item in items if dataset.get_item_id(item) in reusable_item_ids] + + ORACLE_PLANNER_V1_WEIGHTS = { "single-session-user": 0.25, "single-session-assistant": 0.25, @@ -2011,7 +2059,13 @@ async def run_benchmark( ingest_only=ingest_only, force_reingest=force_reingest, ) - current_model_config = get_model_config() + if skip_ingestion or only_ingested: + retain_execution = "not_executed" + elif ingest_only and not force_reingest: + retain_execution = "partial_or_skipped" + else: + retain_execution = "executed" + current_model_config = get_artifact_model_config(retain_execution=retain_execution) validate_output_target( output_path, merge_with_existing=merge_with_existing, @@ -2061,29 +2115,21 @@ async def run_benchmark( items_to_check = filtered_items if filtered_items is not None else original_dataset_items - # Check which items have existing banks - ingested_items = [] - pool = await memory._get_pool() - - for item in items_to_check: - item_id = dataset.get_item_id(item) - agent_id = f"longmemeval_{item_id}" - - # A retained bank is reusable when its source chunks are durable; - # fact extraction may legitimately produce zero facts. - async with pool.acquire() as conn: - has_chunks = await conn.fetchval( - f"SELECT EXISTS(SELECT 1 FROM {fq_table('chunks')} WHERE bank_id = $1)", - agent_id, - ) - if has_chunks: - ingested_items.append(item) - - filtered_items = ingested_items - console.print(f"[green]Found {len(filtered_items)} items with existing memory banks[/green]") + audit_runner = BenchmarkRunner( + dataset=dataset, + answer_generator=answer_generator, + answer_evaluator=answer_evaluator, + memory=memory, + ) + filtered_items = await _filter_items_with_reusable_banks( + items_to_check, + dataset=dataset, + runner=audit_runner, + ) + console.print(f"[green]Found {len(filtered_items)} exact reusable memory banks[/green]") if not filtered_items: - raise RuntimeError("No items with durable retained source chunks were found") + raise RuntimeError("No exact dataset-matching reusable memory banks were found") # Determine query rewriting strategy if query_expansion_enabled: @@ -2187,6 +2233,7 @@ def filtered_load(path: Path, max_items: Optional[int] = None): force_reingest=force_reingest, # Force re-ingest even if data already exists rerun_invalid_existing=resume, run_manifest=current_manifest, + model_config=current_model_config, ) results["run_manifest"] = current_manifest runner.save_results(results, output_path) @@ -2195,7 +2242,10 @@ def filtered_load(path: Path, max_items: Optional[int] = None): console.print("\n[green]✓[/green] Ingest-only mode completed. Data is ready for evaluation.") console.print(" To run evaluation later with a different model:") console.print(" 1. Update .env with your preferred model") - console.print(" 2. Run: HMS_BENCHMARK=longmemeval bash .aaaSCRIPT/run_benchmark.sh --only-ingested --fill") + console.print( + " 2. Choose a new HMS_RESULTS_FILENAME and run: " + "HMS_BENCHMARK=longmemeval bash .aaaSCRIPT/run_benchmark.sh --only-ingested" + ) return results full_run_requested = ( @@ -2378,7 +2428,7 @@ def generate_markdown_table(results: dict, json_output_path: Path): lines.append("") lines.append(f"- **HMS**: {config['hms']['provider']}/{config['hms']['model']}") if "retain" in config: - lines.append(f"- **Retain**: {config['retain']['provider']}/{config['retain']['model']}") + lines.append(f"- **Retain**: {format_retain_model_config(config['retain'])}") lines.append( f"- **Answer Generation**: {config['answer_generation']['provider']}/{config['answer_generation']['model']}" ) diff --git a/lab/evaluation/benchmarks/longmemeval/test_release_integrity.py b/lab/evaluation/benchmarks/longmemeval/test_release_integrity.py index 467e7ca..e8f9c5b 100644 --- a/lab/evaluation/benchmarks/longmemeval/test_release_integrity.py +++ b/lab/evaluation/benchmarks/longmemeval/test_release_integrity.py @@ -4,6 +4,7 @@ import os import subprocess from pathlib import Path +from types import SimpleNamespace from unittest.mock import AsyncMock import pytest @@ -28,7 +29,7 @@ async def test_answer_provider_failure_propagates_to_the_runner(): def _manifest() -> dict: return { - "artifact_schema_version": 1, + "artifact_schema_version": 2, "dataset": { "path": "datasets/dataset.json", "sha256": "dataset-sha", @@ -36,6 +37,7 @@ def _manifest() -> dict: "expected_full_items": 500, }, "pipeline": { + "stages": ["retain", "recall", "answer", "judge"], "planner": "ledger", "context_format": "structured_source", "thinking_budget": 500, @@ -49,8 +51,15 @@ def _manifest() -> dict: "schema": "public", "vector_extension": "pgvector", }, + "ingestion_provenance": { + "mode": "current_run", + "status": "current_run", + "retain_execution": "executed", + "content_identity": "verified_after_ingestion", + }, "concurrency": {"items": 1, "questions": 1, "judge": 1}, "runtime": { + "identity_scope": "executed_stages", "git_commit": "abc123", "git_dirty": False, "source_tree_fingerprint": None, @@ -79,7 +88,13 @@ def _model_config() -> dict: } -def _built_manifest(dataset_path: Path) -> dict: +def _built_manifest( + dataset_path: Path, + *, + skip_ingestion: bool = False, + ingest_only: bool = False, + force_reingest: bool = False, +) -> dict: return benchmark.build_run_manifest( dataset_path=dataset_path, context_format="structured_source", @@ -99,9 +114,9 @@ def _built_manifest(dataset_path: Path) -> dict: query_expansion_enabled=False, query_rewriting_strategy="noop", session_expansion_weight=0.3, - skip_ingestion=False, - ingest_only=False, - force_reingest=False, + skip_ingestion=skip_ingestion, + ingest_only=ingest_only, + force_reingest=force_reingest, ) @@ -128,6 +143,110 @@ def test_run_manifest_never_serializes_an_absolute_dataset_path(tmp_path: Path, assert not Path(external_manifest["dataset"]["path"]).is_absolute() +def test_run_manifest_distinguishes_fresh_and_reused_retain_provenance(tmp_path: Path, monkeypatch): + dataset_path = tmp_path / "dataset.json" + dataset_path.write_text("dataset", encoding="utf-8") + monkeypatch.setattr(benchmark, "_git_value", lambda *args: "abc123" if args == ("rev-parse", "HEAD") else "") + monkeypatch.setattr(benchmark, "_source_tree_fingerprint", lambda: None) + + fresh = _built_manifest(dataset_path) + reused = _built_manifest(dataset_path, skip_ingestion=True) + + assert fresh["artifact_schema_version"] == 2 + assert fresh["pipeline"]["stages"] == ["retain", "recall", "answer", "judge"] + assert fresh["ingestion_provenance"]["mode"] == "current_run" + assert reused["pipeline"]["stages"] == ["recall", "answer", "judge"] + assert reused["ingestion_provenance"] == { + "mode": "reused_bank", + "status": "unverifiable", + "retain_execution": "not_executed", + "content_identity": "verified_at_runtime", + "unverifiable_fields": ["retain_pipeline", "retain_model", "retain_code"], + } + assert reused["runtime"]["identity_scope"] == "executed_stages" + + +def test_ingest_only_manifest_marks_possible_reuse_as_mixed(tmp_path: Path, monkeypatch): + dataset_path = tmp_path / "dataset.json" + dataset_path.write_text("dataset", encoding="utf-8") + monkeypatch.setattr(benchmark, "_git_value", lambda *args: "") + monkeypatch.setattr(benchmark, "_source_tree_fingerprint", lambda: None) + + manifest = _built_manifest(dataset_path, ingest_only=True) + + assert manifest["pipeline"]["stages"] == ["retain"] + assert manifest["ingestion_provenance"] == { + "mode": "mixed_or_reused_bank", + "status": "unverifiable", + "retain_execution": "partial_or_skipped", + "content_identity": "verified_at_runtime", + "unverifiable_fields": [ + "reused_retain_pipeline", + "reused_retain_model", + "reused_retain_code", + ], + } + + +def test_force_reingest_ingest_only_manifest_is_current_run(tmp_path: Path, monkeypatch): + dataset_path = tmp_path / "dataset.json" + dataset_path.write_text("dataset", encoding="utf-8") + monkeypatch.setattr(benchmark, "_git_value", lambda *args: "") + monkeypatch.setattr(benchmark, "_source_tree_fingerprint", lambda: None) + + manifest = _built_manifest(dataset_path, ingest_only=True, force_reingest=True) + + assert manifest["pipeline"]["stages"] == ["retain"] + assert manifest["ingestion_provenance"]["mode"] == "current_run" + + +@pytest.mark.asyncio +async def test_only_ingested_filter_uses_exact_shared_preflight(): + items = [{"question_id": "a"}, {"question_id": "b"}] + dataset = SimpleNamespace(get_item_id=lambda item: item["question_id"]) + runner = SimpleNamespace( + _preflight_reusable_items=AsyncMock(return_value={"b"}), + ) + + filtered = await benchmark._filter_items_with_reusable_banks( + items, + dataset=dataset, + runner=runner, + ) + + assert filtered == [{"question_id": "b"}] + runner._preflight_reusable_items.assert_awaited_once_with( + items, + "longmemeval", + clear_agent_per_item=True, + require_all=False, + ) + + +def test_markdown_report_handles_unverifiable_reused_retain_identity(tmp_path: Path): + model_config = _model_config() + model_config["retain"] = { + "execution": "not_executed", + "bank_creator_identity": "unverifiable", + } + output_path = tmp_path / "results.json" + + benchmark.generate_markdown_table( + { + "model_config": model_config, + "item_results": [], + "overall_accuracy": 0.0, + "total_correct": 0, + "total_questions": 0, + "total_invalid": 0, + }, + output_path, + ) + + markdown = output_path.with_suffix(".md").read_text(encoding="utf-8") + assert "- **Retain**: not executed; reused-bank creator identity is unverifiable" in markdown + + def test_resume_compatibility_ignores_concurrency_but_rejects_model_changes(tmp_path: Path): output_path = tmp_path / "results.json" manifest = _manifest()