Files
2026-07-22 13:48:46 +08:00

587 lines
20 KiB
Python

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