206 lines
9.0 KiB
Python
206 lines
9.0 KiB
Python
"""记忆抽取:模型严格 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:
|
|||
|
|
return _ExtractionPayload.model_validate(json.loads(candidate))
|
|||
|
|
except (json.JSONDecodeError, 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())
|