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

1"""Memory 保存链 — Strategy 通过它写入 MemoryStore""" 

2from __future__ import annotations 

3 

4import logging 

5from abc import ABC, abstractmethod 

6from dataclasses import dataclass, field 

7from typing import Literal 

8 

9logger = logging.getLogger("merco.memory.save_pipeline") 

10 

11 

12MemorySource = Literal["user", "extracted", "system"] 

13 

14 

15SOURCE_PRIORITY: dict[str, int] = { 

16 "user": 3, 

17 "extracted": 2, 

18 "system": 1, 

19} 

20 

21 

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) 

31 

32 

33class MemorySaveProcessor(ABC): 

34 """保存链处理器基类""" 

35 name: str = "" 

36 

37 @abstractmethod 

38 async def process(self, item: SaveItem) -> SaveItem | None: 

39 """返回 None = 跳过该 item""" 

40 ... 

41 

42 

43class SourceEnricher(MemorySaveProcessor): 

44 """自动补 [source] 前缀到 tags""" 

45 name = "source_enricher" 

46 

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 

52 

53 

54class DedupProcessor(MemorySaveProcessor): 

55 """按 source 优先级 skip 已有 key""" 

56 name = "dedup" 

57 

58 def __init__(self, store): 

59 self._store = store 

60 

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 

72 

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" 

82 

83 

84class MemorySavePipeline: 

85 """统一的 Memory 保存链 — Strategy 通过它写入""" 

86 

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 ] 

94 

95 def use(self, processor: MemorySaveProcessor) -> "MemorySavePipeline": 

96 self._processors.append(processor) 

97 return self 

98 

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