## 现象(2026-09-12 跑真实 Worker 时发现)
`episode_worker` 反复报 `模型记忆抽取输出不是有效 JSON` 并重试到失败:
```
pydantic_core.ValidationError: Input should be a valid dictionary or instance of _ExtractionPayload
input_value=[{'memory_key': None, 'value': None, 'memory_type': None, 'confidence': 0}]
input_type=list
```
## 根因
契约是**对象** `{...}`,而模型在 episode 抽取路径会返回**单元素数组** `[{...}]`。
`_parse` 直接 `json.loads` 后交给 pydantic,数组自然过不了 `model_validate`,
于是被归入"输出不是有效 JSON"这一条 —— 但它其实是**合法的空结果**
(`_validate` 已能把"三字段为 null 且 confidence=0"正确识别为"无持久事实",返回 None)。
后果:这类 episode 白跑一遍模型调用、重试到 `retry_count` 上限后判失败,
**该片段的记忆永远抽不出来**。
## 修法
`_parse` 里只对"**恰好一个对象**的数组"做归一化:
- `[{...}]` → 取 `{...}`(空结果照常返回 None;有事实照常解析)
- 多元素数组、元素非对象、空数组 → **不猜**,仍交给校验失败关闭
(多元素时无法判断哪个是答案,猜错会把错误记忆写进库,比失败更糟)
## 测试
`tests/unit/service/test_memory_extraction_service.py` 追加 2 个用例:
- `test_single_element_array_is_normalized`:数组包空结果 → None;数组包有事实 → 正常解析
- `test_multi_element_or_non_dict_array_still_fails_closed`:多元素 / 非对象元素 / 空数组 → 仍失败
## 验证
- `pytest tests/unit/service/test_memory_extraction_service.py tests/unit/worker/test_memory_extraction_worker.py` → 28 passed
- `mypy app` → 0 错 / 245 文件
> 说明:`memory_extraction_service.py` 属主干线代码。此处是从**实际运行日志**里发现的
> 健壮性缺陷,改动限定在输出形状归一化,不改变抽取契约与校验规则。
218 lines
10 KiB
Python
218 lines
10 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:
|
||
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())
|