Files
group_fqcd_jr/app/service/memory_extraction_service.py
T

218 lines
10 KiB
Python
Raw Normal View History

"""记忆抽取:模型严格 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())