Skip to content
Open
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
7 changes: 7 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -122,6 +122,13 @@ HMS_API_RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICY=fixed_fallback
HMS_API_RETAIN_SEMANTIC_CHUNKING_MAX_COMPLETION_TOKENS=1024
HMS_API_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES=1

# Optional per-fact emotion and sentiment recognition. This enriches Retain's
# existing structured extraction call and stores affect metadata; Recall does
# not consume the field. Invalid affect output degrades to NULL without
# rejecting the memory write.
HMS_API_RETAIN_AFFECT_ENABLED=false
HMS_API_RETAIN_AFFECT_VERSION=affect-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
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
"""Add Retain affect metadata to memory units.

Revision ID: s7t8u9v0w1x2
Revises: r6s7t8u9v0w1
Create Date: 2026-09-02
"""

from collections.abc import Sequence

from alembic import context, op

from hms_api.alembic._dialect import run_for_dialect

revision: str = "s7t8u9v0w1x2"
down_revision: str | Sequence[str] | None = "r6s7t8u9v0w1"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None


def _get_schema_prefix() -> str:
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""


def _pg_upgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"""
ALTER TABLE {schema}memory_units
ADD COLUMN IF NOT EXISTS affect JSONB
""")
op.execute(f"""
DO $$
BEGIN
ALTER TABLE {schema}memory_units
ADD CONSTRAINT ck_memory_units_affect_object
CHECK (affect IS NULL OR jsonb_typeof(affect) = 'object');
EXCEPTION
WHEN duplicate_object THEN NULL;
END
$$
""")


def _pg_downgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS affect")


def _oracle_upgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"""
BEGIN
EXECUTE IMMEDIATE 'ALTER TABLE {schema}memory_units ADD (
affect CLOB
CONSTRAINT ck_mu_affect_json CHECK (affect IS NULL OR affect IS JSON)
)';
EXCEPTION
WHEN OTHERS THEN
IF SQLCODE != -1430 THEN
RAISE;
END IF;
END;
""")


def _oracle_downgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"""
BEGIN
EXECUTE IMMEDIATE 'ALTER TABLE {schema}memory_units DROP COLUMN affect';
EXCEPTION
WHEN OTHERS THEN
IF SQLCODE != -904 THEN
RAISE;
END IF;
END;
""")


def upgrade() -> None:
run_for_dialect(pg=_pg_upgrade, oracle=_oracle_upgrade)


def downgrade() -> None:
run_for_dialect(pg=_pg_downgrade, oracle=_oracle_downgrade)
14 changes: 14 additions & 0 deletions core/dataplane/hms_api/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -363,6 +363,8 @@ def normalize_config_dict(config: dict[str, Any]) -> dict[str, Any]:
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"
ENV_RETAIN_AFFECT_ENABLED = "HMS_API_RETAIN_AFFECT_ENABLED"
ENV_RETAIN_AFFECT_VERSION = "HMS_API_RETAIN_AFFECT_VERSION"

# File storage configuration
ENV_FILE_STORAGE_TYPE = "HMS_API_FILE_STORAGE_TYPE"
Expand Down Expand Up @@ -695,6 +697,8 @@ def normalize_config_dict(config: dict[str, Any]) -> dict[str, Any]:
RETAIN_SEMANTIC_CHUNKING_FAILURE_POLICIES = ("fixed_fallback", "raise")
DEFAULT_RETAIN_SEMANTIC_CHUNKING_MAX_COMPLETION_TOKENS = 1024
DEFAULT_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES = 1
DEFAULT_RETAIN_AFFECT_ENABLED = False
DEFAULT_RETAIN_AFFECT_VERSION = "affect-v1"

# File storage defaults
DEFAULT_FILE_STORAGE_TYPE = "native" # PostgreSQL BYTEA storage
Expand Down Expand Up @@ -1358,6 +1362,8 @@ class HMSConfig:
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
retain_affect_enabled: bool = DEFAULT_RETAIN_AFFECT_ENABLED
retain_affect_version: str = DEFAULT_RETAIN_AFFECT_VERSION
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 @@ -1461,6 +1467,7 @@ class HMSConfig:
"retain_default_strategy",
"retain_strategies",
"retain_chunk_batch_size",
"retain_affect_enabled",
# Entity labels (controlled vocabulary for entity classification)
"entity_labels",
"entities_allow_free_form",
Expand Down Expand Up @@ -2204,6 +2211,13 @@ def from_env(cls) -> "HMSConfig":
str(DEFAULT_RETAIN_SEMANTIC_CHUNKING_MAX_RETRIES),
)
),
retain_affect_enabled=os.getenv(
ENV_RETAIN_AFFECT_ENABLED,
str(DEFAULT_RETAIN_AFFECT_ENABLED),
).lower()
== "true",
retain_affect_version=os.getenv(ENV_RETAIN_AFFECT_VERSION, DEFAULT_RETAIN_AFFECT_VERSION).strip()
or DEFAULT_RETAIN_AFFECT_VERSION,
# 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/db/ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,7 @@ async def insert_facts_batch(
observation_scopes_list: list,
text_signals_list: list,
projection_jsons: list[str],
affect_jsons: list[str | None] | None = None,
text_search_extension: str = "native",
) -> list[str]:
"""Batch-insert facts, returning IDs.
Expand Down
8 changes: 6 additions & 2 deletions core/dataplane/hms_api/engine/db/ops_oracle.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,12 +66,15 @@ async def insert_facts_batch(
observation_scopes_list: list,
text_signals_list: list,
projection_jsons: list[str],
affect_jsons: list[str | None] | None = None,
text_search_extension: str = "native",
) -> list[str]:
table = self._get_mu_table()
# Generate UUIDs client-side so we can use executemany (single network
# round-trip) instead of N individual INSERT+RETURNING calls.
unit_ids = [str(uuid_mod.uuid4()) for _ in range(len(fact_texts))]
if affect_jsons is None:
affect_jsons = [None] * len(fact_texts)
rows_data = []
for i in range(len(fact_texts)):
tags_value = json.loads(tags_list[i]) if tags_list[i] else []
Expand All @@ -94,14 +97,15 @@ async def insert_facts_batch(
observation_scopes_list[i],
text_signals_list[i],
projection_jsons[i] or "{}",
affect_jsons[i],
)
)
await conn.executemany(
f"""
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)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17)
tags, observation_scopes, text_signals, projection, affect)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18)
""",
rows_data,
)
Expand Down
20 changes: 13 additions & 7 deletions core/dataplane/hms_api/engine/db/ops_postgresql.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,7 @@ async def insert_facts_batch(
observation_scopes_list: list,
text_signals_list: list,
projection_jsons: list[str],
affect_jsons: list[str | None] | None = None,
text_search_extension: str = "native",
) -> list[str]:
from ...config import get_config
Expand All @@ -81,21 +82,23 @@ async def insert_facts_batch(
# 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 affect_jsons is None:
affect_jsons = [None] * len(fact_texts)

if config.text_search_extension == "vchord":
query = f"""
WITH input_data AS (
SELECT * FROM unnest(
$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[]
$15::jsonb[], $16::text[], $17::jsonb[], $18::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)
observation_scopes_json, text_signals, projection, affect)
)
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)
observation_scopes, text_signals, projection, affect, search_vector)
SELECT
id, $1,
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
Expand All @@ -107,6 +110,7 @@ async def insert_facts_batch(
observation_scopes_json,
text_signals,
COALESCE(projection, '{{}}'::jsonb),
affect,
tokenize(
COALESCE(text, '') || ' ' || COALESCE(context, '') || ' ' || COALESCE(text_signals, ''),
'llmlingua2'
Expand All @@ -119,14 +123,14 @@ async def insert_facts_batch(
SELECT * FROM unnest(
$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[]
$15::jsonb[], $16::text[], $17::jsonb[], $18::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)
observation_scopes_json, text_signals, projection, affect)
)
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)
observation_scopes, text_signals, projection, affect)
SELECT
id, $1,
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
Expand All @@ -137,7 +141,8 @@ async def insert_facts_batch(
),
observation_scopes_json,
text_signals,
COALESCE(projection, '{{}}'::jsonb)
COALESCE(projection, '{{}}'::jsonb),
affect
FROM input_data
"""

Expand All @@ -160,6 +165,7 @@ async def insert_facts_batch(
observation_scopes_list,
text_signals_list,
projection_jsons,
affect_jsons,
)
return unit_ids

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
from typing import Any

from ...response_models import TokenUsage
from ...retain.affect import AffectSignals
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
Expand Down Expand Up @@ -478,6 +479,11 @@ def _convert_facts(
getattr(fact, "entities", None),
field_name=f"fact[{fact_index}].entities",
),
affect=(
getattr(fact, "affect", None)
if isinstance(getattr(fact, "affect", None), AffectSignals)
else None
),
causal_relations=(),
)
except ExtractionContractError:
Expand Down
4 changes: 4 additions & 0 deletions core/dataplane/hms_api/engine/ingestion/extraction/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from dataclasses import dataclass
from datetime import datetime

from ...retain.affect import AffectSignals
from ..domain import FrozenJson, ObservationScopes

FACT_KEY_VERSION = "retain-fact-v1"
Expand Down Expand Up @@ -98,6 +99,7 @@ class FactCandidate:
# objects. Passthrough extraction deliberately leaves these empty, which
# preserves the existing projection manifest semantics.
entity_mentions: tuple[str, ...]
affect: AffectSignals | None
causal_relations: tuple[CausalFactRelation, ...]

def __post_init__(self) -> None:
Expand Down Expand Up @@ -137,6 +139,8 @@ def __post_init__(self) -> None:
not isinstance(entity, str) for entity in self.entity_mentions
):
raise TypeError("entity_mentions must be a tuple of strings")
if self.affect is not None and not isinstance(self.affect, AffectSignals):
raise TypeError("affect must be AffectSignals or None")
if not isinstance(self.causal_relations, tuple) or any(
not isinstance(relation, CausalFactRelation) for relation in self.causal_relations
):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,7 @@ def extract_passthrough(
tags=item.tags,
observation_scopes=item.observation_scopes,
entity_mentions=(),
affect=None,
causal_relations=(),
)
)
Expand Down
2 changes: 2 additions & 0 deletions core/dataplane/hms_api/engine/ingestion/projection/records.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@ def from_candidate(
tags=candidate.tags,
observation_scopes=candidate.observation_scopes,
entity_mentions=candidate.entity_mentions,
affect=candidate.affect,
causal_relations=candidate.causal_relations,
embedding=embedding,
projection=projection,
Expand Down Expand Up @@ -169,6 +170,7 @@ def to_processed_fact(
where=record.where,
entities=[EntityRef(name=name) for name in record.entity_mentions],
causal_relations=causal_relations,
affect=record.affect,
chunk_id=chunk_id,
document_id=document_id,
content_index=content_index,
Expand Down
Loading
Loading