"""会话片段(episode)抽取:把会话消息按"会话片段"聚合后写入 `episodes`。 为什么需要它(对应 P3 缺口 2):基线把 `episodes` 定义为中期记忆的抽取单位 ("会话结束、达到 Token 阈值或转人工时生成片段"),而现有实现按**单条消息**抽取, `episodes` 表此前零 ORM、零使用。这里补上片段粒度与幂等落库。 聚合粒度:同一 `customer_id + session_id` 的消息按 `created_at` 升序,相邻消息 间隔超过 `gap_minutes`(默认 30 分钟)即切分为新片段;单个片段不足 `min_messages` 条时并入上一片段,避免尾部产生无意义碎片。摘要在片段内按 `[消息序号 | 角色] 正文` 拼接,只截断不改造原文。 幂等键:`content_hash = sha256(客户 + 会话 + 每条消息的 id 与正文)` 直接落到表上的 唯一键 `content_hash`。同一片段内容不变时哈希不变 → 重复调用不会产生重复片段; 片段内容变化(新消息并入或编辑)时哈希变化 → 视为新片段,旧片段保留为历史。 冲突时依赖唯一键拦截(`IntegrityError` 回滚到保存点),不使用"先查后写"作为唯一防线。 消费环节(对应 P3 缺口 4):`EpisodeExtractionConsumer` 把已落库、状态为待提取的片段 接入记忆抽取,写入 `memory_unit` 后把片段置为已处理(`已完成` + `promoted_to_ltm`)。 失败只累加 `retry_count` 并把状态置为 `失败`,片段保持可重试、记忆一行不写;重复消费 由 `content_hash` 派生的 `MemoryEvidence.idempotency_key` 兜底,不产生第二条证据。 """ import hashlib import json import logging from dataclasses import dataclass, field from datetime import UTC, datetime, timedelta from uuid import NAMESPACE_URL, uuid4, uuid5 from sqlalchemy import select from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession from app.model.conversation import ConversationMessage from app.model.episode import Episode from app.model.platform import DomainEventOutbox from app.service.memory_extraction_service import ( MemoryExtractionService, get_memory_extraction_service, ) from app.service.memory_service import CacheDeleteAdapter, MemoryService logger = logging.getLogger(__name__) DEFAULT_GAP_MINUTES = 30 DEFAULT_MIN_MESSAGES = 2 SUMMARY_LIMIT = 4000 STATUS_PENDING = "待提取" STATUS_PROCESSING = "处理中" STATUS_DONE = "已完成" STATUS_FAILED = "失败" RETRYABLE_STATUSES = frozenset({STATUS_FAILED, STATUS_PENDING, STATUS_PROCESSING}) # 消费环节可选中的状态:只处理"待提取"与"失败","处理中"留给并发消费者自行收尾。 CONSUMABLE_STATUSES = (STATUS_PENDING, STATUS_FAILED) # 失败片段的重试上限:达到上限后停在"失败",不再无限重试(避免毒片段拖住批处理)。 MAX_EXTRACTION_RETRY = 3 USER_ROLE = "user" EVIDENCE_TYPE = "会话片段" SOURCE_TYPE = "用户自述" EVIDENCE_WEIGHT = 0.05 @dataclass class EpisodeResult: """一次抽取的结果;`inserted` 为新片段,`skipped` 为命中幂等边界。""" inserted: list[int] = field(default_factory=list) skipped: list[int] = field(default_factory=list) @property def total(self) -> int: return len(self.inserted) + len(self.skipped) class EpisodeWorker: """按会话片段聚合消息并落库;可重复调用,失败片段可由调用方重试。""" def __init__( self, session: AsyncSession, *, gap_minutes: int = DEFAULT_GAP_MINUTES, min_messages: int = DEFAULT_MIN_MESSAGES, message_limit: int = 500, ) -> None: self.session = session self.gap = timedelta(minutes=max(1, gap_minutes)) self.min_messages = max(1, min_messages) self.message_limit = max(1, min(message_limit, 2000)) async def aggregate( self, customer_id: int, *, session_id: str | None = None, ended_before: datetime | None = None, status: str = STATUS_PENDING, ) -> EpisodeResult: """聚合并落库;`ended_before` 用于只处理已结束(静默超过间隔)的片段。""" now = datetime.now(UTC).replace(tzinfo=None) cutoff = ended_before or (now - self.gap) result = EpisodeResult() for key_session, messages in await self._pending_sessions(customer_id, session_id, cutoff): for segment in self._segment(messages): if len(segment) < self.min_messages: continue episode = await self._persist(customer_id, key_session, segment, status) target = result.inserted if episode is not None else result.skipped target.append(segment[0].id) return result async def _pending_sessions( self, customer_id: int, session_id: str | None, cutoff: datetime ) -> list[tuple[str, list[ConversationMessage]]]: conditions = [ ConversationMessage.customer_id == customer_id, ConversationMessage.created_at <= cutoff, ] if session_id is not None: conditions.append(ConversationMessage.session_id == session_id) rows = await self.session.scalars( select(ConversationMessage) .where(*conditions) .order_by(ConversationMessage.session_id, ConversationMessage.created_at, ConversationMessage.id) .limit(self.message_limit) ) grouped: dict[str, list[ConversationMessage]] = {} for message in rows: grouped.setdefault(message.session_id, []).append(message) return list(grouped.items()) def _segment(self, messages: list[ConversationMessage]) -> list[list[ConversationMessage]]: """按时间间隔切分;尾部不足 `min_messages` 的片段并入上一片段。""" segments: list[list[ConversationMessage]] = [] current: list[ConversationMessage] = [] for message in messages: if current and message.created_at - current[-1].created_at > self.gap: segments.append(current) current = [] current.append(message) if current: segments.append(current) if len(segments) > 1 and len(segments[-1]) < self.min_messages: tail = segments.pop() segments[-1].extend(tail) return segments async def _persist( self, customer_id: int, session_id: str, segment: list[ConversationMessage], status: str, ) -> Episode | None: content_hash = self.content_hash(customer_id, session_id, segment) existing = await self.session.scalar( select(Episode).where(Episode.content_hash == content_hash) ) if existing is not None: # ⚠️ 这里**不要**碰 `retry_count`。 # # 此前会调 `_touch_retry(existing)` 做 `retry_count += 1`(注释写的是 # "仅在未完成时累加重试计数")。但 `retry_count` 的语义是**抽取失败**的次数 # (由 `_mark_failed` 累加),而"分片逻辑每轮重新看到同一段会话"根本不是失败 # —— 片段内容没变,`content_hash` 才会相同。 # # 两者混用一个字段的后果是实测出来的:客户 9001 有 **40 条片段 # `retry_count=1405`**,远超 `MAX_EXTRACTION_RETRY`(3),于是 # `consume_pending` 的 `retry_count < max_retry` 永远筛不中它们 —— # **片段永久滞留 → 新记忆进不来 → 画像停在旧值** # (`profile_snapshots` 停在 2026-09-10 13:57,而 `memory_unit.updated_at` # 已经是 2026-09-13 11:03)。 # # 重复看到同一片段就只是看到,不改变任何状态。 logger.debug("episode already persisted content_hash=%s", content_hash) return None episode = self._build(customer_id, session_id, segment, content_hash, status) self.session.add(episode) try: async with self.session.begin_nested(): await self.session.flush() except IntegrityError: # 并发抽取命中 uk_episodes_content_hash:片段已由另一事务写入,本轮跳过。 logger.warning("episode already exists content_hash=%s", content_hash) return None return episode def _build( self, customer_id: int, session_id: str, segment: list[ConversationMessage], content_hash: str, status: str, ) -> Episode: started_at = segment[0].created_at ended_at = segment[-1].created_at portals = sorted({message.portal for message in segment if message.portal}) now = datetime.now(UTC).replace(tzinfo=None) numbers = [message.message_no for message in segment if message.message_no is not None] return Episode( # 跨存储稳定标识由幂等键派生,重试与重复消费得到同一个 uuid。 episode_uuid=str(uuid5(NAMESPACE_URL, f"jr:episode:{session_id}:{content_hash}")), customer_id=customer_id, session_id=session_id[:64], start_message_no=min(numbers) if numbers else None, end_message_no=max(numbers) if numbers else None, portals_involved=portals, summary=self.summarize(segment), extraction_status=status, content_hash=content_hash, retry_count=0, started_at=started_at, ended_at=ended_at, start_at=started_at, end_at=ended_at, promoted_to_ltm=False, created_at=now, ) @staticmethod def content_hash( customer_id: int, session_id: str, segment: list[ConversationMessage] ) -> str: """片段指纹:消息 id 与正文共同决定,保证内容变化即产生新指纹。""" payload = json.dumps( { "customer_id": customer_id, "session_id": session_id, "messages": [ {"id": message.id, "role": message.role, "content": (message.content or "").strip()} for message in segment ], }, ensure_ascii=False, sort_keys=True, separators=(",", ":"), ) return hashlib.sha256(payload.encode("utf-8")).hexdigest() @staticmethod def summarize(segment: list[ConversationMessage]) -> str: """片段摘要:只做拼接与截断,不改写、不脱敏(脱敏由上游写入前完成)。""" parts: list[str] = [] for message in segment: number = message.message_no if message.message_no is not None else message.id content = (message.content or "").strip().replace("\r\n", " ").replace("\n", " ") parts.append(f"[{number} | {message.role}] {content}") text = "\n".join(parts) return text[:SUMMARY_LIMIT] @dataclass class EpisodeConsumptionResult: """一次消费的结果;`extracted` 为写出记忆的片段,`failed` 为保持可重试的片段。""" extracted: list[int] = field(default_factory=list) no_fact: list[int] = field(default_factory=list) failed: list[int] = field(default_factory=list) @property def processed(self) -> int: """已终结(不再重试)的片段数:写出记忆的与判定无持久事实的。""" return len(self.extracted) + len(self.no_fact) class EpisodeExtractionConsumer: """消费待提取片段:复用 `MemoryExtractionService` 抽取后写入 `memory_unit`。 幂等由三层保证,都不需要新增库表或列: 1. 选中条件排除 `已完成`/`promoted_to_ltm` 的片段,成功片段不会被二次处理; 2. 证据幂等键 `episode.extraction:{content_hash}`(表上唯一键 `content_hash` 派生) 命中 `memory_evidence` 唯一键时视为已消费,不重复写证据; 3. `memory_unit` 按 `(customer_id, memory_key)` 语义更新,重复抽取同一受控键 不会堆积重复记忆行。 失败(模型不可用、输出违约)不清空已有进度:只累加 `retry_count` 并把状态置为 `失败`,记忆一行不写,片段在下一次调度继续可被选中;达到 `max_retry` 后停在该 状态,由运维按 `retry_count` 排查毒片段,而不是无限重试拖住批处理。 """ def __init__( self, session: AsyncSession, *, extractor: MemoryExtractionService | None = None, cache: CacheDeleteAdapter | None = None, max_retry: int = MAX_EXTRACTION_RETRY, limit: int = 20, ) -> None: self.session = session self.extractor = extractor if extractor is not None else get_memory_extraction_service() # 召回热缓存适配器:写入生效后必须失效,否则新记忆在 TTL 内召回不到。 self.cache = cache self.max_retry = max(1, max_retry) self.limit = max(1, min(limit, 200)) async def consume_pending(self, *, limit: int | None = None) -> EpisodeConsumptionResult: """消费一批待提取片段;单个片段失败不影响同批其它片段。""" bounded = max(1, min(limit if limit is not None else self.limit, 200)) episodes = list(await self.session.scalars( select(Episode) .where( Episode.extraction_status.in_(CONSUMABLE_STATUSES), Episode.retry_count < self.max_retry, Episode.promoted_to_ltm.is_(False), ) .order_by(Episode.id) .limit(bounded) )) result = EpisodeConsumptionResult() for episode in episodes: await self._consume(episode, result) return result async def _consume(self, episode: Episode, result: EpisodeConsumptionResult) -> None: episode_id = int(episode.id) text = self.user_content(episode.summary) if not text: # 片段里没有用户陈述(例如纯助手片段):没有可抽取的事实,直接置为已处理。 await self._mark_done(episode, promoted=False) result.no_fact.append(episode_id) return try: extracted = await self.extractor.extract(message=text) except Exception: # 失败保持可重试:状态置失败并累加重试计数,记忆一行不写。 await self._mark_failed(episode) logger.warning( "episode extraction failed episode_id=%s retry_count=%s", episode_id, episode.retry_count, exc_info=True, ) result.failed.append(episode_id) return if extracted is None: # 模型判定该片段没有持久事实:不是错误,但也没有可写的记忆。 await self._mark_done(episode, promoted=False) result.no_fact.append(episode_id) return service = MemoryService(self.session, cache=self.cache) memory = await service.upsert( episode.customer_id, extracted.memory_key, extracted.value, memory_type=extracted.memory_type, confidence=extracted.confidence, source_type=SOURCE_TYPE, structured_value={ "memory_key": extracted.memory_key, "value": extracted.value, "memory_type": extracted.memory_type, "confidence": extracted.confidence, "session_id": episode.session_id, "episode_uuid": episode.episode_uuid, }, ) recorded = await service.record_evidence( memory, # 幂等键落在片段指纹上:同一片段重复消费不会写出第二条证据。 idempotency_key=self.idempotency_key(episode), evidence_type=EVIDENCE_TYPE, excerpt=text, snapshot={ "episode_uuid": episode.episode_uuid, "content_hash": episode.content_hash, "customer_id": episode.customer_id, "session_id": episode.session_id, "start_message_no": episode.start_message_no, "end_message_no": episode.end_message_no, "memory_key": extracted.memory_key, "value": extracted.value, "confidence": extracted.confidence, }, weight=EVIDENCE_WEIGHT, source_table="episodes", source_record_id=str(episode.id), occurred_at=episode.ended_at or episode.created_at, ) if recorded: # ⚠️ 证据写入后必须把「画像重建」**投成事件**,不能直接调用重建: # 本方法的记忆写入还在当前事务里、尚未提交,另开 session 去重建看不到 # 这条新记忆(`memory_extraction_worker` 里记录了这条实测结论)。 # # 这条路径此前**完全没有**投重建事件,后果是:从会话片段抽取出来的记忆 # 永远到不了画像。实测证据 —— 客户 9001 的 `profile_snapshots` current # 停在 `2026-09-10 13:57`(v7),而它的 `memory_unit.updated_at` 已经是 # `2026-09-13 11:03`、`evidence_count` 涨到 4;那三条新证据的 # `source_table` 正是 `episodes`。也就是说**片段链路记住了,画像不知道**。 # # 只在 `recorded=True` 时投:幂等命中(同一片段重复消费)时证据与计数 # 都没有净变化,投一次重建是白跑。 now = datetime.now(UTC).replace(tzinfo=None) self.session.add(DomainEventOutbox( id=0, event_id=str(uuid4()), event_type="profile.rebuild_requested", aggregate_type="customer_profile", aggregate_id=str(episode.customer_id), trace_id=episode.episode_uuid, payload={"customer_id": episode.customer_id, "trigger": "episode_extraction"}, status="pending", retry_count=0, occurred_at=now, created_at=now, updated_at=now, )) await self._mark_done(episode, promoted=True) result.extracted.append(episode_id) async def _mark_done(self, episode: Episode, *, promoted: bool) -> None: episode.extraction_status = STATUS_DONE episode.promoted_to_ltm = promoted await self.session.flush() async def _mark_failed(self, episode: Episode) -> None: episode.extraction_status = STATUS_FAILED episode.retry_count += 1 await self.session.flush() @staticmethod def idempotency_key(episode: Episode) -> str: """片段级幂等键;`content_hash` 缺失(列可空)时退化为 episode_uuid / 主键。""" fingerprint = episode.content_hash or episode.episode_uuid or str(episode.id) return f"episode.extraction:{fingerprint}" @staticmethod def user_content(summary: str) -> str: """从片段摘要里取出**用户陈述**:助手回复不是记忆来源,不参与抽取。 摘要行格式由 `EpisodeWorker.summarize` 固定为 `[序号 | 角色] 正文`; 这里只按该格式取 `user` 行,取不到就当作"没有可抽取的用户事实"。 """ marker = f"| {USER_ROLE}] " parts: list[str] = [] for line in (summary or "").splitlines(): index = line.find(marker) if index < 0: continue content = line[index + len(marker):].strip() if content: parts.append(content) return "\n".join(parts)