587 lines
20 KiB
Python
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
|
|
|