Coverage for src\core\sync.py: 0%

258 statements  

« prev     ^ index     » next       coverage.py v7.3.4, created at 2026-04-21 14:54 +0800

1""" 

2记忆同步系统 

3实现跨设备的记忆同步和冲突解决 

4""" 

5 

6import json 

7import time 

8import hashlib 

9from typing import Dict, List, Optional, Any, Set, Tuple 

10from dataclasses import dataclass, field 

11from enum import Enum 

12import asyncio 

13from datetime import datetime, timedelta 

14 

15from .memory import MemorySystem, MemoryItem, MemoryQuery, MemoryType 

16from .device import DeviceInfo 

17 

18 

19class SyncStatus(str, Enum): 

20 """同步状态""" 

21 PENDING = "pending" # 等待同步 

22 SYNCING = "syncing" # 同步中 

23 COMPLETED = "completed" # 同步完成 

24 FAILED = "failed" # 同步失败 

25 CONFLICT = "conflict" # 冲突需要解决 

26 

27 

28class SyncDirection(str, Enum): 

29 """同步方向""" 

30 PUSH = "push" # 推送(本地→远程) 

31 PULL = "pull" # 拉取(远程→本地) 

32 BIDIRECTIONAL = "bidirectional" # 双向同步 

33 

34 

35class SyncStrategy(str, Enum): 

36 """同步策略""" 

37 LAST_WRITE_WINS = "last_write_wins" # 最后写入获胜 

38 MANUAL_MERGE = "manual_merge" # 手动合并 

39 AUTO_MERGE = "auto_merge" # 自动合并 

40 SOURCE_PRIORITY = "source_priority" # 源设备优先 

41 

42 

43@dataclass 

44class SyncOperation: 

45 """同步操作""" 

46 

47 id: str 

48 memory_id: str 

49 device_id: str 

50 direction: SyncDirection 

51 status: SyncStatus = SyncStatus.PENDING 

52 timestamp: float = field(default_factory=time.time) 

53 retry_count: int = 0 

54 error_message: Optional[str] = None 

55 

56 # 冲突相关 

57 conflict_resolution: Optional[Dict[str, Any]] = None 

58 merged_memory_id: Optional[str] = None 

59 

60 

61@dataclass 

62class SyncSession: 

63 """同步会话""" 

64 

65 id: str 

66 source_device_id: str 

67 target_device_id: str 

68 start_time: float = field(default_factory=time.time) 

69 end_time: Optional[float] = None 

70 status: SyncStatus = SyncStatus.PENDING 

71 strategy: SyncStrategy = SyncStrategy.LAST_WRITE_WINS 

72 

73 # 同步统计 

74 total_memories: int = 0 

75 synced_memories: int = 0 

76 failed_memories: int = 0 

77 conflict_memories: int = 0 

78 

79 # 操作列表 

80 operations: List[SyncOperation] = field(default_factory=list) 

81 

82 

83class ConflictResolver: 

84 """冲突解决器""" 

85 

86 @staticmethod 

87 def resolve_last_write_wins(local_memory: MemoryItem, remote_memory: MemoryItem) -> MemoryItem: 

88 """最后写入获胜策略""" 

89 if local_memory.updated_at > remote_memory.updated_at: 

90 return local_memory 

91 else: 

92 return remote_memory 

93 

94 @staticmethod 

95 def resolve_source_priority(local_memory: MemoryItem, remote_memory: MemoryItem, source_device_id: str) -> MemoryItem: 

96 """源设备优先策略""" 

97 if local_memory.device_id == source_device_id: 

98 return local_memory 

99 elif remote_memory.device_id == source_device_id: 

100 return remote_memory 

101 else: 

102 # 回退到最后写入获胜 

103 return ConflictResolver.resolve_last_write_wins(local_memory, remote_memory) 

104 

105 @staticmethod 

106 def auto_merge(local_memory: MemoryItem, remote_memory: MemoryItem) -> MemoryItem: 

107 """自动合并策略""" 

108 # 创建合并后的记忆 

109 merged = MemoryItem( 

110 id=local_memory.id, 

111 type=local_memory.type, 

112 content=local_memory.content.copy(), 

113 priority=max(local_memory.priority, remote_memory.priority), 

114 metadata=local_memory.metadata, 

115 device_id=local_memory.device_id, # 保持本地设备ID 

116 related_memories=list(set(local_memory.related_memories + remote_memory.related_memories)) 

117 ) 

118 

119 # 合并内容(简单合并策略) 

120 for key, value in remote_memory.content.items(): 

121 if key not in merged.content: 

122 merged.content[key] = value 

123 elif isinstance(value, dict) and isinstance(merged.content[key], dict): 

124 # 递归合并字典 

125 merged.content[key] = {**merged.content[key], **value} 

126 elif isinstance(value, list) and isinstance(merged.content[key], list): 

127 # 合并列表(去重) 

128 merged.content[key] = list(set(merged.content[key] + value)) 

129 

130 # 合并元数据 

131 merged.metadata.access_count = max( 

132 local_memory.metadata.access_count, 

133 remote_memory.metadata.access_count 

134 ) 

135 merged.metadata.last_accessed = max( 

136 local_memory.metadata.last_accessed, 

137 remote_memory.metadata.last_accessed 

138 ) 

139 

140 # 合并标签 

141 merged.metadata.tags = list(set(local_memory.metadata.tags + remote_memory.metadata.tags)) 

142 

143 # 更新最后访问时间 

144 merged.metadata.last_accessed = time.time() 

145 merged.updated_at = time.time() 

146 

147 return merged 

148 

149 @staticmethod 

150 def detect_conflict(local_memory: MemoryItem, remote_memory: MemoryItem) -> bool: 

151 """检测冲突""" 

152 if local_memory.type != remote_memory.type: 

153 return True 

154 

155 # 检查内容是否相同 

156 if local_memory.content != remote_memory.content: 

157 return True 

158 

159 # 检查优先级是否相同 

160 if local_memory.priority != remote_memory.priority: 

161 return True 

162 

163 return False 

164 

165 

166class SyncEngine: 

167 """同步引擎""" 

168 

169 def __init__(self, memory_system: MemorySystem, local_device_id: str): 

170 self.memory_system = memory_system 

171 self.local_device_id = local_device_id 

172 

173 # 同步会话管理 

174 self.sessions: Dict[str, SyncSession] = {} 

175 self.pending_operations: Dict[str, SyncOperation] = {} 

176 

177 # 同步配置 

178 self.sync_interval: int = 300 # 同步间隔(秒) 

179 self.max_retries: int = 3 

180 self.batch_size: int = 50 

181 

182 # 设备连接管理 

183 self.connected_devices: Dict[str, DeviceInfo] = {} 

184 

185 # 运行状态 

186 self.is_running = False 

187 self.sync_task: Optional[asyncio.Task] = None 

188 

189 async def start(self) -> None: 

190 """启动同步引擎""" 

191 if self.is_running: 

192 return 

193 

194 self.is_running = True 

195 

196 # 启动定期同步任务 

197 self.sync_task = asyncio.create_task(self._periodic_sync()) 

198 

199 async def stop(self) -> None: 

200 """停止同步引擎""" 

201 if not self.is_running: 

202 return 

203 

204 self.is_running = False 

205 

206 # 取消同步任务 

207 if self.sync_task: 

208 self.sync_task.cancel() 

209 try: 

210 await self.sync_task 

211 except asyncio.CancelledError: 

212 pass 

213 

214 # 等待所有同步会话完成 

215 await self._wait_for_sessions() 

216 

217 async def sync_with_device(self, device_id: str, strategy: SyncStrategy = SyncStrategy.LAST_WRITE_WINS) -> SyncSession: 

218 """ 

219 与指定设备同步 

220  

221 Args: 

222 device_id: 设备ID 

223 strategy: 同步策略 

224  

225 Returns: 

226 同步会话 

227 """ 

228 if device_id not in self.connected_devices: 

229 raise ValueError(f"设备未连接: {device_id}") 

230 

231 # 创建同步会话 

232 session_id = hashlib.md5(f"{self.local_device_id}_{device_id}_{time.time()}".encode()).hexdigest()[:12] 

233 session = SyncSession( 

234 id=session_id, 

235 source_device_id=self.local_device_id, 

236 target_device_id=device_id, 

237 strategy=strategy 

238 ) 

239 

240 self.sessions[session_id] = session 

241 

242 # 开始同步 

243 asyncio.create_task(self._execute_sync_session(session)) 

244 

245 return session 

246 

247 async def get_sync_status(self, session_id: str) -> Optional[SyncSession]: 

248 """获取同步状态""" 

249 return self.sessions.get(session_id) 

250 

251 async def list_sync_sessions(self, device_id: Optional[str] = None) -> List[SyncSession]: 

252 """列出同步会话""" 

253 if device_id: 

254 return [ 

255 session for session in self.sessions.values() 

256 if session.source_device_id == device_id or session.target_device_id == device_id 

257 ] 

258 

259 return list(self.sessions.values()) 

260 

261 async def add_device(self, device: DeviceInfo) -> None: 

262 """添加设备""" 

263 self.connected_devices[device.device_id] = device 

264 

265 async def remove_device(self, device_id: str) -> None: 

266 """移除设备""" 

267 if device_id in self.connected_devices: 

268 del self.connected_devices[device_id] 

269 

270 async def _execute_sync_session(self, session: SyncSession) -> None: 

271 """执行同步会话""" 

272 session.status = SyncStatus.SYNCING 

273 

274 try: 

275 # 获取需要同步的记忆 

276 local_memories = await self._get_memories_for_sync(session.target_device_id) 

277 session.total_memories = len(local_memories) 

278 

279 # 分批同步 

280 for i in range(0, len(local_memories), self.batch_size): 

281 batch = local_memories[i:i + self.batch_size] 

282 

283 # 同步批次 

284 await self._sync_batch(session, batch) 

285 

286 # 更新进度 

287 session.synced_memories = min(session.total_memories, i + self.batch_size) 

288 

289 # 短暂休眠,避免过载 

290 await asyncio.sleep(0.1) 

291 

292 # 标记完成 

293 session.status = SyncStatus.COMPLETED 

294 session.end_time = time.time() 

295 

296 except Exception as e: 

297 session.status = SyncStatus.FAILED 

298 session.end_time = time.time() 

299 print(f"同步会话失败: {e}") 

300 

301 async def _sync_batch(self, session: SyncSession, memories: List[MemoryItem]) -> None: 

302 """同步一批记忆""" 

303 for memory in memories: 

304 try: 

305 # 创建同步操作 

306 operation = SyncOperation( 

307 id=hashlib.md5(f"{memory.id}_{session.id}".encode()).hexdigest()[:8], 

308 memory_id=memory.id, 

309 device_id=session.target_device_id, 

310 direction=SyncDirection.BIDIRECTIONAL 

311 ) 

312 

313 # 执行同步 

314 success = await self._sync_single_memory(session, memory, operation) 

315 

316 if success: 

317 operation.status = SyncStatus.COMPLETED 

318 session.synced_memories += 1 

319 else: 

320 operation.status = SyncStatus.FAILED 

321 session.failed_memories += 1 

322 

323 # 添加到会话 

324 session.operations.append(operation) 

325 

326 except Exception as e: 

327 print(f"同步记忆失败 {memory.id}: {e}") 

328 session.failed_memories += 1 

329 

330 async def _sync_single_memory(self, session: SyncSession, memory: MemoryItem, operation: SyncOperation) -> bool: 

331 """同步单个记忆""" 

332 try: 

333 # 这里应该调用远程设备的API来获取远程记忆 

334 # 由于第二阶段只实现本地同步,这里模拟远程记忆 

335 remote_memory = await self._simulate_remote_memory(memory, session.target_device_id) 

336 

337 if not remote_memory: 

338 # 远程不存在,直接推送 

339 return await self._push_memory(memory, session.target_device_id) 

340 

341 # 检测冲突 

342 if ConflictResolver.detect_conflict(memory, remote_memory): 

343 session.conflict_memories += 1 

344 operation.status = SyncStatus.CONFLICT 

345 

346 # 根据策略解决冲突 

347 resolved_memory = await self._resolve_conflict( 

348 memory, remote_memory, session.strategy, session.source_device_id 

349 ) 

350 

351 # 更新本地和远程 

352 await self.memory_system.save(resolved_memory) 

353 await self._push_memory(resolved_memory, session.target_device_id) 

354 

355 operation.conflict_resolution = {"strategy": session.strategy.value} 

356 operation.merged_memory_id = resolved_memory.id 

357 

358 else: 

359 # 无冲突,更新较新的版本 

360 if memory.updated_at > remote_memory.updated_at: 

361 await self._push_memory(memory, session.target_device_id) 

362 else: 

363 await self.memory_system.save(remote_memory) 

364 

365 return True 

366 

367 except Exception as e: 

368 operation.error_message = str(e) 

369 operation.retry_count += 1 

370 

371 if operation.retry_count < self.max_retries: 

372 # 重试 

373 await asyncio.sleep(2 ** operation.retry_count) # 指数退避 

374 return await self._sync_single_memory(session, memory, operation) 

375 

376 return False 

377 

378 async def _push_memory(self, memory: MemoryItem, target_device_id: str) -> bool: 

379 """推送记忆到远程设备""" 

380 # 第二阶段:模拟推送成功 

381 # 实际实现中应该通过WebSocket或HTTP API发送到远程设备 

382 print(f"推送记忆 {memory.id} 到设备 {target_device_id}") 

383 return True 

384 

385 async def _pull_memory(self, memory_id: str, source_device_id: str) -> Optional[MemoryItem]: 

386 """从远程设备拉取记忆""" 

387 # 第二阶段:模拟拉取 

388 # 实际实现中应该通过WebSocket或HTTP API从远程设备获取 

389 print(f"从设备 {source_device_id} 拉取记忆 {memory_id}") 

390 

391 # 模拟返回一个记忆 

392 return None 

393 

394 async def _get_memories_for_sync(self, target_device_id: str) -> List[MemoryItem]: 

395 """获取需要同步的记忆""" 

396 # 查询所有需要同步到目标设备的记忆 

397 query = MemoryQuery( 

398 device_id=self.local_device_id, 

399 limit=1000 

400 ) 

401 

402 memories = await self.memory_system.query(query) 

403 

404 # 过滤:只同步最近修改过的记忆 

405 recent_memories = [ 

406 memory for memory in memories 

407 if memory.updated_at > time.time() - 86400 # 24小时内 

408 ] 

409 

410 return recent_memories 

411 

412 async def _simulate_remote_memory(self, memory: MemoryItem, device_id: str) -> Optional[MemoryItem]: 

413 """模拟远程记忆(用于测试)""" 

414 # 50%的概率返回模拟的远程记忆 

415 import random 

416 if random.random() > 0.5: 

417 return None 

418 

419 # 创建模拟的远程记忆 

420 remote_memory = MemoryItem( 

421 id=memory.id, 

422 type=memory.type, 

423 content=memory.content.copy(), 

424 priority=memory.priority, 

425 metadata=memory.metadata, 

426 device_id=device_id, 

427 related_memories=memory.related_memories.copy() 

428 ) 

429 

430 # 稍微修改内容以模拟冲突 

431 if isinstance(remote_memory.content, dict) and len(remote_memory.content) > 0: 

432 first_key = list(remote_memory.content.keys())[0] 

433 if isinstance(remote_memory.content[first_key], str): 

434 remote_memory.content[first_key] += " (远程修改)" 

435 

436 # 设置不同的更新时间(可能更旧或更新) 

437 remote_memory.updated_at = memory.updated_at + random.uniform(-3600, 3600) 

438 

439 return remote_memory 

440 

441 async def _resolve_conflict(self, local_memory: MemoryItem, remote_memory: MemoryItem, 

442 strategy: SyncStrategy, source_device_id: str) -> MemoryItem: 

443 """解决冲突""" 

444 if strategy == SyncStrategy.LAST_WRITE_WINS: 

445 return ConflictResolver.resolve_last_write_wins(local_memory, remote_memory) 

446 

447 elif strategy == SyncStrategy.SOURCE_PRIORITY: 

448 return ConflictResolver.resolve_source_priority(local_memory, remote_memory, source_device_id) 

449 

450 elif strategy == SyncStrategy.AUTO_MERGE: 

451 return ConflictResolver.auto_merge(local_memory, remote_memory) 

452 

453 else: 

454 # 默认使用最后写入获胜 

455 return ConflictResolver.resolve_last_write_wins(local_memory, remote_memory) 

456 

457 async def _periodic_sync(self) -> None: 

458 """定期同步""" 

459 while self.is_running: 

460 try: 

461 # 等待同步间隔 

462 await asyncio.sleep(self.sync_interval) 

463 

464 # 与所有连接的设备同步 

465 for device_id in list(self.connected_devices.keys()): 

466 try: 

467 await self.sync_with_device(device_id) 

468 except Exception as e: 

469 print(f"与设备 {device_id} 定期同步失败: {e}") 

470 

471 except asyncio.CancelledError: 

472 break 

473 except Exception as e: 

474 print(f"定期同步任务出错: {e}") 

475 await asyncio.sleep(60) # 出错后等待1分钟 

476 

477 async def _wait_for_sessions(self, timeout: float = 30.0) -> None: 

478 """等待所有同步会话完成""" 

479 start_time = time.time() 

480 

481 while self.sessions: 

482 # 检查超时 

483 if time.time() - start_time > timeout: 

484 print(f"等待同步会话超时,剩余 {len(self.sessions)} 个会话") 

485 break 

486 

487 # 检查是否有活跃的会话 

488 active_sessions = [ 

489 session for session in self.sessions.values() 

490 if session.status in [SyncStatus.PENDING, SyncStatus.SYNCING] 

491 ] 

492 

493 if not active_sessions: 

494 break 

495 

496 # 等待一段时间 

497 await asyncio.sleep(1) 

498 

499 def get_stats(self) -> Dict[str, Any]: 

500 """获取同步统计信息""" 

501 total_sessions = len(self.sessions) 

502 completed_sessions = len([s for s in self.sessions.values() if s.status == SyncStatus.COMPLETED]) 

503 failed_sessions = len([s for s in self.sessions.values() if s.status == SyncStatus.FAILED]) 

504 

505 total_operations = 0 

506 completed_operations = 0 

507 failed_operations = 0 

508 conflict_operations = 0 

509 

510 for session in self.sessions.values(): 

511 total_operations += len(session.operations) 

512 completed_operations += len([op for op in session.operations if op.status == SyncStatus.COMPLETED]) 

513 failed_operations += len([op for op in session.operations if op.status == SyncStatus.FAILED]) 

514 conflict_operations += len([op for op in session.operations if op.status == SyncStatus.CONFLICT]) 

515 

516 return { 

517 "connected_devices": len(self.connected_devices), 

518 "total_sessions": total_sessions, 

519 "completed_sessions": completed_sessions, 

520 "failed_sessions": failed_sessions, 

521 "total_operations": total_operations, 

522 "completed_operations": completed_operations, 

523 "failed_operations": failed_operations, 

524 "conflict_operations": conflict_operations 

525 }