Files
group_fqcd_jr/app/worker/episode_worker.py
T

445 lines
20 KiB
Python
Raw Normal View History

"""会话片段(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)