"""记忆抽取:模型严格 JSON + Pydantic 校验 + 受控词表。 抽取是"从用户消息里读出结构化事实",不是"把用户原文存下来":模型必须返回 `memory_key` / `value` / `memory_type` / `confidence` 四字段的严格 JSON,任何 契约违规(JSON 非法、字段缺失、置信度越界、键不在受控词表、键与类型不一致) 一律失败关闭,抛 `RecoverableAgentError`,由事件层决定重试,绝不落库。 端点解析复用现有 resolver 风格(`DatabaseModelEndpointResolver`),task_type 固定为 `memory_extraction`,与意图分类走同一套模型路由入口。 """ import json import logging from collections.abc import Awaitable, Callable from functools import lru_cache from typing import Any, Protocol from pydantic import BaseModel, ConfigDict, Field, ValidationError from sqlalchemy import select from app.core.errors import RecoverableAgentError from app.infrastructure.db import SessionFactory from app.model.configuration import ConfigRelease, PlatformConfigItem from app.service.memory_taxonomy import ( DEFAULT_MEMORY_KEYS, MEMORY_KEY_CONFIG_KEY, MEMORY_KEY_FIELD, MEMORY_KEY_NAMESPACE, MEMORY_TYPES, ) from app.service.model_gateway import DatabaseModelEndpointResolver, ModelGenerationService logger = logging.getLogger(__name__) TASK_TYPE = "memory_extraction" DEFAULT_AGENT_TYPE = "customer_service" DEFAULT_MIN_CONFIDENCE = 0.6 VALUE_MAX_LENGTH = 200 class ExtractionEndpointResolver(Protocol): async def resolve(self, *, agent_type: str, task_type: str) -> list[Any]: ... class _ExtractionPayload(BaseModel): """模型输出的唯一合法形状;额外字段一律拒绝(extra="forbid")。""" model_config = ConfigDict(extra="forbid") memory_key: str | None value: str | None memory_type: str | None confidence: float = Field(ge=0, le=1) class ExtractedMemory(BaseModel): """校验通过的结构化记忆;`value` 是短语,不是用户原文整句。""" model_config = ConfigDict(extra="forbid", frozen=True) memory_key: str = Field(min_length=1, max_length=128) value: str = Field(min_length=1, max_length=VALUE_MAX_LENGTH) memory_type: str = Field(min_length=1, max_length=24) confidence: float = Field(ge=0, le=1) async def load_memory_keys() -> frozenset[str]: """从配置中心 `namespace=memory` 扩展受控词表,读不到就用代码默认词表。 配置中心异常绝不能让记忆整体不可用,因此这里不向上抛错,只记录告警。 """ try: async with SessionFactory() as session: release = await session.scalar( select(ConfigRelease).where(ConfigRelease.status == "active") ) if release is None: return DEFAULT_MEMORY_KEYS item = await session.scalar(select(PlatformConfigItem).where( PlatformConfigItem.release_id == release.id, PlatformConfigItem.namespace == MEMORY_KEY_NAMESPACE, PlatformConfigItem.config_key == MEMORY_KEY_CONFIG_KEY, )) except Exception: # 配置中心不可用(网络/表缺失)时降级,不影响记忆可用性。 logger.warning("memory key vocabulary unavailable; using code defaults") return DEFAULT_MEMORY_KEYS if item is None: return DEFAULT_MEMORY_KEYS raw = item.value_json.get(MEMORY_KEY_FIELD, []) if not isinstance(raw, list): logger.warning("memory key vocabulary config ignored: %s is not a list", MEMORY_KEY_FIELD) return DEFAULT_MEMORY_KEYS extra = frozenset( entry.strip() for entry in raw if isinstance(entry, str) and entry.strip() ) return DEFAULT_MEMORY_KEYS | extra if extra else DEFAULT_MEMORY_KEYS class MemoryExtractionService: """把用户消息抽取为受控键 + 结构化值;不合规输出一律失败关闭。""" def __init__( self, model_service: ModelGenerationService, endpoint_resolver: ExtractionEndpointResolver, *, min_confidence: float = DEFAULT_MIN_CONFIDENCE, vocabulary_loader: Callable[[], Awaitable[frozenset[str]]] | None = None, ) -> None: if not 0 <= min_confidence <= 1: raise ValueError("min_confidence must be between 0 and 1") self.model_service = model_service self.endpoint_resolver = endpoint_resolver self.min_confidence = min_confidence self.vocabulary_loader = vocabulary_loader or load_memory_keys async def extract( self, *, message: str, agent_type: str = DEFAULT_AGENT_TYPE ) -> ExtractedMemory | None: """返回校验通过的记忆;模型显式声明"无持久事实"时返回 None。 失败关闭:端点缺失、模型调用失败、输出违约一律抛 `RecoverableAgentError`, 调用方不得写入任何记忆。 """ text = message.strip() if not text: return None vocabulary = await self._vocabulary() endpoints = await self.endpoint_resolver.resolve( agent_type=agent_type, task_type=TASK_TYPE ) if not endpoints: raise RecoverableAgentError("没有可用的记忆抽取模型端点") execution = await self.model_service.generate(endpoints, self._prompt(text, vocabulary)) return self._validate(self._parse(execution.text), vocabulary) async def _vocabulary(self) -> frozenset[str]: keys = await self.vocabulary_loader() return keys if keys else DEFAULT_MEMORY_KEYS @staticmethod def _prompt(message: str, vocabulary: frozenset[str]) -> str: keys = ", ".join(sorted(vocabulary)) types = ", ".join(sorted(MEMORY_TYPES)) return ( "你是金融客户记忆抽取器。只输出一个 JSON 对象,不要 Markdown,不要解释。" "字段固定为 memory_key、value、memory_type、confidence,不得增删字段。" f"memory_key 只能取以下之一:{keys}。" f"memory_type 取 memory_key 冒号前的前缀,只能是以下之一:{types}。" "confidence 是 0 到 1 的数字。" f"value 必须是指代明确的结构化短语(例如 \"稳健型\"、\"随时可赎回\")," f"不得复制用户原句,长度不超过 {VALUE_MAX_LENGTH} 字。" "只有用户明确陈述的长期偏好、约束、身份或目标才抽取;" "普通提问、寒暄、一次性操作意图一律不抽取,此时 memory_key、value、" "memory_type 三个字段输出 null,confidence 输出 0。" f"用户消息:{message}" ) @staticmethod def _parse(text: str) -> _ExtractionPayload: candidate = text.strip() if candidate.startswith("```"): candidate = candidate.removeprefix("```") candidate = candidate.removeprefix("json").removesuffix("```").strip() try: data = json.loads(candidate) except (json.JSONDecodeError, TypeError) as exc: raise RecoverableAgentError("模型记忆抽取输出不是有效 JSON") from exc # 模型有时把结果包在**单元素数组**里,而契约是**对象**。实测出现于 episode 抽取路径: # [{"memory_key": null, "value": null, "memory_type": null, "confidence": 0}] # 该形状此前直接进 pydantic 校验、报 `Input should be a valid dictionary`, # 被记成"输出不是有效 JSON"并反复重试直到 episode 判失败 —— 而它其实是**合法的空结果** # (`_validate` 已能把"三字段为 null 且 confidence=0"识别为"无持久事实")。 # 只归一化"恰好一个对象的数组";多元素或元素非对象时**不猜**,仍交给校验失败关闭。 if isinstance(data, list) and len(data) == 1 and isinstance(data[0], dict): data = data[0] try: return _ExtractionPayload.model_validate(data) except (ValidationError, TypeError) as exc: raise RecoverableAgentError("模型记忆抽取输出不是有效 JSON") from exc def _validate( self, payload: _ExtractionPayload, vocabulary: frozenset[str] ) -> ExtractedMemory | None: if payload.memory_key is None and payload.value is None and payload.memory_type is None: # 模型显式声明"没有持久事实":合法的空结果,不是失败。 if payload.confidence != 0: raise RecoverableAgentError("模型空结果的置信度必须为 0") return None if payload.memory_key is None or payload.value is None or payload.memory_type is None: raise RecoverableAgentError("模型记忆抽取字段缺失") if payload.memory_key not in vocabulary: raise RecoverableAgentError("模型返回的记忆键不在受控词表中") if payload.memory_type not in MEMORY_TYPES: raise RecoverableAgentError("模型返回的记忆类型不在受控集合中") prefix = payload.memory_key.split(":", 1)[0] if prefix != payload.memory_type: raise RecoverableAgentError("模型返回的记忆键与记忆类型不一致") if payload.confidence < self.min_confidence: raise RecoverableAgentError("模型抽取置信度低于阈值") try: return ExtractedMemory( memory_key=payload.memory_key, value=payload.value.strip(), memory_type=payload.memory_type, confidence=payload.confidence, ) except ValidationError as exc: raise RecoverableAgentError("模型记忆抽取内容不合规") from exc @lru_cache(maxsize=1) def get_memory_extraction_service() -> MemoryExtractionService: """生产装配:记忆抽取与业务 Agent 共用同一个 ModelGenerationService。""" from app.service.agent.bootstrap import get_model_service return MemoryExtractionService(get_model_service(), DatabaseModelEndpointResolver())