From 8fc394b7edbe8ad178e3018467d5a42ed5e7b328 Mon Sep 17 00:00:00 2001 From: xuanmingShang <1260379304@qq.com> Date: Wed, 2 Sep 2026 02:12:16 +0800 Subject: [PATCH] feat(retain): add affect recognition metadata --- .env.example | 7 + ...s7t8u9v0w1x2_add_affect_to_memory_units.py | 85 ++++++ core/dataplane/hms_api/config.py | 14 + core/dataplane/hms_api/engine/db/ops.py | 1 + .../dataplane/hms_api/engine/db/ops_oracle.py | 8 +- .../hms_api/engine/db/ops_postgresql.py | 20 +- .../engine/ingestion/extraction/extractor.py | 6 + .../engine/ingestion/extraction/models.py | 4 + .../ingestion/extraction/passthrough.py | 1 + .../engine/ingestion/projection/records.py | 2 + .../dataplane/hms_api/engine/retain/affect.py | 82 ++++++ .../hms_api/engine/retain/fact_extraction.py | 103 +++++-- .../hms_api/engine/retain/fact_storage.py | 3 + core/dataplane/hms_api/engine/retain/types.py | 6 + core/dataplane/tests/test_db_abstraction.py | 15 +- core/dataplane/tests/test_retain_affect.py | 271 ++++++++++++++++++ docker-compose.yml | 2 + 17 files changed, 591 insertions(+), 39 deletions(-) create mode 100644 core/dataplane/hms_api/alembic/versions/s7t8u9v0w1x2_add_affect_to_memory_units.py create mode 100644 core/dataplane/hms_api/engine/retain/affect.py create mode 100644 core/dataplane/tests/test_retain_affect.py diff --git a/.env.example b/.env.example index 92fe7b6..ec78458 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/core/dataplane/hms_api/alembic/versions/s7t8u9v0w1x2_add_affect_to_memory_units.py b/core/dataplane/hms_api/alembic/versions/s7t8u9v0w1x2_add_affect_to_memory_units.py new file mode 100644 index 0000000..4bfc3f9 --- /dev/null +++ b/core/dataplane/hms_api/alembic/versions/s7t8u9v0w1x2_add_affect_to_memory_units.py @@ -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) diff --git a/core/dataplane/hms_api/config.py b/core/dataplane/hms_api/config.py index 064db1c..60efd0b 100644 --- a/core/dataplane/hms_api/config.py +++ b/core/dataplane/hms_api/config.py @@ -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" @@ -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 @@ -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 @@ -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", @@ -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, diff --git a/core/dataplane/hms_api/engine/db/ops.py b/core/dataplane/hms_api/engine/db/ops.py index bc8904e..599571e 100644 --- a/core/dataplane/hms_api/engine/db/ops.py +++ b/core/dataplane/hms_api/engine/db/ops.py @@ -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. diff --git a/core/dataplane/hms_api/engine/db/ops_oracle.py b/core/dataplane/hms_api/engine/db/ops_oracle.py index 8b6f8bf..82c4b87 100644 --- a/core/dataplane/hms_api/engine/db/ops_oracle.py +++ b/core/dataplane/hms_api/engine/db/ops_oracle.py @@ -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 [] @@ -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, ) diff --git a/core/dataplane/hms_api/engine/db/ops_postgresql.py b/core/dataplane/hms_api/engine/db/ops_postgresql.py index 695ec6f..e8f3532 100644 --- a/core/dataplane/hms_api/engine/db/ops_postgresql.py +++ b/core/dataplane/hms_api/engine/db/ops_postgresql.py @@ -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 @@ -81,6 +82,8 @@ 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""" @@ -88,14 +91,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, search_vector) + observation_scopes, text_signals, projection, affect, search_vector) SELECT id, $1, text, embedding, event_date, occurred_start, occurred_end, mentioned_at, @@ -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' @@ -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, @@ -137,7 +141,8 @@ async def insert_facts_batch( ), observation_scopes_json, text_signals, - COALESCE(projection, '{{}}'::jsonb) + COALESCE(projection, '{{}}'::jsonb), + affect FROM input_data """ @@ -160,6 +165,7 @@ async def insert_facts_batch( observation_scopes_list, text_signals_list, projection_jsons, + affect_jsons, ) return unit_ids diff --git a/core/dataplane/hms_api/engine/ingestion/extraction/extractor.py b/core/dataplane/hms_api/engine/ingestion/extraction/extractor.py index 4e4d6f4..6765b23 100644 --- a/core/dataplane/hms_api/engine/ingestion/extraction/extractor.py +++ b/core/dataplane/hms_api/engine/ingestion/extraction/extractor.py @@ -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 @@ -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: diff --git a/core/dataplane/hms_api/engine/ingestion/extraction/models.py b/core/dataplane/hms_api/engine/ingestion/extraction/models.py index 90297af..6d7d68f 100644 --- a/core/dataplane/hms_api/engine/ingestion/extraction/models.py +++ b/core/dataplane/hms_api/engine/ingestion/extraction/models.py @@ -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" @@ -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: @@ -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 ): diff --git a/core/dataplane/hms_api/engine/ingestion/extraction/passthrough.py b/core/dataplane/hms_api/engine/ingestion/extraction/passthrough.py index 09ca19f..80d40dd 100644 --- a/core/dataplane/hms_api/engine/ingestion/extraction/passthrough.py +++ b/core/dataplane/hms_api/engine/ingestion/extraction/passthrough.py @@ -106,6 +106,7 @@ def extract_passthrough( tags=item.tags, observation_scopes=item.observation_scopes, entity_mentions=(), + affect=None, causal_relations=(), ) ) diff --git a/core/dataplane/hms_api/engine/ingestion/projection/records.py b/core/dataplane/hms_api/engine/ingestion/projection/records.py index 9d67917..1f89520 100644 --- a/core/dataplane/hms_api/engine/ingestion/projection/records.py +++ b/core/dataplane/hms_api/engine/ingestion/projection/records.py @@ -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, @@ -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, diff --git a/core/dataplane/hms_api/engine/retain/affect.py b/core/dataplane/hms_api/engine/retain/affect.py new file mode 100644 index 0000000..facf8c1 --- /dev/null +++ b/core/dataplane/hms_api/engine/retain/affect.py @@ -0,0 +1,82 @@ +"""Validated affect signals produced during Retain fact extraction.""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import asdict, dataclass +from typing import Any, Literal + +from pydantic import BaseModel, ConfigDict, Field + +Sentiment = Literal["positive", "negative", "neutral"] +Emotion = Literal["joy", "sadness", "anger", "fear", "surprise", "disgust", "neutral"] + + +class AffectAnnotation(BaseModel): + """Strict provider-facing schema embedded in the fact extraction response.""" + + model_config = ConfigDict(extra="forbid") + + sentiment: Sentiment = Field(description="Overall affective polarity expressed in the fact") + emotion: Emotion = Field(description="Primary expressed emotion, or neutral when none is evidenced") + intensity: float = Field( + ge=0.0, + le=1.0, + description="Strength of the expressed emotion from 0.0 (none) to 1.0 (very strong)", + ) + + +@dataclass(frozen=True, slots=True) +class AffectSignals: + """Provider-neutral, versioned affect attached to one memory fact.""" + + sentiment: Sentiment + emotion: Emotion + intensity: float + version: str + + def __post_init__(self) -> None: + if self.sentiment not in {"positive", "negative", "neutral"}: + raise ValueError("sentiment must be positive, negative, or neutral") + if self.emotion not in {"joy", "sadness", "anger", "fear", "surprise", "disgust", "neutral"}: + raise ValueError("emotion is not supported by the affect-v1 taxonomy") + if isinstance(self.intensity, bool) or not isinstance(self.intensity, (int, float)): + raise TypeError("intensity must be a number") + if not 0.0 <= float(self.intensity) <= 1.0: + raise ValueError("intensity must be between 0.0 and 1.0") + if not isinstance(self.version, str) or not self.version.strip(): + raise ValueError("version must be a non-empty string") + + object.__setattr__(self, "intensity", float(self.intensity)) + object.__setattr__(self, "version", self.version.strip()) + + def to_json(self) -> dict[str, str | float]: + """Return the JSON object persisted with the memory unit.""" + + return asdict(self) + + +def parse_affect(value: Any, *, version: str) -> AffectSignals | None: + """Parse untrusted model output without making affect fatal to Retain.""" + + if isinstance(value, AffectSignals): + return value + if isinstance(value, AffectAnnotation): + annotation = value + elif isinstance(value, Mapping): + try: + annotation = AffectAnnotation.model_validate(dict(value)) + except (TypeError, ValueError): + return None + else: + return None + + try: + return AffectSignals( + sentiment=annotation.sentiment, + emotion=annotation.emotion, + intensity=annotation.intensity, + version=version, + ) + except (TypeError, ValueError): + return None diff --git a/core/dataplane/hms_api/engine/retain/fact_extraction.py b/core/dataplane/hms_api/engine/retain/fact_extraction.py index 2a42eb5..4803a94 100644 --- a/core/dataplane/hms_api/engine/retain/fact_extraction.py +++ b/core/dataplane/hms_api/engine/retain/fact_extraction.py @@ -17,6 +17,7 @@ from ...config import get_config from ..llm_wrapper import LLMConfig, OutputTooLongError, sanitize_llm_output from ..response_models import TokenUsage +from .affect import AffectAnnotation, AffectSignals, parse_affect from .entity_labels import ( EntityLabelsConfig, MapField, @@ -194,6 +195,7 @@ class Fact(BaseModel): # Optional structured data entities: list[Entity] | None = None causal_relations: list["CausalRelation"] | None = None + affect: AffectSignals | None = None class CausalRelation(BaseModel): @@ -885,6 +887,23 @@ def _chunk_conversation(turns: list[dict], max_chars: int) -> list[str]: - Fact 2: Moved apartment, causal_relations: [{target_index: 1, relation_type: "caused_by"}]""" +AFFECT_RECOGNITION_SECTION = """ + +══════════════════════════════════════════════════════════════════════════ +AFFECT RECOGNITION +══════════════════════════════════════════════════════════════════════════ + +For every fact, classify the affect explicitly expressed by the person the fact is about: +- sentiment: positive, negative, or neutral +- emotion: joy, sadness, anger, fear, surprise, disgust, or neutral +- intensity: 0.0 (none) to 1.0 (very strong) + +Use evidence in the source text, including multilingual and conversational wording. Do not infer a person's +emotion merely because an event or topic is usually positive or negative. Do not attribute the assistant's +empathetic tone to the user. When no affect is expressed, return neutral sentiment, neutral emotion, and 0.0 +intensity. Choose the primary emotion when several are expressed.""" + + def _append_map_fields_prompt(fields: dict[str, "MapField"], lines: list[str], indent: int = 4) -> None: """Recursively append map field descriptions to the prompt lines.""" pad = " " * indent @@ -1058,41 +1077,55 @@ def _build_extraction_prompt_and_schema(config) -> tuple[str, type]: prompt = prompt + labels_section response_schema = base_response_class + dynamic_fields: dict = {} + required_dynamic_fields: list[str] = [] + + if getattr(config, "retain_affect_enabled", False) is True: + prompt = prompt + AFFECT_RECOGNITION_SECTION + dynamic_fields["affect"] = ( + AffectAnnotation, + Field(description="Affect explicitly expressed in this fact"), + ) + required_dynamic_fields.append("affect") if labels_cfg and labels_cfg.attributes: LabelsModel = build_labels_model(labels_cfg) if LabelsModel is not None: - dynamic_fields: dict = { - "labels": ( - LabelsModel, - Field( - description="Classification labels for this fact. Fill each applicable field; leave others null/empty." - ), - ) - } + dynamic_fields["labels"] = ( + LabelsModel, + Field( + description="Classification labels for this fact. Fill each applicable field; leave others null/empty." + ), + ) + required_dynamic_fields.append("labels") if not free_form_entities: dynamic_fields["entities"] = ( list[Entity] | None, Field(default=None, description="Leave empty — labels-only mode"), ) - # Inherit parent's required fields and add 'labels' so it appears in the JSON schema - # required array (the base class json_schema_extra overrides required entirely) - base_extra = base_fact_class.model_config.get("json_schema_extra") - base_required = cast(dict, base_extra).get("required", []) if isinstance(base_extra, dict) else [] - DynamicFact = create_model( - "LabelsFact", - __base__=base_fact_class, - __config__=ConfigDict( - json_schema_mode="validation", - json_schema_extra={"required": [*base_required, "labels"]}, - ), - **dynamic_fields, - ) - DynamicResponse = create_model( - "LabelsResponse", - facts=(list[DynamicFact], ...), # type: ignore[valid-type] # ty: ignore[invalid-type-form] - ) - response_schema = DynamicResponse + + if dynamic_fields: + has_labels = bool(labels_cfg and labels_cfg.attributes) + dynamic_fact_name = "LabelsFact" if has_labels else "AffectFact" + dynamic_response_name = "LabelsResponse" if has_labels else "AffectResponse" + # The base class schema explicitly owns its required array, so carry it + # forward and append every required enrichment field. + base_extra = base_fact_class.model_config.get("json_schema_extra") + base_required = cast(dict, base_extra).get("required", []) if isinstance(base_extra, dict) else [] + DynamicFact = create_model( + dynamic_fact_name, + __base__=base_fact_class, + __config__=ConfigDict( + json_schema_mode="validation", + json_schema_extra={"required": [*base_required, *required_dynamic_fields]}, + ), + **dynamic_fields, + ) + DynamicResponse = create_model( + dynamic_response_name, + facts=(list[DynamicFact], ...), # type: ignore[valid-type] # ty: ignore[invalid-type-form] + ) + response_schema = DynamicResponse return prompt, response_schema @@ -1421,6 +1454,14 @@ def get_value(field_name): if validated_entities: fact_data["entities"] = validated_entities + if getattr(config, "retain_affect_enabled", False) is True: + affect = parse_affect( + get_value("affect"), + version=getattr(config, "retain_affect_version", "affect-v1"), + ) + if affect is not None: + fact_data["affect"] = affect + # Add per-fact causal relations (only if enabled in config) if extract_causal_links: validated_relations = _remap_causal_relations( @@ -2102,6 +2143,14 @@ def get_value(field_name): if validated_entities: fact_data["entities"] = validated_entities + if getattr(config, "retain_affect_enabled", False) is True: + affect = parse_affect( + get_value("affect"), + version=getattr(config, "retain_affect_version", "affect-v1"), + ) + if affect is not None: + fact_data["affect"] = affect + # Causal relations if extract_causal_links: validated_relations = _remap_causal_relations( @@ -2173,6 +2222,7 @@ def get_value(field_name): fact_from_llm.causal_relations or [], chunk_fact_start_idx, ), + affect=fact_from_llm.affect, content_index=chunk_meta.content_index, chunk_index=chunk_meta.chunk_index, context=content.context, @@ -2372,6 +2422,7 @@ async def extract_facts_from_contents( fact_from_llm.causal_relations or [], chunk_fact_start_idx, ), + affect=fact_from_llm.affect, content_index=content_index, chunk_index=chunk_global_idx, context=content.context, diff --git a/core/dataplane/hms_api/engine/retain/fact_storage.py b/core/dataplane/hms_api/engine/retain/fact_storage.py index 5412e05..dda30c1 100644 --- a/core/dataplane/hms_api/engine/retain/fact_storage.py +++ b/core/dataplane/hms_api/engine/retain/fact_storage.py @@ -78,6 +78,7 @@ async def insert_facts_batch( observation_scopes_list = [] text_signals_list = [] projection_jsons = [] + affect_jsons = [] for fact in facts: fact_texts.append(_sanitize_text(fact.fact_text)) @@ -117,6 +118,7 @@ async def insert_facts_batch( pass text_signals_list.append(" ".join(signal_parts) if signal_parts else None) projection_jsons.append(json.dumps(fact.projection or {})) + affect_jsons.append(json.dumps(fact.affect.to_json()) if fact.affect is not None else None) # Batch insert all facts — delegates to DataAccessOps which handles # unnest (PG) vs row-by-row (Oracle) transparently. @@ -140,6 +142,7 @@ async def insert_facts_batch( observation_scopes_list, text_signals_list, projection_jsons, + affect_jsons, text_search_extension=config.text_search_extension, ) return unit_ids diff --git a/core/dataplane/hms_api/engine/retain/types.py b/core/dataplane/hms_api/engine/retain/types.py index 8c4246c..50357b6 100644 --- a/core/dataplane/hms_api/engine/retain/types.py +++ b/core/dataplane/hms_api/engine/retain/types.py @@ -11,6 +11,7 @@ from uuid import UUID from ..entity_resolution_contracts import EntityResolutionReadPlan +from .affect import AffectSignals class RetainContentDict(TypedDict, total=False): @@ -118,6 +119,7 @@ class ExtractedFact: occurred_end: datetime | None = None where: str | None = None # WHERE the fact occurred or is about causal_relations: list[CausalRelation] = field(default_factory=list) + affect: AffectSignals | None = None # Context from the content item content_index: int = 0 # Which content this fact came from @@ -162,6 +164,9 @@ class ProcessedFact: # Causal relations causal_relations: list[CausalRelation] = field(default_factory=list) + # Affect recognized during Retain extraction. Recall does not consume it. + affect: AffectSignals | None = None + # Chunk reference chunk_id: str | None = None @@ -235,6 +240,7 @@ def from_extracted_fact( metadata=extracted_fact.metadata, entities=entities, causal_relations=extracted_fact.causal_relations, + affect=extracted_fact.affect, chunk_id=chunk_id, content_index=extracted_fact.content_index, tags=extracted_fact.tags, diff --git a/core/dataplane/tests/test_db_abstraction.py b/core/dataplane/tests/test_db_abstraction.py index 677746d..cbf1095 100644 --- a/core/dataplane/tests/test_db_abstraction.py +++ b/core/dataplane/tests/test_db_abstraction.py @@ -498,6 +498,7 @@ def _make_batch(n: int = 3) -> dict: observation_scopes_list=[None] * n, text_signals_list=[None] * n, projection_jsons=["{}"] * n, + affect_jsons=[None] * n, ) @pytest.mark.asyncio @@ -539,7 +540,8 @@ async def test_client_generated_ids_are_inserted_and_returned_in_input_order( assert bank_id == "bank-pg" assert inserted_ids == generated assert fact_texts == batch["fact_texts"] - assert remaining[-1] == batch["projection_jsons"] + assert remaining[-2] == batch["projection_jsons"] + assert remaining[-1] == batch["affect_jsons"] assert result == [str(value) for value in generated] assert "$2::uuid[]" in query assert "AS t(id, text, embedding" in query @@ -592,6 +594,7 @@ def _make_batch(self, n: int = 2) -> dict: observation_scopes_list=[None] * n, text_signals_list=[None] * n, projection_jsons=["{}"] * n, + affect_jsons=[None] * n, ) @pytest.mark.asyncio @@ -653,6 +656,7 @@ async def test_column_values_correctly_mapped(self, ops, mock_conn): observation_scopes_list=["global"], text_signals_list=["positive"], projection_jsons=['{"embedding":{"ok":true}}'], + affect_jsons=['{"sentiment":"positive","emotion":"joy","intensity":0.8,"version":"affect-v1"}'], ) query, rows_data = mock_conn.executemany.call_args.args @@ -661,7 +665,7 @@ async def test_column_values_correctly_mapped(self, ops, mock_conn): # Verify column order matches: 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 + # chunk_id, document_id, tags, observation_scopes, text_signals, projection, affect assert row[0] == result[0], "row[0] should be the generated UUID" assert row[1] == "bank-42", "row[1] should be bank_id" assert row[2] == "The sky is blue", "row[2] should be text" @@ -679,17 +683,20 @@ async def test_column_values_correctly_mapped(self, ops, mock_conn): assert row[14] == "global", "row[14] should be observation_scopes" assert row[15] == "positive", "row[15] should be text_signals" assert row[16] == '{"embedding":{"ok":true}}', "row[16] should be projection JSON string" + assert row[17] == ( + '{"sentiment":"positive","emotion":"joy","intensity":0.8,"version":"affect-v1"}' + ), "row[17] should be affect JSON string" @pytest.mark.asyncio async def test_sql_column_count_matches_values(self, ops, mock_conn): - """The INSERT column list and VALUES placeholders must both have 17 entries.""" + """The INSERT column list and VALUES placeholders must both have 18 entries.""" batch = self._make_batch(1) await ops.insert_facts_batch(conn=mock_conn, **batch) query, _ = mock_conn.executemany.call_args.args # Extract the column list between "(" and ")" after INSERT INTO ... ( # and count the $N placeholders in VALUES - assert query.count("$") == 17, "VALUES clause must have 17 placeholders" + assert query.count("$") == 18, "VALUES clause must have 18 placeholders" @pytest.mark.asyncio async def test_tags_json_decoded_to_list(self, ops, mock_conn): diff --git a/core/dataplane/tests/test_retain_affect.py b/core/dataplane/tests/test_retain_affect.py new file mode 100644 index 0000000..8cc7dba --- /dev/null +++ b/core/dataplane/tests/test_retain_affect.py @@ -0,0 +1,271 @@ +"""Retain-only affect recognition tests.""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from hms_api.engine.response_models import TokenUsage +from hms_api.engine.retain.affect import AffectSignals, parse_affect +from hms_api.engine.retain.fact_extraction import ( + _build_extraction_prompt_and_schema, + _extract_facts_from_chunk, +) + + +def _config(*, enabled: bool) -> SimpleNamespace: + return SimpleNamespace( + retain_extraction_mode="concise", + retain_extract_causal_links=False, + retain_mission=None, + retain_custom_instructions=None, + entity_labels=None, + entities_allow_free_form=True, + retain_affect_enabled=enabled, + retain_affect_version="affect-v1", + retain_batch_enabled=False, + retain_llm_max_retries=1, + llm_max_retries=1, + retain_llm_initial_backoff=None, + llm_initial_backoff=0.0, + retain_llm_max_backoff=None, + llm_max_backoff=0.0, + retain_max_completion_tokens=8192, + ) + + +def test_affect_schema_is_opt_in_and_strict() -> None: + disabled_prompt, disabled_schema = _build_extraction_prompt_and_schema(_config(enabled=False)) + assert "AFFECT RECOGNITION" not in disabled_prompt + assert "AffectFact" not in disabled_schema.model_json_schema().get("$defs", {}) + + enabled_prompt, enabled_schema = _build_extraction_prompt_and_schema(_config(enabled=True)) + schema = enabled_schema.model_json_schema() + fact_schema = schema["$defs"]["AffectFact"] + affect_schema = schema["$defs"]["AffectAnnotation"] + + assert "AFFECT RECOGNITION" in enabled_prompt + assert "Do not attribute the assistant" in enabled_prompt + assert "empathetic tone to the user" in enabled_prompt + assert "affect" in fact_schema["required"] + assert set(affect_schema["properties"]["sentiment"]["enum"]) == { + "positive", + "negative", + "neutral", + } + assert set(affect_schema["properties"]["emotion"]["enum"]) == { + "joy", + "sadness", + "anger", + "fear", + "surprise", + "disgust", + "neutral", + } + assert affect_schema["properties"]["intensity"]["minimum"] == 0.0 + assert affect_schema["properties"]["intensity"]["maximum"] == 1.0 + + +def test_parse_affect_is_versioned_and_fail_soft() -> None: + affect = parse_affect( + {"sentiment": "negative", "emotion": "sadness", "intensity": 0.85}, + version="affect-v1", + ) + + assert affect == AffectSignals( + sentiment="negative", + emotion="sadness", + intensity=0.85, + version="affect-v1", + ) + assert ( + parse_affect( + {"sentiment": "negative", "emotion": "nostalgia", "intensity": 0.5}, + version="affect-v1", + ) + is None + ) + assert ( + parse_affect( + {"sentiment": "positive", "emotion": "joy", "intensity": 1.1}, + version="affect-v1", + ) + is None + ) + + +@pytest.mark.asyncio +async def test_extract_chunk_attaches_affect_when_enabled() -> None: + llm = MagicMock() + llm.call = AsyncMock( + return_value=( + { + "facts": [ + { + "what": "The user is delighted about shipping the release", + "when": None, + "who": "the user", + "why": "the release shipped", + "fact_type": "world", + "affect": { + "sentiment": "positive", + "emotion": "joy", + "intensity": 0.9, + }, + } + ] + }, + TokenUsage(), + ) + ) + + facts, _usage = await _extract_facts_from_chunk( + chunk="I am thrilled that we finally shipped!", + chunk_index=0, + total_chunks=1, + event_date=None, + context="chat message", + llm_config=llm, + config=_config(enabled=True), + agent_name="chatbot", + ) + + assert len(facts) == 1 + assert facts[0].affect == AffectSignals( + sentiment="positive", + emotion="joy", + intensity=0.9, + version="affect-v1", + ) + + +@pytest.mark.asyncio +async def test_extract_chunk_ignores_affect_when_disabled() -> None: + llm = MagicMock() + llm.call = AsyncMock( + return_value=( + { + "facts": [ + { + "what": "The user is angry", + "fact_type": "world", + "affect": { + "sentiment": "negative", + "emotion": "anger", + "intensity": 1.0, + }, + } + ] + }, + TokenUsage(), + ) + ) + + facts, _usage = await _extract_facts_from_chunk( + chunk="I am furious.", + chunk_index=0, + total_chunks=1, + event_date=None, + context="chat message", + llm_config=llm, + config=_config(enabled=False), + agent_name="chatbot", + ) + + assert len(facts) == 1 + assert facts[0].affect is None + + +@pytest.mark.asyncio +async def test_affect_survives_ingestion_projection_pipeline() -> None: + from hms_api.engine.ingestion.chunking import compute_content_hash + from hms_api.engine.ingestion.domain import ChunkPlan, freeze_json + from hms_api.engine.ingestion.extraction import ( + ExtractionMode, + ExtractionPolicy, + FactExtractorAdapter, + build_prechunked_extraction_layout, + ) + from hms_api.engine.ingestion.normalization import normalize_contents + from hms_api.engine.ingestion.projection.records import MemoryRecord, to_processed_fact + from hms_api.engine.retain.types import ChunkMetadata, ExtractedFact + + affect = AffectSignals( + sentiment="negative", + emotion="fear", + intensity=0.7, + version="affect-v1", + ) + item = normalize_contents( + ( + { + "content": "I am worried the deployment might fail.", + "context": "chat message", + "document_id": "document-1", + "event_date": None, + }, + ) + )[0] + chunk = ChunkPlan( + chunk_key="chunk-key", + source_index=item.source_index, + global_index=0, + local_index=0, + text=item.content, + content_hash=compute_content_hash(item.content), + ) + request = build_prechunked_extraction_layout((item,), (chunk,)).extraction_request( + ExtractionPolicy(mode=ExtractionMode.CONCISE) + ) + + async def primitive(**_kwargs): + return ( + [ + ExtractedFact( + fact_text="The user is worried the deployment might fail.", + fact_type="world", + affect=affect, + content_index=0, + chunk_index=0, + context="chat message", + ) + ], + [ + ChunkMetadata( + chunk_text=item.content, + fact_count=1, + content_index=0, + chunk_index=0, + ) + ], + TokenUsage(), + ) + + result = await FactExtractorAdapter( + llm_config=object(), + config=SimpleNamespace( + retain_extraction_mode="concise", + retain_batch_enabled=False, + retain_chunk_size=3000, + ), + agent_name="chatbot", + sync_primitive=primitive, + batch_primitive=primitive, + ).extract(request) + + candidate = result.candidates[0] + record = MemoryRecord.from_candidate( + candidate, + embedding=None, + projection=freeze_json({"extraction": {"v": "test"}}), + ) + processed = to_processed_fact( + record, + document_id="document-1", + chunk_id="chunk-1", + content_index=0, + ) + + assert candidate.affect is affect + assert record.affect is affect + assert processed.affect is affect diff --git a/docker-compose.yml b/docker-compose.yml index 592d1c2..c16627c 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -46,6 +46,8 @@ services: 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_RETAIN_AFFECT_ENABLED: ${HMS_API_RETAIN_AFFECT_ENABLED:-false} + HMS_API_RETAIN_AFFECT_VERSION: ${HMS_API_RETAIN_AFFECT_VERSION:-affect-v1} 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}