from abc import ABC, abstractmethod from typing import Any, Dict, List, Optional, Union, Set, Tuple, override, cast from datetime import datetime, timedelta import threading import json import os import hashlib from context.schemas import ( SketchPadItem, SketchPadStatistics, SketchPadListItem, ) from redis import Redis class SketchPadBackend(ABC): """ SketchPad base interface. Defines the operations that any SketchPad backend implementation must support. Every sketch item has the following attributes: value: Any = Field(..., description="Stored value") timestamp: datetime = Field(default_factory=datetime.now, description="Creation time") summary: Optional[str] = Field(default=None, description="Content summary") expires_at: Optional[datetime] = Field(default=None, description="Expiration time") access_count: int = Field(default=0, description="Access count") last_accessed: Optional[datetime] = Field(default=None, description="Last access time") tags: Set[str] = Field(default_factory=set, description="Tag set") content_type: str = Field(default="text", description="Content type") content_hash: Optional[str] = Field(default=None, description="Content hash value") """ @abstractmethod def __init__( self, sketch_pad_id: str, file_path: Optional[str] = None, ): """ Initialize the SketchPad backend. Args: sketch_pad_id: Unique identifier for the SketchPad backend file_path: File path used for persisting data """ pass @abstractmethod async def set_item( self, key: str, value: Any, ttl: Optional[int] = None, summary: Optional[str] = None, tags: Optional[Set[str]] = None, ) -> str: """Set a key-value pair.""" pass @abstractmethod def get_item(self, key: str) -> Optional[SketchPadItem]: """Get complete item information.""" pass @abstractmethod def get_value(self, key: str) -> Any: """Get only the value.""" pass @abstractmethod def search_by_tags( self, tags: Set[str], match_all: bool = False ) -> List[Tuple[str, SketchPadItem]]: """Search by tags.""" pass @abstractmethod def search_by_content( self, query: str, limit: int = 5 ) -> List[Tuple[str, SketchPadItem]]: """Simple content-based search.""" pass @abstractmethod def delete(self, key: str) -> bool: """Delete a key-value pair.""" pass @abstractmethod def exists(self, key: str) -> bool: """Check whether a key exists.""" pass @abstractmethod def keys(self, pattern: Optional[str] = None) -> List[str]: """Get all key names.""" pass @abstractmethod def clear(self) -> None: """Clear all data.""" pass @abstractmethod def serialize(self) -> Dict[str, Any]: """Serialize to a dictionary for saving to file.""" pass @abstractmethod def deserialize(self, data: Dict[str, Any]) -> None: """Deserialize from a dictionary for loading from file.""" pass @abstractmethod def persist(self) -> None: """Persist data.""" pass @abstractmethod def restore(self) -> None: """Restore from persisted data.""" pass @abstractmethod def get_statistics(self) -> SketchPadStatistics: """Get statistics.""" pass @abstractmethod def list_items(self, include_value: bool = False) -> List[SketchPadListItem]: """List all items.""" pass class RedisFileSketchPadBackend(SketchPadBackend): """ RedisFileSketchPadBackend combines immediate Redis storage with file-system persistence for SketchPad backend implementation. Features: 1. Redis provides high-performance immediate access 2. The file system provides reliable persistence 3. Supports automatic synchronization and recovery 4. Uses Redis AOF + RDB mechanisms """ @override def __init__( self, sketch_pad_id: str, redis_host: str = "localhost", redis_port: int = 6379, redis_db: int = 0, file_path: Optional[str] = None, ): """ Initialize RedisFileSketchPadBackend. Args: sketch_pad_id: Unique SketchPad identifier redis_host: Redis host redis_port: Redis port redis_db: Redis database file_path: File path used for persisting data """ self.sketch_pad_id = sketch_pad_id self.file_path = file_path or f"sketch_pads/sketch_{sketch_pad_id}.json" self.redis_host = redis_host self.redis_port = redis_port self.redis_db = redis_db self.redis: Redis = Redis(host=self.redis_host, port=self.redis_port, db=self.redis_db) self._lock = threading.RLock() self._restore_from_storage() # ---- Redis typed helpers (to avoid Awaitable union types in stubs) ---- def _redis_get(self, key: str) -> Optional[bytes]: raw = cast(Any, self.redis.get(key)) return cast(Optional[bytes], raw) def _redis_keys(self, pattern: str) -> List[bytes]: raw = cast(Any, self.redis.keys(pattern)) return cast(List[bytes], raw) def _redis_smembers(self, key: str) -> Set[bytes]: raw = cast(Any, self.redis.smembers(key)) return cast(Set[bytes], raw) def _redis_delete(self, *keys: Union[str, bytes]) -> int: raw = cast(Any, self.redis.delete(*keys)) return cast(int, raw) def _redis_exists(self, key: str) -> int: raw = cast(Any, self.redis.exists(key)) return cast(int, raw) def _restore_from_storage(self) -> None: """Restore from storage.""" with self._lock: if self.file_path is None: return if os.path.exists(self.file_path): try: with open(self.file_path, "r", encoding="utf-8") as f: data = json.load(f) self.deserialize(data) except Exception as e: print(f"Warning: Failed to restore from file: {e}") def _get_redis_key(self, key: str) -> str: """Get the Redis key name.""" return f"sketch_pad:{self.sketch_pad_id}:{key}" def _get_content_hash(self, value: Any) -> str: """Compute the hash value of the content.""" content_str = json.dumps(value, sort_keys=True, ensure_ascii=False) return hashlib.md5(content_str.encode('utf-8')).hexdigest() @override async def set_item( self, key: str, value: Any, ttl: Optional[int] = None, summary: Optional[str] = None, tags: Optional[Set[str]] = None, ) -> str: """ Set a key-value pair. Args: key: Key value: Value ttl: Expiration time in seconds summary: Summary tags: Tags """ with self._lock: # Create SketchPadItem. item = SketchPadItem( value=value, timestamp=datetime.now(), summary=summary, tags=tags or set(), expires_at=datetime.now() + timedelta(seconds=ttl) if ttl else None, content_hash=self._get_content_hash(value), ) # Store in Redis. item_json = item.model_dump_json() redis_key = self._get_redis_key(key) self.redis.set(redis_key, item_json) # Set expiration time. if ttl: self.redis.expire(redis_key, ttl) # Update tag index. if tags: for tag in tags: tag_key = self._get_redis_key(f"tag:{tag}") self.redis.sadd(tag_key, key) return key @override def get_item(self, key: str) -> Optional[SketchPadItem]: """Get complete item information.""" with self._lock: item_json_opt = self._redis_get(self._get_redis_key(key)) if item_json_opt is None: return None try: item_bytes = cast(bytes, item_json_opt) item = SketchPadItem.model_validate_json(item_bytes) # Update access information. item.access_count += 1 item.last_accessed = datetime.now() # Update access information in Redis. item_json_str: str = item.model_dump_json() self.redis.set(self._get_redis_key(key), item_json_str) return item except Exception as e: print(f"Warning: Failed to deserialize item: {e}") return None @override def get_value(self, key: str) -> Any: """Get only the value.""" item = self.get_item(key) return item.value if item else None @override def search_by_tags( self, tags: Set[str], match_all: bool = False ) -> List[Tuple[str, SketchPadItem]]: """Search by tags.""" with self._lock: results: List[Tuple[str, SketchPadItem]] = [] if match_all: # Must match all tags. if not tags: return results # Get all keys for the first tag. first_tag = list(tags)[0] tag_key = self._get_redis_key(f"tag:{first_tag}") candidate_keys = self._redis_smembers(tag_key) # Check whether each candidate key contains all tags. for cand_key in candidate_keys: cand_key_str: str = cand_key.decode('utf-8') item = self.get_item(cand_key_str) if item and tags.issubset(item.tags): results.append((cand_key_str, item)) else: # Match any tag. for tag in tags: tag_key = self._get_redis_key(f"tag:{tag}") keys = self._redis_smembers(tag_key) for member_key in keys: member_key_str: str = member_key.decode('utf-8') item = self.get_item(member_key_str) if item and (member_key_str, item) not in results: results.append((member_key_str, item)) return results @override def search_by_content( self, query: str, limit: int = 5 ) -> List[Tuple[str, SketchPadItem]]: """Simple content-based search.""" with self._lock: results: List[Tuple[str, SketchPadItem]] = [] query_lower = query.lower() # Get all keys, but filter out tag index keys. pattern = self._get_redis_key("*") all_keys = self._redis_keys(pattern) for key in all_keys: redis_key_string: str = key.decode('utf-8') # Filter out tag index keys. if ":tag:" in redis_key_string: continue # Extract the original key name. original_key = redis_key_string.split(":", 2)[-1] item = self.get_item(original_key) if item: # Search value, summary, and tags. searchable_text = "" if isinstance(item.value, str): searchable_text += item.value + " " if item.summary: searchable_text += item.summary + " " if item.tags: searchable_text += " ".join(item.tags) + " " if query_lower in searchable_text.lower(): results.append((original_key, item)) if len(results) >= limit: break return results @override def delete(self, key: str) -> bool: """Delete a key-value pair.""" with self._lock: # Get the item to delete tag indexes. item = self.get_item(key) if item and item.tags: for tag in item.tags: tag_key = self._get_redis_key(f"tag:{tag}") self.redis.srem(tag_key, key) # Delete the primary key. redis_key = self._get_redis_key(key) result = self._redis_delete(redis_key) return result > 0 @override def exists(self, key: str) -> bool: """Check whether a key exists.""" with self._lock: exists_count = self._redis_exists(self._get_redis_key(key)) return exists_count > 0 @override def keys(self, pattern: Optional[str] = None) -> List[str]: """Get all key names.""" with self._lock: redis_pattern = self._get_redis_key(pattern or "*") keys = self._redis_keys(redis_pattern) # Extract the original key name. result: List[str] = [] for key in keys: redis_key_str_list: str = key.decode('utf-8') original_key = redis_key_str_list.split(":", 2)[-1] result.append(original_key) return result @override def clear(self) -> None: """Clear all data.""" with self._lock: # Get all keys. pattern = self._get_redis_key("*") keys = self._redis_keys(pattern) # Delete all keys. if keys: self._redis_delete(*keys) @override def serialize(self) -> Dict[str, Any]: """Serialize to a dictionary for saving to file.""" with self._lock: data: Dict[str, Any] = { "sketch_pad_id": self.sketch_pad_id, "items": {}, "serialization_timestamp": datetime.now().isoformat(), } # Serialize all items, but filter out tag index keys. pattern = self._get_redis_key("*") all_keys = self._redis_keys(pattern) for key in all_keys: redis_key_str_ser: str = key.decode('utf-8') # Filter out tag index keys. if ":tag:" in redis_key_str_ser: continue # Extract the original key name. original_key = redis_key_str_ser.split(":", 2)[-1] item = self.get_item(original_key) if item: data["items"][original_key] = item.model_dump() return data @override def deserialize(self, data: Dict[str, Any]) -> None: """Deserialize from a dictionary for loading from file.""" with self._lock: if "items" in data: for key, item_data in data["items"].items(): try: item = SketchPadItem(**item_data) item_json = item.model_dump_json() self.redis.set(self._get_redis_key(key), item_json) # Restore tag indexes. if item.tags: for tag in item.tags: tag_key = self._get_redis_key(f"tag:{tag}") self.redis.sadd(tag_key, key) except Exception as e: print(f"Warning: Failed to deserialize item {key}: {e}") @override def persist(self) -> None: """Persist data.""" try: # Ensure the directory exists. dir_path = os.path.dirname(self.file_path) if dir_path: os.makedirs(dir_path, exist_ok=True) # Serialize data. data = self.serialize() # Write to file. with open(self.file_path, "w", encoding="utf-8") as f: json.dump(data, f, ensure_ascii=False, indent=2, default=str) except Exception as e: print(f"Warning: Failed to persist sketch pad: {e}") @override def restore(self) -> None: """Restore from persisted data.""" if not os.path.exists(self.file_path): return try: with open(self.file_path, "r", encoding="utf-8") as f: data = json.load(f) self.deserialize(data) except Exception as e: print(f"Warning: Failed to restore sketch pad: {e}") def get_statistics(self) -> SketchPadStatistics: """Get statistics.""" with self._lock: # Get only actual data keys, excluding tag index keys. pattern = self._get_redis_key("*") all_keys = self._redis_keys(pattern) data_keys: List[str] = [] for key in all_keys: redis_key_str_stats: str = key.decode('utf-8') # Filter out tag index keys. if not redis_key_str_stats.endswith(":tag:") and ":tag:" not in redis_key_str_stats: original_key = redis_key_str_stats.split(":", 2)[-1] data_keys.append(original_key) total_items = len(data_keys) total_accesses = 0 items_with_summary = 0 popular_tags: Dict[str, int] = {} content_types: Dict[str, int] = {} for data_key in data_keys: item = self.get_item(data_key) if item: total_accesses += item.access_count if item.summary: items_with_summary += 1 # Count tags. for tag in item.tags: popular_tags[tag] = popular_tags.get(tag, 0) + 1 # Count content types. content_types[item.content_type] = content_types.get(item.content_type, 0) + 1 avg_access_per_item = total_accesses / total_items if total_items > 0 else 0 memory_usage_percent = (total_items / 1000) * 100 # Assume a maximum of 1000 items. return SketchPadStatistics( total_items=total_items, max_items=1000, items_with_summary=items_with_summary, total_accesses=total_accesses, popular_tags=popular_tags, content_types=content_types, avg_access_per_item=avg_access_per_item, memory_usage_percent=memory_usage_percent, ) def list_items(self, include_value: bool = False) -> List[SketchPadListItem]: """List all items.""" with self._lock: items: List[SketchPadListItem] = [] # Get all keys, but filter out tag index keys. pattern = self._get_redis_key("*") all_keys = self._redis_keys(pattern) for key in all_keys: redis_key_str_list_items: str = key.decode('utf-8') # Filter out tag index keys. if ":tag:" in redis_key_str_list_items: continue # Extract the original key name. original_key = redis_key_str_list_items.split(":", 2)[-1] item = self.get_item(original_key) if item: list_item = SketchPadListItem( key=original_key, summary=item.summary, timestamp=item.timestamp.isoformat(), tags=list(item.tags), content_type=item.content_type, access_count=item.access_count, content_hash=item.content_hash, value=item.value if include_value else None, ) items.append(list_item) return items