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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,23 @@ 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

# Retain uses semantic boundary planning by default for long JSON conversations.
# The planner asks the Retain model for topic boundaries, then materializes
# chunks from the original exchanges without rewriting source text. Short
# content and trusted pre-chunked input bypass the planner. Non-conversation
# content uses deterministic structural chunking, and planning failures follow
# the failure policy below.
HMS_API_RETAIN_CHUNK_SIZE=3000
HMS_API_RETAIN_SEMANTIC_CHUNKING_ENABLED=true
# Set this to false when using provider Batch extraction, which does not yet
# bind Batch checkpoints to semantic plan digests.
# "fixed_fallback" preserves ingestion when boundary planning fails; "raise"
# fails the Retain request and is useful for controlled experiments.
HMS_API_RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICY=fixed_fallback
# These limits apply only to the boundary-planning call, not fact extraction.
HMS_API_RETAIN_SEMANTIC_CHUNKING_MAX_COMPLETION_TOKENS=1024
HMS_API_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES=1

# 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
Expand Down
15 changes: 12 additions & 3 deletions .github/workflows/retain-offline.yml
Original file line number Diff line number Diff line change
Expand Up @@ -70,7 +70,10 @@ jobs:
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
core/dataplane/tests/test_op_cancellation.py \
core/dataplane/tests/test_prechunked_extraction_boundaries.py \
core/dataplane/tests/test_semantic_segmentation.py \
core/dataplane/tests/test_semantic_segmentation_service.py
uv run --no-sync --package hms-api-slim ruff format --check \
--config core/dataplane/pyproject.toml \
core/dataplane/hms_api/config.py \
Expand All @@ -95,7 +98,10 @@ jobs:
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
core/dataplane/tests/test_op_cancellation.py \
core/dataplane/tests/test_prechunked_extraction_boundaries.py \
core/dataplane/tests/test_semantic_segmentation.py \
core/dataplane/tests/test_semantic_segmentation_service.py

- name: Check LongMemEval source style
run: |
Expand Down Expand Up @@ -167,7 +173,10 @@ jobs:
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
core/dataplane/tests/test_multimodal_security.py \
core/dataplane/tests/test_prechunked_extraction_boundaries.py \
core/dataplane/tests/test_semantic_segmentation.py \
core/dataplane/tests/test_semantic_segmentation_service.py

- name: Run LongMemEval offline tests
run: |
Expand Down
44 changes: 44 additions & 0 deletions core/dataplane/hms_api/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -359,6 +359,10 @@ def normalize_config_dict(config: dict[str, Any]) -> dict[str, Any]:
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"
ENV_RETAIN_SEMANTIC_CHUNKING_ENABLED = "HMS_API_RETAIN_SEMANTIC_CHUNKING_ENABLED"
ENV_RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICY = "HMS_API_RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICY"
ENV_RETAIN_SEMANTIC_CHUNKING_MAX_COMPLETION_TOKENS = "HMS_API_RETAIN_SEMANTIC_CHUNKING_MAX_COMPLETION_TOKENS"
ENV_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES = "HMS_API_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES"

# File storage configuration
ENV_FILE_STORAGE_TYPE = "HMS_API_FILE_STORAGE_TYPE"
Expand Down Expand Up @@ -686,6 +690,11 @@ def normalize_config_dict(config: dict[str, Any]) -> dict[str, Any]:
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")
DEFAULT_RETAIN_SEMANTIC_CHUNKING_ENABLED = True
DEFAULT_RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICY = "fixed_fallback"
RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICIES = ("fixed_fallback", "raise")
DEFAULT_RETAIN_SEMANTIC_CHUNKING_MAX_COMPLETION_TOKENS = 1024
DEFAULT_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES = 1

# File storage defaults
DEFAULT_FILE_STORAGE_TYPE = "native" # PostgreSQL BYTEA storage
Expand Down Expand Up @@ -1345,6 +1354,10 @@ class HMSConfig:
# 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
retain_semantic_chunking_enabled: bool = DEFAULT_RETAIN_SEMANTIC_CHUNKING_ENABLED
retain_semantic_chunking_failure_policy: str = DEFAULT_RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICY
retain_semantic_chunking_max_completion_tokens: int = DEFAULT_RETAIN_SEMANTIC_CHUNKING_MAX_COMPLETION_TOKENS
retain_semantic_chunking_max_retries: int = DEFAULT_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES
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
Expand Down Expand Up @@ -1762,6 +1775,16 @@ def validate(self) -> None:
f"\n (current model: {self.retain_llm_model or self.llm_model}, "
f"provider: {self.retain_llm_provider or self.llm_provider})"
)
if self.retain_semantic_chunking_failure_policy not in RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICIES:
choices = ", ".join(RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICIES)
raise ValueError(f"{ENV_RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICY} must be one of: {choices}")
if (
isinstance(self.retain_semantic_chunking_max_completion_tokens, bool)
or self.retain_semantic_chunking_max_completion_tokens <= 0
):
raise ValueError(f"{ENV_RETAIN_SEMANTIC_CHUNKING_MAX_COMPLETION_TOKENS} must be a positive integer")
if isinstance(self.retain_semantic_chunking_max_retries, bool) or self.retain_semantic_chunking_max_retries < 0:
raise ValueError(f"{ENV_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES} must be a non-negative integer")

# Warn if local ML dependencies are missing when configured.
# Don't hard-fail here — the actual ImportError fires at model init time
Expand Down Expand Up @@ -2160,6 +2183,27 @@ def from_env(cls) -> "HMSConfig":
DEFAULT_RETAIN_EMBEDDING_FAILURE_POLICY,
)
),
retain_semantic_chunking_enabled=os.getenv(
ENV_RETAIN_SEMANTIC_CHUNKING_ENABLED,
str(DEFAULT_RETAIN_SEMANTIC_CHUNKING_ENABLED),
).lower()
== "true",
retain_semantic_chunking_failure_policy=os.getenv(
ENV_RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICY,
DEFAULT_RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICY,
),
retain_semantic_chunking_max_completion_tokens=int(
os.getenv(
ENV_RETAIN_SEMANTIC_CHUNKING_MAX_COMPLETION_TOKENS,
str(DEFAULT_RETAIN_SEMANTIC_CHUNKING_MAX_COMPLETION_TOKENS),
)
),
retain_semantic_chunking_max_retries=int(
os.getenv(
ENV_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES,
str(DEFAULT_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES),
)
),
# 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,
Expand Down
1 change: 1 addition & 0 deletions core/dataplane/hms_api/engine/ingestion/contracts.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ class RetainInvocation:
outbox_callback: OutboxCallback | None = None
strategy: str | None = None
sanitize_log_identifiers: bool = False
trusted_prechunked_input: bool = False


@dataclass(frozen=True, slots=True)
Expand Down
31 changes: 31 additions & 0 deletions core/dataplane/hms_api/engine/ingestion/extraction/extractor.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,6 +234,36 @@ def _sync_fallback_config(self) -> Any:
setattr(fallback, "retain_batch_enabled", False)
return fallback

@staticmethod
def _boundary_preserving_config(request: ExtractionRequest, config: Any) -> Any:
"""Prevent the extraction primitive from re-splitting planned chunks.

A pre-chunked layout represents every planned chunk as one temporary
content item. The extraction primitive still applies
``retain_chunk_size`` to each item, so a semantic segment containing a
complete oversized exchange could otherwise be divided between its
user and assistant turns. Raising the limit on a request-local shallow
copy preserves the planner boundary without mutating shared bank
configuration. Output-overflow recovery inside the primitive remains
available because it operates after this initial split.
"""

if not request.preserve_chunk_boundaries or not request.items:
return config

required_chunk_size = max(len(item.content) for item in request.items)
configured_chunk_size = getattr(config, "retain_chunk_size", None)
if (
isinstance(configured_chunk_size, int)
and not isinstance(configured_chunk_size, bool)
and configured_chunk_size >= required_chunk_size
):
return config

boundary_config = copy(config)
setattr(boundary_config, "retain_chunk_size", required_chunk_size)
return boundary_config

async def _batch_primitive_if_supported(
self,
mode: ExtractionMode,
Expand Down Expand Up @@ -281,6 +311,7 @@ async def extract(self, request: ExtractionRequest) -> ExtractionResult:
primitive_config = self._config
if getattr(self._config, "retain_batch_enabled", False):
primitive, primitive_config = await self._batch_primitive_if_supported(request.policy.mode)
primitive_config = self._boundary_preserving_config(request, primitive_config)

storage_contents = [_storage_content(item) for item in request.items]
primitive_result = await primitive(
Expand Down
7 changes: 6 additions & 1 deletion core/dataplane/hms_api/engine/ingestion/extraction/layout.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,7 +114,12 @@ def __post_init__(self) -> None:
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)
return ExtractionRequest(
items=self.temporary_items,
chunks=self.temporary_chunks,
policy=policy,
preserve_chunk_boundaries=True,
)

def remap_result(self, result: ExtractionResult) -> ExtractionResult:
"""Restore original source indices and every derived fact identity."""
Expand Down
3 changes: 3 additions & 0 deletions core/dataplane/hms_api/engine/ingestion/extraction/ports.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@ class ExtractionRequest:
items: tuple[ContentItem, ...]
chunks: tuple[ChunkPlan, ...]
policy: ExtractionPolicy
preserve_chunk_boundaries: bool = False

def __post_init__(self) -> None:
if not isinstance(self.items, tuple) or any(not isinstance(item, ContentItem) for item in self.items):
Expand All @@ -44,6 +45,8 @@ def __post_init__(self) -> None:
raise TypeError("chunks must be a tuple of ChunkPlan values")
if not isinstance(self.policy, ExtractionPolicy):
raise TypeError("policy must be an ExtractionPolicy")
if not isinstance(self.preserve_chunk_boundaries, bool):
raise TypeError("preserve_chunk_boundaries must be a bool")


@dataclass(frozen=True, slots=True)
Expand Down
4 changes: 2 additions & 2 deletions core/dataplane/hms_api/engine/ingestion/persistence/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,8 +100,8 @@ def unit_ids_for_document(self, document_id: str) -> tuple[str, ...] | None:

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.
mapping was absent or did not contain the document, so callers must
fail closed unless they have another operation-local proof.
"""

if not isinstance(document_id, str) or not document_id:
Expand Down
59 changes: 59 additions & 0 deletions core/dataplane/hms_api/engine/ingestion/segmentation/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
"""Semantic conversation segmentation for Retain chunk planning.

This package is intentionally independent from the Retain application service.
It exposes an asynchronous planner that can be integrated between normalized
document planning and ``ChunkPlan`` construction without changing extraction or
persistence contracts.
"""

from .adapters import build_chunk_plans_from_segmentation
from .models import (
SEMANTIC_POLICY_VERSION,
SEMANTIC_PROMPT_VERSION,
BoundaryResponse,
ConversationExchange,
EffectiveSegmentationStrategy,
MaterializedSegment,
ParsedConversation,
SegmentationFailurePolicy,
SegmentationManifest,
SegmentationMode,
SegmentationResult,
SegmentManifestEntry,
SemanticSegmentationPolicy,
)
from .planner import (
SegmentationReuseError,
SemanticBoundaryValidationError,
SemanticSegmentationError,
SemanticSegmenter,
UnsplittableExchangeError,
materialize_semantic_boundaries,
parse_conversation,
validate_boundary_response,
)

__all__ = [
"BoundaryResponse",
"ConversationExchange",
"EffectiveSegmentationStrategy",
"MaterializedSegment",
"ParsedConversation",
"SEMANTIC_POLICY_VERSION",
"SEMANTIC_PROMPT_VERSION",
"SegmentManifestEntry",
"SegmentationFailurePolicy",
"SegmentationManifest",
"SegmentationMode",
"SegmentationResult",
"SegmentationReuseError",
"SemanticBoundaryValidationError",
"SemanticSegmentationError",
"SemanticSegmentationPolicy",
"SemanticSegmenter",
"UnsplittableExchangeError",
"build_chunk_plans_from_segmentation",
"materialize_semantic_boundaries",
"parse_conversation",
"validate_boundary_response",
]
61 changes: 61 additions & 0 deletions core/dataplane/hms_api/engine/ingestion/segmentation/adapters.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
"""Adapters from per-item segmentation results to Retain chunk plans."""

from __future__ import annotations

from collections.abc import Sequence

from ..chunking import compute_content_hash
from ..domain import ChunkPlan, ContentItem
from .models import SegmentationResult


def build_chunk_plans_from_segmentation(
document_id: str,
items: Sequence[ContentItem],
results: Sequence[SegmentationResult],
) -> tuple[ChunkPlan, ...]:
"""Build stable document-wide plans in original item order.

Callers may plan items concurrently, but ``results`` must be restored to
the same order as ``items`` before calling this function. ``local_index``
restarts for each item; ``global_index`` and ``chunk_key`` are assigned
synchronously across the complete document, matching the existing Retain
chunk identity contract.
"""

if not isinstance(document_id, str) or not document_id:
raise ValueError("document_id must be a non-empty string")
if isinstance(items, (str, bytes)) or not isinstance(items, Sequence):
raise TypeError("items must be a sequence of ContentItem values")
if isinstance(results, (str, bytes)) or not isinstance(results, Sequence):
raise TypeError("results must be a sequence of SegmentationResult values")
if len(items) != len(results):
raise ValueError("items and results must have the same length")

plans: list[ChunkPlan] = []
global_index = 0
for item, result in zip(items, results, strict=True):
if not isinstance(item, ContentItem):
raise TypeError("items must contain only ContentItem values")
if not isinstance(result, SegmentationResult):
raise TypeError("results must contain only SegmentationResult values")
for local_index, segment in enumerate(result.segments):
content_hash = compute_content_hash(segment.text)
if content_hash != segment.content_hash:
raise ValueError("segment content_hash does not match its text")
chunk_key = f"chunk:{len(document_id)}:{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=segment.text,
content_hash=content_hash,
)
)
global_index += 1
return tuple(plans)


__all__ = ["build_chunk_plans_from_segmentation"]
Loading