Coverage for merco/memory/save_pipeline.py: 89%
81 statements
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-07 14:04 +0800
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-07 14:04 +0800
1"""Memory 保存链 — Strategy 通过它写入 MemoryStore"""
2from __future__ import annotations
4import logging
5from abc import ABC, abstractmethod
6from dataclasses import dataclass, field
7from typing import Literal
9logger = logging.getLogger("merco.memory.save_pipeline")
12MemorySource = Literal["user", "extracted", "system"]
15SOURCE_PRIORITY: dict[str, int] = {
16 "user": 3,
17 "extracted": 2,
18 "system": 1,
19}
22@dataclass
23class SaveItem:
24 """Pipeline 输入单元"""
25 key: str
26 value: str
27 source: MemorySource
28 tags: list[str] = field(default_factory=list)
29 session_id: str = ""
30 metadata: dict = field(default_factory=dict)
33class MemorySaveProcessor(ABC):
34 """保存链处理器基类"""
35 name: str = ""
37 @abstractmethod
38 async def process(self, item: SaveItem) -> SaveItem | None:
39 """返回 None = 跳过该 item"""
40 ...
43class SourceEnricher(MemorySaveProcessor):
44 """自动补 [source] 前缀到 tags"""
45 name = "source_enricher"
47 async def process(self, item: SaveItem) -> SaveItem:
48 prefix = f"[{item.source}]"
49 if prefix not in item.tags:
50 item.tags.insert(0, prefix)
51 return item
54class DedupProcessor(MemorySaveProcessor):
55 """按 source 优先级 skip 已有 key"""
56 name = "dedup"
58 def __init__(self, store):
59 self._store = store
61 async def process(self, item: SaveItem) -> SaveItem | None:
62 existing = self._store.load(item.key)
63 if not existing:
64 return item
65 existing_tags = existing.get("tags", []) or []
66 existing_source = self._infer_source(existing_tags)
67 new_priority = SOURCE_PRIORITY.get(item.source, 0)
68 existing_priority = SOURCE_PRIORITY.get(existing_source, 0)
69 if new_priority <= existing_priority:
70 return None
71 return item
73 @staticmethod
74 def _infer_source(tags: list[str]) -> str:
75 """从 tags 推断 source。无 [source] 标记视为 system(最低,向后兼容旧记录)"""
76 for t in tags:
77 if t.startswith("[") and t.endswith("]"):
78 inner = t[1:-1]
79 if inner in SOURCE_PRIORITY:
80 return inner
81 return "system"
84class MemorySavePipeline:
85 """统一的 Memory 保存链 — Strategy 通过它写入"""
87 def __init__(self, store, hooks):
88 self.store = store
89 self.hooks = hooks
90 self._processors: list[MemorySaveProcessor] = [
91 SourceEnricher(),
92 DedupProcessor(store),
93 ]
95 def use(self, processor: MemorySaveProcessor) -> "MemorySavePipeline":
96 self._processors.append(processor)
97 return self
99 async def save(self, item: SaveItem) -> bool:
100 """返回 True=写入成功,False=被 dedup skip"""
101 for p in self._processors:
102 try:
103 item = await p.process(item)
104 except Exception as e:
105 logger.warning("MemorySaveProcessor '%s' 异常: %s", p.name, e)
106 return False
107 if item is None:
108 return False
109 try:
110 self.store.save(item.key, item.value, tags=item.tags)
111 except Exception as e:
112 logger.warning("MemoryStore.save 失败 [%s]: %s", item.key, e)
113 try:
114 await self.hooks.emit("memory.failed", key=item.key, error=str(e))
115 except Exception as hook_err:
116 logger.debug("hooks.emit('memory.failed') 失败: %s", hook_err)
117 return False
118 try:
119 await self.hooks.emit(
120 "memory.saved",
121 key=item.key, value=item.value,
122 source=item.source, tags=item.tags,
123 )
124 except Exception as hook_err:
125 logger.debug("hooks.emit('memory.saved') 失败: %s", hook_err)
126 return True