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
« prev ^ index » next coverage.py v7.3.4, created at 2026-04-21 14:54 +0800
1"""
2记忆同步系统
3实现跨设备的记忆同步和冲突解决
4"""
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
15from .memory import MemorySystem, MemoryItem, MemoryQuery, MemoryType
16from .device import DeviceInfo
19class SyncStatus(str, Enum):
20 """同步状态"""
21 PENDING = "pending" # 等待同步
22 SYNCING = "syncing" # 同步中
23 COMPLETED = "completed" # 同步完成
24 FAILED = "failed" # 同步失败
25 CONFLICT = "conflict" # 冲突需要解决
28class SyncDirection(str, Enum):
29 """同步方向"""
30 PUSH = "push" # 推送(本地→远程)
31 PULL = "pull" # 拉取(远程→本地)
32 BIDIRECTIONAL = "bidirectional" # 双向同步
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" # 源设备优先
43@dataclass
44class SyncOperation:
45 """同步操作"""
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
56 # 冲突相关
57 conflict_resolution: Optional[Dict[str, Any]] = None
58 merged_memory_id: Optional[str] = None
61@dataclass
62class SyncSession:
63 """同步会话"""
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
73 # 同步统计
74 total_memories: int = 0
75 synced_memories: int = 0
76 failed_memories: int = 0
77 conflict_memories: int = 0
79 # 操作列表
80 operations: List[SyncOperation] = field(default_factory=list)
83class ConflictResolver:
84 """冲突解决器"""
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
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)
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 )
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))
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 )
140 # 合并标签
141 merged.metadata.tags = list(set(local_memory.metadata.tags + remote_memory.metadata.tags))
143 # 更新最后访问时间
144 merged.metadata.last_accessed = time.time()
145 merged.updated_at = time.time()
147 return merged
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
155 # 检查内容是否相同
156 if local_memory.content != remote_memory.content:
157 return True
159 # 检查优先级是否相同
160 if local_memory.priority != remote_memory.priority:
161 return True
163 return False
166class SyncEngine:
167 """同步引擎"""
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
173 # 同步会话管理
174 self.sessions: Dict[str, SyncSession] = {}
175 self.pending_operations: Dict[str, SyncOperation] = {}
177 # 同步配置
178 self.sync_interval: int = 300 # 同步间隔(秒)
179 self.max_retries: int = 3
180 self.batch_size: int = 50
182 # 设备连接管理
183 self.connected_devices: Dict[str, DeviceInfo] = {}
185 # 运行状态
186 self.is_running = False
187 self.sync_task: Optional[asyncio.Task] = None
189 async def start(self) -> None:
190 """启动同步引擎"""
191 if self.is_running:
192 return
194 self.is_running = True
196 # 启动定期同步任务
197 self.sync_task = asyncio.create_task(self._periodic_sync())
199 async def stop(self) -> None:
200 """停止同步引擎"""
201 if not self.is_running:
202 return
204 self.is_running = False
206 # 取消同步任务
207 if self.sync_task:
208 self.sync_task.cancel()
209 try:
210 await self.sync_task
211 except asyncio.CancelledError:
212 pass
214 # 等待所有同步会话完成
215 await self._wait_for_sessions()
217 async def sync_with_device(self, device_id: str, strategy: SyncStrategy = SyncStrategy.LAST_WRITE_WINS) -> SyncSession:
218 """
219 与指定设备同步
221 Args:
222 device_id: 设备ID
223 strategy: 同步策略
225 Returns:
226 同步会话
227 """
228 if device_id not in self.connected_devices:
229 raise ValueError(f"设备未连接: {device_id}")
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 )
240 self.sessions[session_id] = session
242 # 开始同步
243 asyncio.create_task(self._execute_sync_session(session))
245 return session
247 async def get_sync_status(self, session_id: str) -> Optional[SyncSession]:
248 """获取同步状态"""
249 return self.sessions.get(session_id)
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 ]
259 return list(self.sessions.values())
261 async def add_device(self, device: DeviceInfo) -> None:
262 """添加设备"""
263 self.connected_devices[device.device_id] = device
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]
270 async def _execute_sync_session(self, session: SyncSession) -> None:
271 """执行同步会话"""
272 session.status = SyncStatus.SYNCING
274 try:
275 # 获取需要同步的记忆
276 local_memories = await self._get_memories_for_sync(session.target_device_id)
277 session.total_memories = len(local_memories)
279 # 分批同步
280 for i in range(0, len(local_memories), self.batch_size):
281 batch = local_memories[i:i + self.batch_size]
283 # 同步批次
284 await self._sync_batch(session, batch)
286 # 更新进度
287 session.synced_memories = min(session.total_memories, i + self.batch_size)
289 # 短暂休眠,避免过载
290 await asyncio.sleep(0.1)
292 # 标记完成
293 session.status = SyncStatus.COMPLETED
294 session.end_time = time.time()
296 except Exception as e:
297 session.status = SyncStatus.FAILED
298 session.end_time = time.time()
299 print(f"同步会话失败: {e}")
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 )
313 # 执行同步
314 success = await self._sync_single_memory(session, memory, operation)
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
323 # 添加到会话
324 session.operations.append(operation)
326 except Exception as e:
327 print(f"同步记忆失败 {memory.id}: {e}")
328 session.failed_memories += 1
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)
337 if not remote_memory:
338 # 远程不存在,直接推送
339 return await self._push_memory(memory, session.target_device_id)
341 # 检测冲突
342 if ConflictResolver.detect_conflict(memory, remote_memory):
343 session.conflict_memories += 1
344 operation.status = SyncStatus.CONFLICT
346 # 根据策略解决冲突
347 resolved_memory = await self._resolve_conflict(
348 memory, remote_memory, session.strategy, session.source_device_id
349 )
351 # 更新本地和远程
352 await self.memory_system.save(resolved_memory)
353 await self._push_memory(resolved_memory, session.target_device_id)
355 operation.conflict_resolution = {"strategy": session.strategy.value}
356 operation.merged_memory_id = resolved_memory.id
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)
365 return True
367 except Exception as e:
368 operation.error_message = str(e)
369 operation.retry_count += 1
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)
376 return False
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
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}")
391 # 模拟返回一个记忆
392 return None
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 )
402 memories = await self.memory_system.query(query)
404 # 过滤:只同步最近修改过的记忆
405 recent_memories = [
406 memory for memory in memories
407 if memory.updated_at > time.time() - 86400 # 24小时内
408 ]
410 return recent_memories
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
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 )
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] += " (远程修改)"
436 # 设置不同的更新时间(可能更旧或更新)
437 remote_memory.updated_at = memory.updated_at + random.uniform(-3600, 3600)
439 return remote_memory
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)
447 elif strategy == SyncStrategy.SOURCE_PRIORITY:
448 return ConflictResolver.resolve_source_priority(local_memory, remote_memory, source_device_id)
450 elif strategy == SyncStrategy.AUTO_MERGE:
451 return ConflictResolver.auto_merge(local_memory, remote_memory)
453 else:
454 # 默认使用最后写入获胜
455 return ConflictResolver.resolve_last_write_wins(local_memory, remote_memory)
457 async def _periodic_sync(self) -> None:
458 """定期同步"""
459 while self.is_running:
460 try:
461 # 等待同步间隔
462 await asyncio.sleep(self.sync_interval)
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}")
471 except asyncio.CancelledError:
472 break
473 except Exception as e:
474 print(f"定期同步任务出错: {e}")
475 await asyncio.sleep(60) # 出错后等待1分钟
477 async def _wait_for_sessions(self, timeout: float = 30.0) -> None:
478 """等待所有同步会话完成"""
479 start_time = time.time()
481 while self.sessions:
482 # 检查超时
483 if time.time() - start_time > timeout:
484 print(f"等待同步会话超时,剩余 {len(self.sessions)} 个会话")
485 break
487 # 检查是否有活跃的会话
488 active_sessions = [
489 session for session in self.sessions.values()
490 if session.status in [SyncStatus.PENDING, SyncStatus.SYNCING]
491 ]
493 if not active_sessions:
494 break
496 # 等待一段时间
497 await asyncio.sleep(1)
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])
505 total_operations = 0
506 completed_operations = 0
507 failed_operations = 0
508 conflict_operations = 0
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])
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 }