first commit

This commit is contained in:
2026-07-22 13:48:46 +08:00
commit c87751c3dc
2820 changed files with 726976 additions and 0 deletions
+634
View File
@@ -0,0 +1,634 @@
from typing import Dict, List, Optional, Any, Union, override
from SimpleLLMFunc import async_llm_function, OpenAICompatible
import json
import os
import redis
import threading
from datetime import datetime
from abc import ABC, abstractmethod
from context.schemas import Message, ChatMessages
class ContextBackend(ABC):
"""
ContextBackend is the backend interface for context storage, defining the interfaces used by backend implementations.
Main responsibilities:
1. Define core interfaces for storage, querying, serialization, and persistence
2. Provide a unified abstraction layer that supports different storage implementations
3. Manage core data such as conversation history, summaries, and metadata
"""
@abstractmethod
def __init__(
self,
context_id: str,
llm_interface: Optional[OpenAICompatible] = None,
max_history_length: int = 5,
auto_summarize_trigger: int = 1000000,
file_path: Optional[str] = None,
):
"""
Initialize the context backend.
Args:
context_id: Unique context identifier
llm_interface: LLM interface used for history summarization
max_history_length: Maximum history record length
auto_summarize_trigger: Automatic summary trigger threshold
file_path: File persistence path (optional)
"""
pass
# ===== Core storage interface =====
@abstractmethod
async def store_message(self, message: Message) -> None:
"""Store one message."""
pass
@abstractmethod
def retrieve_messages(self, limit: Optional[int] = None) -> List[Message]:
"""Retrieve message history."""
pass
@abstractmethod
def retrieve_full_messages(self, limit: Optional[int] = None) -> List[Message]:
"""Retrieve the complete persisted message history."""
pass
@abstractmethod
def update_summary(self, summary: str) -> None:
"""Update the conversation summary."""
pass
@abstractmethod
def get_summary(self) -> Optional[str]:
"""Get the conversation summary."""
pass
@abstractmethod
def update_metadata(self, metadata: Dict[str, Any]) -> None:
"""Update metadata."""
pass
@abstractmethod
def get_metadata(self) -> Dict[str, Any]:
"""Get metadata."""
pass
# ===== Query interface =====
@abstractmethod
def search_messages(self, query: str, limit: int = 5) -> List[Message]:
"""Search messages."""
pass
@abstractmethod
def get_message_count(self) -> int:
"""Get the message count."""
pass
@abstractmethod
def get_total_message_count(self) -> int:
"""Get the message count in the complete persisted history."""
pass
@abstractmethod
def clear_messages(self, keep_summary: bool = True) -> None:
"""Clear message history."""
pass
# ===== Serialization interface =====
@abstractmethod
def serialize(self) -> Dict[str, Any]:
"""Serialize to a dictionary."""
pass
@abstractmethod
def deserialize(self, data: Dict[str, Any]) -> None:
"""Deserialize from a dictionary."""
pass
# ===== Persistence interface =====
@abstractmethod
async def persist(self) -> bool:
"""Persist to storage."""
pass
@abstractmethod
async def restore(self) -> bool:
"""Restore from storage."""
pass
# ===== Advanced feature interface =====
@abstractmethod
async def auto_summarize(self) -> str:
"""Automatically summarize history records."""
pass
@abstractmethod
def get_context_for_llm(self) -> str:
"""Get a context string suitable for the LLM."""
pass
class RedisFileContextBackend(ContextBackend):
"""
## RedisFileContextBackend combines immediate Redis storage with file-system persistence for context 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,
context_id: str,
llm_interface: Optional[OpenAICompatible] = None,
max_history_length: int = 5,
auto_summarize_trigger: int = 1000000,
redis_host: str = "localhost",
redis_port: int = 6379,
redis_db: int = 0,
file_path: Optional[str] = None,
):
"""
Initialize the Redis file backend.
Args:
context_id: Unique context identifier
llm_interface: LLM interface used for history summarization
max_history_length: Maximum history record length
auto_summarize_trigger: Automatic summary trigger threshold
redis_host: Redis host address
redis_port: Redis port
redis_db: Redis database number
file_path: File persistence path
"""
self.context_id = context_id
self.llm_interface = llm_interface
self.max_history_length = max_history_length
self.auto_summarize_trigger = auto_summarize_trigger
self.file_path = file_path or f"contexts/ctx_{context_id}.json"
# Redis connection.
self.redis_client = redis.Redis(
host=redis_host, port=redis_port, db=redis_db, decode_responses=True
)
# Thread lock.
self._lock = threading.RLock()
# Initialize the history summarization function.
self._summarize_func = None
if self.llm_interface:
self._summarize_func = async_llm_function(
llm_interface=self.llm_interface,
toolkit=[],
timeout=600,
)(self._summarize_history_impl)
# Initialize metadata.
self._init_metadata()
# Try to restore data from storage.
self._restore_from_storage()
def _init_metadata(self) -> None:
"""Initialize metadata."""
self._metadata = {
"context_id": self.context_id,
"session_id": self.context_id,
"start_time": datetime.now().isoformat(),
"last_activity": datetime.now().isoformat(),
"total_messages": 0,
"max_history_length": self.max_history_length,
"auto_summarize_trigger": self.auto_summarize_trigger,
}
def _normalize_metadata(self) -> None:
"""Normalize session metadata to ensure conversation and tracing session alignment."""
self._metadata["context_id"] = self.context_id
self._metadata["session_id"] = self.context_id
self._metadata.setdefault("start_time", datetime.now().isoformat())
self._metadata.setdefault("last_activity", datetime.now().isoformat())
self._metadata.setdefault("total_messages", 0)
self._metadata.setdefault("max_history_length", self.max_history_length)
self._metadata.setdefault("auto_summarize_trigger", self.auto_summarize_trigger)
def _generate_session_id(self) -> str:
"""Generate a session ID."""
return f"session_{datetime.now().strftime('%Y%m%d_%H%M%S')}"
def _get_redis_key(self, key: str) -> str:
"""Get the Redis key name."""
return f"context:{self.context_id}:{key}"
def _serialize_message(self, message: Message) -> str:
return message.model_dump_json()
def _append_to_working_messages(self, message: Message) -> None:
messages_key = self._get_redis_key("messages")
self.redis_client.lpush(messages_key, self._serialize_message(message))
def _append_to_full_messages(self, message: Message) -> None:
messages_key = self._get_redis_key("full_messages")
self.redis_client.rpush(messages_key, self._serialize_message(message))
def _replace_working_messages(self, messages: List[Message]) -> None:
messages_key = self._get_redis_key("messages")
self.redis_client.delete(messages_key)
for message in messages:
self._append_to_working_messages(message)
def _store_metadata_snapshot(self) -> None:
metadata_key = self._get_redis_key("metadata")
self.redis_client.set(metadata_key, json.dumps(self._metadata))
@override
async def store_message(self, message: Message) -> None:
"""
Store one message.
If the number of messages exceeds max_history_length, automatically trigger the summarization strategy and update the conversation records according to that strategy.
Args:
message: Message to store
Returns:
None
"""
with self._lock:
# Ensure the message has a timestamp.
if message.timestamp is None:
message.timestamp = datetime.now().isoformat()
# Store into working memory and full history.
self._append_to_working_messages(message)
self._append_to_full_messages(message)
# Automatic memory management.
await self._auto_memory_manage()
# Limit history length.
messages_key = self._get_redis_key("messages")
self.redis_client.ltrim(messages_key, 0, self.max_history_length - 1)
# Update metadata.
current_total = self._metadata.get("total_messages", 0)
if isinstance(current_total, (int, float)):
self._metadata["total_messages"] = int(current_total) + 1
else:
self._metadata["total_messages"] = 1
self._metadata["last_activity"] = datetime.now().isoformat()
self._store_metadata_snapshot()
# Automatic persistence.
await self.persist()
@override
def retrieve_messages(self, limit: Optional[int] = None) -> List[Message]:
"""Retrieve message history."""
with self._lock:
messages_key = self._get_redis_key("messages")
message_data_list = self.redis_client.lrange(messages_key, 0, -1)
messages = []
for message_data in message_data_list:
try:
message_dict = json.loads(message_data)
message = Message(**message_dict)
messages.append(message)
except Exception as e:
print(f"Warning: Failed to deserialize message: {e}")
# Sort by time, newest first.
messages.reverse()
if limit is not None:
messages = messages[-limit:]
return messages
@override
def retrieve_full_messages(self, limit: Optional[int] = None) -> List[Message]:
"""Retrieve the complete persisted message history."""
with self._lock:
messages_key = self._get_redis_key("full_messages")
message_data_list = self.redis_client.lrange(messages_key, 0, -1)
messages = []
for message_data in message_data_list:
try:
message_dict = json.loads(message_data)
message = Message(**message_dict)
messages.append(message)
except Exception as e:
print(f"Warning: Failed to deserialize full history message: {e}")
if not messages:
messages = self.retrieve_messages()
if limit is not None:
messages = messages[-limit:]
return messages
@override
def update_summary(self, summary: str) -> None:
"""Update the conversation summary."""
with self._lock:
summary_key = self._get_redis_key("summary")
self.redis_client.set(summary_key, summary)
@override
def get_summary(self) -> Optional[str]:
"""Get the conversation summary."""
with self._lock:
summary_key = self._get_redis_key("summary")
return self.redis_client.get(summary_key)
@override
def update_metadata(self, metadata: Dict[str, Any]) -> None:
"""Update metadata."""
with self._lock:
self._metadata.update(metadata)
self._store_metadata_snapshot()
@override
def get_metadata(self) -> Dict[str, Any]:
"""Get metadata."""
with self._lock:
return self._metadata.copy()
@override
def search_messages(self, query: str, limit: int = 5) -> List[Message]:
"""
Search messages.
Args:
query: Search keyword
limit: Search result count limit
Returns:
List[Message]: Search result list
"""
messages = self.retrieve_messages()
results = []
query_lower = query.lower()
for message in reversed(messages):
content = message.content
if isinstance(content, str) and query_lower in content.lower():
results.append(message)
if len(results) >= limit:
break
return list(reversed(results))
@override
def get_message_count(self) -> int:
"""Get the message count."""
with self._lock:
messages_key = self._get_redis_key("messages")
return self.redis_client.llen(messages_key)
@override
def get_total_message_count(self) -> int:
"""Get the message count in the complete persisted history."""
with self._lock:
messages_key = self._get_redis_key("full_messages")
total = self.redis_client.llen(messages_key)
if total == 0:
return self.get_message_count()
return total
@override
def clear_messages(self, keep_summary: bool = True) -> None:
"""Clear message history."""
with self._lock:
messages_key = self._get_redis_key("messages")
full_messages_key = self._get_redis_key("full_messages")
self.redis_client.delete(messages_key)
self.redis_client.delete(full_messages_key)
if not keep_summary:
summary_key = self._get_redis_key("summary")
self.redis_client.delete(summary_key)
self._metadata["total_messages"] = 0
self._metadata["last_activity"] = datetime.now().isoformat()
self._store_metadata_snapshot()
@override
def serialize(self) -> Dict[str, Any]:
"""Serialize to a dictionary."""
with self._lock:
return {
"context_id": self.context_id,
"metadata": self._metadata,
"messages": [msg.model_dump() for msg in self.retrieve_full_messages()],
"working_messages": [
msg.model_dump() for msg in self.retrieve_messages()
],
"summary": self.get_summary(),
"serialization_timestamp": datetime.now().isoformat(),
}
@override
def deserialize(self, data: Dict[str, Any]) -> None:
"""Deserialize from a dictionary."""
with self._lock:
# Restore metadata.
if "metadata" in data:
self._metadata.update(data["metadata"])
self._normalize_metadata()
# Restore messages.
full_history_payload = data.get("messages", [])
full_messages: List[Message] = []
for message_data in full_history_payload:
try:
full_messages.append(Message(**message_data))
except Exception as e:
print(f"Warning: Failed to deserialize full history message: {e}")
full_messages_key = self._get_redis_key("full_messages")
self.redis_client.delete(full_messages_key)
for message in full_messages:
self._append_to_full_messages(message)
working_payload = data.get("working_messages")
working_messages: List[Message] = []
if isinstance(working_payload, list):
for message_data in working_payload:
try:
working_messages.append(Message(**message_data))
except Exception as e:
print(f"Warning: Failed to deserialize working message: {e}")
elif full_messages:
working_messages = full_messages[-self.max_history_length :]
self._replace_working_messages(working_messages)
# Restore summary.
if "summary" in data and data["summary"]:
self.update_summary(data["summary"])
@override
async def persist(self) -> bool:
"""Persist to file."""
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)
return True
except Exception as e:
print(f"Warning: Failed to persist context: {e}")
return False
@override
async def restore(self) -> bool:
"""Restore from file."""
if not os.path.exists(self.file_path):
return False
try:
with open(self.file_path, "r", encoding="utf-8") as f:
data = json.load(f)
self.deserialize(data)
return True
except Exception as e:
print(f"Warning: Failed to restore context: {e}")
return False
def _restore_from_storage(self) -> None:
"""Restore data from storage."""
# Try to restore from Redis.
metadata_key = self._get_redis_key("metadata")
stored_metadata = self.redis_client.get(metadata_key)
if stored_metadata:
try:
self._metadata.update(json.loads(stored_metadata))
self._normalize_metadata()
except Exception as e:
print(f"Warning: Failed to restore metadata from Redis: {e}")
# Try to restore from file.
if os.path.exists(self.file_path):
import asyncio
asyncio.create_task(self.restore())
async def _auto_memory_manage(self) -> None:
"""Automatic memory management."""
if (
self.get_message_count() > self.auto_summarize_trigger
and self.llm_interface
):
# Create summary.
summary = await self.auto_summarize()
# Save summary.
current_summary = self.get_summary()
if current_summary:
self.update_summary(f"{current_summary}\n\n{summary}")
else:
self.update_summary(summary)
# Keep the most recent message.
messages = self.retrieve_messages()
if messages:
self._replace_working_messages([messages[-1]])
@override
async def auto_summarize(self) -> str:
"""Automatically summarize history records."""
if self._summarize_func:
messages = self.retrieve_messages()
return await self._summarize_func(messages)
else:
count = self.get_message_count()
return f"The conversation contains {count} messages."
@override
def get_context_for_llm(self) -> str:
"""Get a context string suitable for the LLM."""
context_parts = []
# Add summary.
summary = self.get_summary()
if summary:
context_parts.append(f"Conversation summary:\n{summary}\n")
# Add recent history records.
messages = self.retrieve_messages()
if messages:
context_parts.append("Recent conversation history:")
for message in messages:
role = message.role
content = message.content
if isinstance(content, str):
context_parts.append(f"{role}: {content}")
return "\n".join(context_parts)
@staticmethod
async def _summarize_history_impl(messages: List[Message]) -> str: # type: ignore
"""
Please extract and summarize key information from the following conversation history. Requirements:
1. Distill the user's core intent and clearly describe it under the [User Intent] field.
2. Extract all key parameters, variable names, keys, file names, and similar information that appeared, and list them under the [Key Information] field. Use one item per line and indicate the type, such as file, key, parameter, and so on.
3. Preserve important operations, decisions, or changes involved in the conversation, and concisely summarize them under the [Conversation Highlights] field.
4. Output all fields strictly in the following format:
[User Intent]
... (briefly describe the user's main requirements and goals)
[Key Information]
- Type: Name
- Type: Name
...
[Conversation Highlights]
- Highlight 1
- Highlight 2
[Files Operated On]
- File 1
- File 2
- File 3
[Next-Step Plan]
- Plan 1
- Plan 2
- Plan 3
[Summary]
- Summary 1
- Summary 2
...
Ensure the summary is accurate and clearly structured, making it easy for later retrieval and context recovery.
Args:
messages: Message list
Returns:
str: Summarized conversation history
"""
@@ -0,0 +1,478 @@
import os
import json
import uuid
import threading
from typing import Dict, Optional, List, Type, Any, Literal, cast
from datetime import datetime
import redis
from SimpleLLMFunc import OpenAICompatible
from context.schemas import Message
from context.context import ContextBackend, RedisFileContextBackend
from config.config import get_config
from SimpleLLMFunc.logger import push_warning, app_log
class ContextManager:
"""
General-purpose context manager that supports different backend implementations.
Main responsibilities:
1. Manage creation and lifecycle of ContextBackend instances
2. Provide advanced convenience interfaces
3. Handle batch operations and cleanup tasks
4. Support pluggable backend implementations
"""
_instance = None
_lock: threading.Lock = threading.Lock()
def __new__(cls, backend_class: Type[ContextBackend] = RedisFileContextBackend):
"""Singleton pattern implementation."""
if cls._instance is None:
with cls._lock:
if cls._instance is None:
cls._instance = super(ContextManager, cls).__new__(cls)
cls._instance.backend_class = backend_class
return cls._instance
def __init__(self, backend_class: Type[ContextBackend]):
"""
Initialize the context manager.
Args:
backend_class: Backend implementation class, defaulting to RedisFileBackend
"""
# Prevent duplicate initialization.
if hasattr(self, "_initialized"):
return
self.backend_class = backend_class
self.config = get_config()
self.context_dir = self.config.CONTEXT_DIR
self._active_contexts: Dict[str, ContextBackend] = {}
# Ensure the directory exists.
os.makedirs(self.context_dir, exist_ok=True)
self._initialized = True
def _redis_client(self) -> redis.Redis:
return redis.Redis(
host=self.config.REDIS_HOST,
port=int(self.config.REDIS_PORT),
db=int(self.config.REDIS_DB),
decode_responses=True,
)
def _list_context_ids_from_redis(self) -> set[str]:
context_ids: set[str] = set()
try:
client = self._redis_client()
raw_keys = cast(Any, client.keys("context:*:*"))
for key in cast(List[str], raw_keys):
parts = key.split(":", 2)
if len(parts) >= 3 and parts[0] == "context" and parts[1]:
context_ids.add(parts[1])
except Exception as e:
print(f"Warning: Failed to list context ids from Redis: {e}")
return context_ids
def _delete_context_redis_keys(self, context_id: str) -> bool:
try:
client = self._redis_client()
raw_keys = cast(Any, client.keys(f"context:{context_id}:*"))
keys = cast(List[str], raw_keys)
if not keys:
return False
deleted = cast(Any, client.delete(*keys))
return int(deleted) > 0
except Exception as e:
print(f"Warning: Failed to delete Redis context keys for {context_id}: {e}")
return False
def create_context(
self,
context_id: Optional[str] = None,
llm_interface: Optional[
OpenAICompatible
] = get_config().CONTEXT_SUMMARY_INTERFACE,
max_history_length: int = get_config().CONTEXT_MAX_HISTORY_LENGTH,
auto_summarize_trigger: int = get_config().CONTEXT_AUTO_SUMMARIZE_TRIGGER,
**backend_kwargs,
) -> ContextBackend:
"""
Create a new context object.
Args:
context_id: Context ID; generated automatically if None
llm_interface: LLM interface
max_history_length: Maximum history length
auto_summarize_trigger: Automatic summary trigger threshold
**backend_kwargs: Extra parameters passed to the backend
Returns:
ContextBackend: Created context object
"""
with self._lock:
if context_id is None:
context_id = str(uuid.uuid4())
# Check whether it already exists.
if context_id in self._active_contexts:
app_log(
f"Context {context_id} already exists, and is in active contexts. Returning the existing context."
)
return self._active_contexts[context_id]
# Generate the file path if the backend needs one.
if "file_path" not in backend_kwargs:
context_file = os.path.join(self.context_dir, f"ctx_{context_id}.json")
backend_kwargs["file_path"] = context_file
push_warning(f"Context file path: {context_file}")
# Create the context object.
context = self.backend_class(
context_id=context_id,
llm_interface=llm_interface,
max_history_length=max_history_length,
auto_summarize_trigger=auto_summarize_trigger,
**backend_kwargs,
)
# Add it to the active context list.
self._active_contexts[context_id] = context
return context
def get_context(self, context_id: str) -> Optional[ContextBackend]:
"""
Get the context object with the specified ID.
Args:
context_id: Context ID
Returns:
ContextBackend: Context object, or None if it does not exist
"""
with self._lock:
# First check active contexts.
if context_id in self._active_contexts:
return self._active_contexts[context_id]
# Try to load from file if the backend supports it.
context_file = os.path.join(self.context_dir, f"ctx_{context_id}.json")
if os.path.exists(context_file):
try:
context = self.backend_class(
context_id=context_id,
llm_interface=self.config.CONTEXT_SUMMARY_INTERFACE, # Can be configured later.
max_history_length=self.config.CONTEXT_MAX_HISTORY_LENGTH,
auto_summarize_trigger=self.config.CONTEXT_AUTO_SUMMARIZE_TRIGGER,
file_path=context_file,
)
self._active_contexts[context_id] = context
return context
except Exception as e:
print(f"Warning: Failed to load context {context_id}: {e}")
return None
def delete_context(self, context_id: str) -> bool:
"""
Delete the context object with the specified ID.
Args:
context_id: Context ID
Returns:
bool: Whether deletion succeeded
"""
with self._lock:
success = False
# Remove it from active contexts.
if context_id in self._active_contexts:
del self._active_contexts[context_id]
success = True
# Delete context keys from Redis.
if self._delete_context_redis_keys(context_id):
success = True
# Delete the file if it exists.
context_file = os.path.join(self.context_dir, f"ctx_{context_id}.json")
if os.path.exists(context_file):
try:
os.remove(context_file)
success = True
except Exception as e:
print(f"Warning: Failed to delete context file {context_file}: {e}")
return success
def list_context_ids(self) -> List[str]:
"""List all known context IDs, including Redis and the file system."""
context_ids = set(self._active_contexts.keys())
context_ids.update(self._list_context_ids_from_redis())
try:
for filename in os.listdir(self.context_dir):
if filename.startswith("ctx_") and filename.endswith(".json"):
context_ids.add(filename[4:-5])
except Exception as e:
print(f"Warning: Failed to scan context dir for ids: {e}")
return sorted(context_ids)
def list_contexts(self) -> List[Dict[str, Any]]:
"""
List all available contexts.
Returns:
List[Dict]: Context information list
"""
contexts = []
# Scan context files in the file system.
try:
for filename in os.listdir(self.context_dir):
if filename.startswith("ctx_") and filename.endswith(".json"):
context_id = filename[4:-5] # Remove the "ctx_" prefix and ".json" suffix.
context_info = {
"context_id": context_id,
"file_path": os.path.join(self.context_dir, filename),
"is_active": context_id in self._active_contexts,
}
# Try to read basic information.
try:
file_path = context_info["file_path"]
if isinstance(file_path, str):
with open(file_path, "r", encoding="utf-8") as f:
data = json.load(f)
metadata = data.get("metadata", {})
context_info.update(
{
"start_time": metadata.get("start_time"),
"last_activity": metadata.get("last_activity"),
"total_messages": metadata.get(
"total_messages", 0
),
}
)
except Exception:
pass # Ignore read errors.
contexts.append(context_info)
except Exception as e:
print(f"Warning: Failed to list contexts: {e}")
return contexts
async def save_context(self, context_id: str) -> bool:
"""
Manually save the specified context to file.
Args:
context_id: Context ID
Returns:
bool: Whether saving succeeded
"""
with self._lock:
context = self._active_contexts.get(context_id)
if context is None:
return False
try:
return await context.persist()
except Exception as e:
print(f"Warning: Failed to save context {context_id}: {e}")
return False
async def save_all_contexts(self) -> int:
"""
Save all active contexts to files.
Returns:
int: Number of contexts successfully saved
"""
saved_count = 0
with self._lock:
context_ids = list(self._active_contexts.keys())
for context_id in context_ids:
if await self.save_context(context_id):
saved_count += 1
return saved_count
async def cleanup_inactive_contexts(self, max_inactive_time: int = 3600) -> int:
"""
Clean up contexts that have been inactive for a long time.
Args:
max_inactive_time: Maximum inactive time in seconds
Returns:
int: Number of cleaned contexts
"""
cleaned_count = 0
current_time = datetime.now()
contexts_to_persist: list[tuple[str, ContextBackend]] = []
with self._lock:
contexts_to_remove = []
for context_id, context in self._active_contexts.items():
try:
metadata = context.get_metadata()
last_activity_str = metadata.get("last_activity")
if last_activity_str:
last_activity = datetime.fromisoformat(last_activity_str)
inactive_time = (current_time - last_activity).total_seconds()
if inactive_time > max_inactive_time:
# Persist the context before removing it.
contexts_to_persist.append((context_id, context))
contexts_to_remove.append(context_id)
cleaned_count += 1
except Exception as e:
print(
f"Warning: Error checking activity for context {context_id}: {e}"
)
for context_id, context in contexts_to_persist:
try:
await context.persist()
except Exception as e:
print(f"Warning: Failed to persist inactive context {context_id}: {e}")
with self._lock:
for context_id in contexts_to_remove:
del self._active_contexts[context_id]
return cleaned_count
# ===== Convenience interfaces =====
async def add_message(
self,
context_id: str,
message: Message,
) -> bool:
"""
Convenience method for adding a message.
Args:
context_id: Context ID
role: Message role
content: Message content
**message_kwargs: Other message parameters
Returns:
bool: Whether adding succeeded
"""
backend = self.get_context(context_id)
if not backend:
return False
try:
await backend.store_message(message)
return True
except Exception as e:
print(f"Warning: Failed to add message: {e}")
return False
def get_history(self, context_id: str, limit: Optional[int] = None) -> List:
"""
Get conversation history.
Args:
context_id: Context ID
limit: Limit on the number of returned messages
Returns:
List: Message history
"""
backend = self.get_context(context_id)
if not backend:
return []
return backend.retrieve_messages(limit)
async def summarize_context(self, context_id: str) -> Optional[str]:
"""
Summarize the context.
Args:
context_id: Context ID
Returns:
Optional[str]: Summary content
"""
backend = self.get_context(context_id)
if not backend:
return None
return await backend.auto_summarize()
# Global instance.
_global_context_manager: Optional[ContextManager] = None
def get_context_manager() -> ContextManager:
"""Get the global ContextManager instance."""
global _global_context_manager
if _global_context_manager is None:
# Get Redis configuration from config.
config = get_config()
# redis config
redis_host = config.REDIS_HOST
redis_port = int(config.REDIS_PORT)
redis_db = int(config.REDIS_DB)
# Create a custom backend class with preconfigured Redis parameters.
class ConfiguredRedisFileBackend(RedisFileContextBackend):
def __init__(
self,
context_id: str,
llm_interface: Optional[
OpenAICompatible
] = config.CONTEXT_SUMMARY_INTERFACE,
max_history_length: int = config.CONTEXT_MAX_HISTORY_LENGTH,
auto_summarize_trigger: int = config.CONTEXT_AUTO_SUMMARIZE_TRIGGER,
redis_host: str = redis_host,
redis_port: int = redis_port,
redis_db: int = redis_db,
file_path: str = "",
):
super().__init__(
context_id=context_id,
llm_interface=llm_interface,
max_history_length=max_history_length,
auto_summarize_trigger=auto_summarize_trigger,
redis_host=redis_host,
redis_port=redis_port,
redis_db=redis_db,
file_path=file_path,
)
self.file_path = file_path
self.context_id = context_id
self.llm_interface = llm_interface
self.max_history_length = max_history_length
self.auto_summarize_trigger = auto_summarize_trigger
self.redis_host = redis_host
self.redis_port = redis_port
self.redis_db = redis_db
# Create ContextManager using the configured backend class.
_global_context_manager = ContextManager(
backend_class=ConfiguredRedisFileBackend
)
return _global_context_manager
@@ -0,0 +1,456 @@
from __future__ import annotations
import uuid
import threading
import os
from typing import Dict, Optional, List, Any
from datetime import datetime
from dataclasses import dataclass
from SimpleLLMFunc import OpenAICompatible
from SimpleLLMFunc.logger import push_warning, push_error, app_log
from context.context_manager import get_context_manager, ContextManager
from context.sketch_manager import get_sketch_manager, SketchManager
from context.sketch_pad import SketchPadBackend
from context.context import ContextBackend
from config.config import get_config
# Global current conversation context variable.
_current_conversation: Optional[Conversation] = None
_conversation_context_lock = threading.RLock()
def get_current_conversation() -> Optional[Conversation]:
"""Get the Conversation in the current context."""
global _current_conversation
with _conversation_context_lock:
return _current_conversation
def get_current_context() -> Optional[ContextBackend]:
"""Get the Context in the current context."""
conversation = get_current_conversation()
return conversation.context if conversation else None
def get_current_sketch_pad() -> Optional[SketchPadBackend]:
"""Get the SketchPad in the current context."""
conversation = get_current_conversation()
return conversation.sketch_pad if conversation else None
@dataclass
class Conversation:
"""
Conversation data class representing a complete conversation session.
Contains a unique UUID, associated Context, and SketchPad.
Supports use as a context manager.
"""
uuid: str
context: ContextBackend
sketch_pad: SketchPadBackend
created_at: datetime
last_accessed: datetime
def update_access_time(self):
"""Update the last access time."""
self.last_accessed = datetime.now()
def __enter__(self):
"""Enter the context manager."""
global _current_conversation
with _conversation_context_lock:
if _current_conversation is not None:
raise RuntimeError("Cannot nest conversation contexts")
_current_conversation = self
self.update_access_time()
return self
def __exit__(self, exc_type, exc_val, exc_tb):
"""Exit the context manager."""
global _current_conversation
with _conversation_context_lock:
_current_conversation = None
return False
class ConversationManager:
"""
Conversation manager responsible for creating, managing, and coordinating Conversation lifecycles.
Each Conversation contains one Context and one SketchPad, and they share the same UUID.
ConversationManager is a global singleton that uses ContextManager and SketchManager
to manage the underlying Context and SketchPad objects.
"""
_instance = None
_lock = threading.Lock()
def __new__(cls):
"""Singleton pattern implementation."""
if cls._instance is None:
with cls._lock:
if cls._instance is None:
cls._instance = super(ConversationManager, cls).__new__(cls)
return cls._instance
def __init__(self):
"""Initialize ConversationManager."""
# Prevent duplicate initialization.
if hasattr(self, "_initialized"):
return
self.config = get_config()
self.context_manager: ContextManager = get_context_manager()
self.sketch_manager: SketchManager = get_sketch_manager()
self._active_conversations: Dict[str, Conversation] = {}
self._lock = threading.RLock()
# Create the conversations directory.
self.conversations_dir = os.path.join(
os.path.dirname(self.config.CONTEXT_DIR), "conversations"
)
os.makedirs(self.conversations_dir, exist_ok=True)
self._initialized = True
def create_conversation(
self,
conversation_id: Optional[str] = None,
llm_interface: Optional[OpenAICompatible] = None,
max_history_length: int = 5,
) -> Conversation:
"""
Create a new Conversation.
Args:
conversation_id: Conversation UUID; generated automatically if None
llm_interface: LLM interface used for Context
max_history_length: Maximum Context history length
Returns:
Conversation: Created Conversation object
"""
with self._lock:
if conversation_id is None:
conversation_id = str(uuid.uuid4())
# Check whether it already exists.
if conversation_id in self._active_conversations:
conversation = self._active_conversations[conversation_id]
conversation.update_access_time()
return conversation
# Create Context with the ctx prefix.
context = self.context_manager.create_context(
context_id=conversation_id,
llm_interface=llm_interface,
max_history_length=max_history_length,
)
# Create SketchPad with the skt prefix.
sketch_pad = self.sketch_manager.create_sketch_pad(
sketch_id=conversation_id
)
# Create the Conversation object.
now = datetime.now()
conversation = Conversation(
uuid=conversation_id,
context=context,
sketch_pad=sketch_pad,
created_at=now,
last_accessed=now,
)
# Add it to the active Conversation list.
self._active_conversations[conversation_id] = conversation
# Immediately persist context and sketch_pad to the file system.
try:
import asyncio
# Create an event loop to run the asynchronous task.
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
loop.run_until_complete(context.persist())
# Synchronously call sketch_pad.persist().
sketch_pad.persist()
finally:
loop.close()
app_log(f"Conversation {conversation_id} was successfully persisted to the file system")
except Exception as e:
push_warning(f"Failed to persist conversation {conversation_id}: {e}")
# Create a persistence marker file.
self._create_conversation_marker(conversation_id)
return conversation
def get_conversation(self, conversation_id: str) -> Optional[Conversation]:
"""
Get the Conversation with the specified ID.
Args:
conversation_id: Conversation UUID
Returns:
Conversation: Conversation object, or None if it does not exist
"""
with self._lock:
# First check active Conversations.
if conversation_id in self._active_conversations:
conversation = self._active_conversations[conversation_id]
conversation.update_access_time()
return conversation
# Try to rebuild from the file system.
context = self.context_manager.get_context(conversation_id)
sketch_pad = self.sketch_manager.get_sketch_pad(conversation_id)
if context is not None and sketch_pad is not None:
# Rebuild the Conversation object.
now = datetime.now()
conversation = Conversation(
uuid=conversation_id,
context=context,
sketch_pad=sketch_pad,
created_at=now, # Use the current time as the rebuild time.
last_accessed=now,
)
self._active_conversations[conversation_id] = conversation
return conversation
return None
def delete_conversation(self, conversation_id: str) -> bool:
"""
Delete the Conversation with the specified ID.
Args:
conversation_id: Conversation UUID
Returns:
bool: Whether deletion succeeded
"""
with self._lock:
success = False
# Remove it from active Conversations.
if conversation_id in self._active_conversations:
del self._active_conversations[conversation_id]
success = True
# Delete the underlying Context and SketchPad.
context_deleted = self.context_manager.delete_context(conversation_id)
sketch_deleted = self.sketch_manager.delete_sketch_pad(conversation_id)
# Delete the marker file.
marker_file = os.path.join(
self.conversations_dir, f"conv_{conversation_id}.marker"
)
if os.path.exists(marker_file):
try:
os.remove(marker_file)
except Exception as e:
push_warning(f"Failed to delete marker file {marker_file}: {e}")
return success or context_deleted or sketch_deleted
def _discover_conversation_ids(self) -> List[str]:
"""Discover all known conversation ids across memory, files, and Redis-backed stores."""
conversation_ids = set(self._active_conversations.keys())
try:
for filename in os.listdir(self.conversations_dir):
if filename.startswith("conv_") and filename.endswith(".marker"):
conversation_ids.add(filename[5:-7])
except Exception as e:
push_warning(f"Failed to scan conversation markers: {e}")
try:
conversation_ids.update(self.context_manager.list_context_ids())
except Exception as e:
push_warning(f"Failed to collect context ids: {e}")
try:
conversation_ids.update(self.sketch_manager.list_sketch_ids())
except Exception as e:
push_warning(f"Failed to collect sketch ids: {e}")
return sorted(conversation_ids)
def delete_all_conversations(self) -> List[str]:
"""Delete all known conversations from memory, files, and Redis-backed stores."""
deleted_ids: List[str] = []
for conversation_id in self._discover_conversation_ids():
if self.delete_conversation(conversation_id):
deleted_ids.append(conversation_id)
return deleted_ids
def list_conversations(self) -> List[Dict[str, Any]]:
"""
List all available Conversations.
Returns:
List[Dict]: Conversation information list
"""
conversations = []
try:
for conversation_id in self._discover_conversation_ids():
marker_file = os.path.join(
self.conversations_dir, f"conv_{conversation_id}.marker"
)
conversation_info = {
"conversation_id": conversation_id,
"marker_file": marker_file if os.path.exists(marker_file) else None,
"is_active": conversation_id in self._active_conversations,
}
context = self.context_manager.get_context(conversation_id)
sketch_pad = self.sketch_manager.get_sketch_pad(conversation_id)
if context:
metadata = context.get_metadata()
conversation_info.update(
{
"context_start_time": metadata.get("start_time"),
"context_last_activity": metadata.get("last_activity"),
"context_total_messages": context.get_total_message_count(),
"context_has_summary": bool(context.get_summary()),
}
)
if sketch_pad:
stats = sketch_pad.get_statistics()
conversation_info.update(
{
"sketch_total_items": stats.total_items,
"sketch_max_items": stats.max_items,
"sketch_memory_usage": stats.memory_usage_percent,
}
)
conversations.append(conversation_info)
except Exception as e:
push_warning(f"Failed to list conversations: {e}")
return conversations
async def save_conversation(self, conversation_id: str) -> bool:
"""
Manually save the specified Conversation to files.
Args:
conversation_id: Conversation UUID
Returns:
bool: Whether saving succeeded
"""
with self._lock:
if conversation_id in self._active_conversations:
try:
conversation = self._active_conversations[conversation_id]
# Save Context.
context_saved = await conversation.context.persist()
# Save SketchPad.
conversation.sketch_pad.persist()
sketch_saved = True
return bool(context_saved and sketch_saved)
except Exception as e:
push_warning(f"Failed to save conversation {conversation_id}: {e}")
return False
async def save_all_conversations(self) -> int:
"""
Save all active Conversations to files.
Returns:
int: Number of Conversations successfully saved
"""
saved_count = 0
with self._lock:
for conversation_id in list(self._active_conversations.keys()):
if await self.save_conversation(conversation_id):
saved_count += 1
return saved_count
async def cleanup_inactive_conversations(
self, max_inactive_time: int = 3600
) -> int:
"""
Clean up Conversations that have been inactive for a long time.
Args:
max_inactive_time: Maximum inactive time in seconds
Returns:
int: Number of cleaned Conversations
"""
cleaned_count = 0
current_time = datetime.now()
with self._lock:
conversations_to_remove = []
for conversation_id, conversation in self._active_conversations.items():
try:
inactive_time = (
current_time - conversation.last_accessed
).total_seconds()
if inactive_time > max_inactive_time:
# Save the Conversation before removing it.
await self.save_conversation(conversation_id)
conversations_to_remove.append(conversation_id)
cleaned_count += 1
except Exception as e:
push_warning(
f"Error checking activity for conversation {conversation_id}: {e}"
)
# Remove inactive Conversations.
for conversation_id in conversations_to_remove:
del self._active_conversations[conversation_id]
return cleaned_count
def _create_conversation_marker(self, conversation_id: str) -> None:
"""
Create a Conversation marker file.
Args:
conversation_id: Conversation UUID
"""
try:
marker_file = os.path.join(
self.conversations_dir, f"conv_{conversation_id}.marker"
)
with open(marker_file, "w") as f:
f.write(
f"Conversation {conversation_id} created at {datetime.now().isoformat()}"
)
except Exception as e:
push_warning(
f"Failed to create marker file for conversation {conversation_id}: {e}"
)
# Global instance.
_global_conversation_manager: Optional[ConversationManager] = None
def get_conversation_manager() -> ConversationManager:
"""Get the global ConversationManager instance."""
global _global_conversation_manager
if _global_conversation_manager is None:
_global_conversation_manager = ConversationManager()
return _global_conversation_manager
+275
View File
@@ -0,0 +1,275 @@
from typing import Literal, List, Optional, Union, Any, Dict, Set
from pydantic import BaseModel, Field, model_validator, field_validator, RootModel
from datetime import datetime
import hashlib
def _content_item_type(value: Any) -> Optional[str]:
if isinstance(value, dict):
item_type = value.get("type")
return item_type if isinstance(item_type, str) else None
item_type = getattr(value, "type", None)
return item_type if isinstance(item_type, str) else None
def _content_item_text(value: Any) -> Optional[str]:
if isinstance(value, dict):
text = value.get("text")
return text if isinstance(text, str) else None
text = getattr(value, "text", None)
return text if isinstance(text, str) else None
def _content_item_image_payload(value: Any) -> Any:
if isinstance(value, dict):
return value.get("image_url")
return getattr(value, "image_url", None)
def _image_payload_string_field(payload: Any, field_name: str) -> Optional[str]:
if isinstance(payload, dict):
value = payload.get(field_name)
return value if isinstance(value, str) else None
value = getattr(payload, field_name, None)
return value if isinstance(value, str) else None
def _normalize_multimodal_content_item(value: Any) -> Dict[str, Any]:
if hasattr(value, "model_dump"):
value = value.model_dump()
item_type = _content_item_type(value)
if item_type == "text":
text = _content_item_text(value)
if text is None:
raise ValueError("Text content item is missing a valid 'text' field")
return {"type": "text", "text": text}
if item_type == "image_url":
image_payload = _content_item_image_payload(value)
url = _image_payload_string_field(image_payload, "url")
if url is None:
raise ValueError(
"Image content item is missing a valid 'image_url.url' field"
)
normalized_payload: Dict[str, Any] = {"url": url}
detail = _image_payload_string_field(image_payload, "detail")
if detail in {"auto", "low", "high"}:
normalized_payload["detail"] = detail
local_path = _image_payload_string_field(image_payload, "local_path")
if local_path:
normalized_payload["local_path"] = local_path
return {"type": "image_url", "image_url": normalized_payload}
raise ValueError(f"Unsupported multimodal content item: {type(value).__name__}")
def _normalize_message_content(value: Any) -> Any:
if value is None or isinstance(value, str):
return value
if isinstance(value, list):
return [_normalize_multimodal_content_item(item) for item in value]
return value
class TextContent(BaseModel):
type: Literal["text"] = Field(..., description="Content block type: plain text")
text: str = Field(..., description="Text content of the message")
class ImageURL(BaseModel):
url: str = Field(..., description="Public access URL for the image")
detail: Optional[Literal["auto", "low", "high"]] = Field(
None, description="Optional image detail level"
)
local_path: Optional[str] = Field(None, description="Local path of the image in the workspace")
class ImageContent(BaseModel):
type: Literal["image_url"] = Field(..., description="Content block type: image URL")
image_url: ImageURL = Field(..., description="Detailed information for the image content")
MessageContent = Union[str, None, List[Union[TextContent, ImageContent]]]
class FunctionCall(BaseModel):
name: str = Field(..., description="Name of the function to call")
arguments: str = Field(..., description="JSON-formatted string of arguments to pass")
class ToolCall(BaseModel):
id: str = Field(..., description="Unique ID of this tool call")
type: Literal["function"] = Field(..., description="Type of called tool (function)")
function: FunctionCall = Field(..., description="Function call specification")
class Message(BaseModel):
role: Literal["system", "user", "assistant", "tool"] = Field(
..., description="Role of the message sender"
)
content: MessageContent = Field(
...,
description=(
"Message content. It can be a string, null (when calling tools), or a list of structured multimodal blocks."
),
)
name: Optional[str] = Field(
default=None,
description="Optional sender name, required when the role is 'user' or 'tool'",
max_length=64,
pattern=r"^[a-zA-Z0-9_]*$",
)
tool_calls: Optional[List[ToolCall]] = Field(
default=None, description="List of tool calls the assistant wants to invoke"
)
tool_call_id: Optional[str] = Field(
default=None, description="Tool call ID that this tool message responds to"
)
timestamp: Optional[str] = Field(default=None, description="Message timestamp (ISO format)")
@field_validator("content", mode="before")
@classmethod
def normalize_content(cls, value: Any):
return _normalize_message_content(value)
@model_validator(mode="after")
def validate_tool_message_consistency(cls, values):
role = values.role
content = values.content
tool_calls = values.tool_calls
tool_call_id = values.tool_call_id
if role == "assistant" and tool_calls and content is not None:
raise ValueError(
"When role is 'assistant' and tool_calls exist, content must be None."
)
if role == "tool" and not tool_call_id:
raise ValueError("When role is 'tool', tool_call_id must be provided.")
return values
class ChatMessages(RootModel[List[Message]]):
"""Chat message list ordered chronologically."""
class SketchPadItem(BaseModel):
"""Data structure for a SketchPad storage item."""
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")
@field_validator("last_accessed", mode="before")
@classmethod
def set_last_accessed(cls, v):
"""If last_accessed is None, set it to the current time."""
if v is None:
return datetime.now()
return v
@field_validator("content_hash", mode="before")
@classmethod
def set_content_hash(cls, v, info):
"""If content_hash is None, compute the hash value."""
if v is None:
value = info.data.get("value")
if value is not None:
content_str = str(value)
return hashlib.md5(content_str.encode()).hexdigest()[:8]
return v
def is_expired(self) -> bool:
"""Check whether the item has expired."""
return self.expires_at is not None and datetime.now() > self.expires_at
def update_access(self):
"""Update access information for LRU caching."""
self.access_count += 1
self.last_accessed = datetime.now()
def to_dict(self) -> Dict[str, Any]:
"""Convert to a dictionary for serialization."""
return {
"value": self.value,
"timestamp": self.timestamp.isoformat(),
"summary": self.summary,
"expires_at": self.expires_at.isoformat() if self.expires_at else None,
"access_count": self.access_count,
"last_accessed": (
self.last_accessed.isoformat() if self.last_accessed else None
),
"tags": list(self.tags),
"content_type": self.content_type,
"content_hash": self.content_hash,
}
@classmethod
def from_dict(cls, data: Dict[str, Any]) -> "SketchPadItem":
"""Create an instance from a dictionary."""
# Process time fields.
if isinstance(data.get("timestamp"), str):
data["timestamp"] = datetime.fromisoformat(data["timestamp"])
if data.get("expires_at") and isinstance(data["expires_at"], str):
data["expires_at"] = datetime.fromisoformat(data["expires_at"])
if data.get("last_accessed") and isinstance(data["last_accessed"], str):
data["last_accessed"] = datetime.fromisoformat(data["last_accessed"])
# Process tags.
if data.get("tags") and isinstance(data["tags"], list):
data["tags"] = set(data["tags"])
return cls(**data)
class SketchPadStatistics(BaseModel):
"""SketchPad statistics."""
total_items: int = Field(..., description="Total number of items")
max_items: int = Field(..., description="Maximum number of items")
items_with_summary: int = Field(..., description="Number of items with summaries")
total_accesses: int = Field(..., description="Total number of accesses")
popular_tags: Dict[str, int] = Field(..., description="Popular tag statistics")
content_types: Dict[str, int] = Field(..., description="Content type statistics")
avg_access_per_item: float = Field(..., description="Average accesses per item")
memory_usage_percent: float = Field(..., description="Memory usage percentage")
class SketchPadSearchResult(BaseModel):
"""SketchPad search result."""
key: str = Field(..., description="Item key")
value: Any = Field(..., description="Item value")
summary: Optional[str] = Field(default=None, description="Item summary")
timestamp: str = Field(..., description="Creation time (ISO format)")
tags: List[str] = Field(default_factory=list, description="Tag list")
content_type: str = Field(..., description="Content type")
access_count: int = Field(..., description="Access count")
class SketchPadListItem(BaseModel):
"""SketchPad list item."""
key: str = Field(..., description="Item key")
summary: Optional[str] = Field(default=None, description="Item summary")
timestamp: str = Field(..., description="Creation time (ISO format)")
tags: List[str] = Field(default_factory=list, description="Tag list")
content_type: str = Field(..., description="Content type")
access_count: int = Field(..., description="Access count")
content_hash: Optional[str] = Field(default=None, description="Content hash value")
value: Optional[Any] = Field(default=None, description="Item value, only present when content is included")
@@ -0,0 +1,540 @@
import os
import json
import uuid
import threading
from typing import Dict, Optional, List, Type, Any, cast
from datetime import datetime
import redis
from SimpleLLMFunc import OpenAICompatible
from context.sketch_pad import SketchPadBackend, RedisFileSketchPadBackend
from config.config import get_config
class SketchManager:
"""
General-purpose SketchPad manager that supports different backend implementations.
Main responsibilities:
1. Manage creation and lifecycle of SketchPadBackend instances
2. Provide advanced convenience interfaces
3. Handle batch operations and cleanup tasks
4. Support pluggable backend implementations
"""
_instance = None
_lock: threading.Lock = threading.Lock()
def __new__(cls, backend_class: Type[SketchPadBackend] = RedisFileSketchPadBackend):
"""Singleton pattern implementation."""
if cls._instance is None:
with cls._lock:
if cls._instance is None:
cls._instance = super(SketchManager, cls).__new__(cls)
cls._instance.backend_class = backend_class
return cls._instance
def __init__(self, backend_class: Type[SketchPadBackend]):
"""
Initialize the SketchPad manager.
Args:
backend_class: Backend implementation class, defaulting to RedisFileSketchPadBackend
"""
# Prevent duplicate initialization.
if hasattr(self, "_initialized"):
return
self.backend_class = backend_class
self.config = get_config()
self.sketch_dir = self.config.SKETCH_DIR
self._active_sketches: Dict[str, SketchPadBackend] = {}
# Ensure the directory exists.
os.makedirs(self.sketch_dir, exist_ok=True)
self._initialized = True
def _redis_client(self) -> redis.Redis:
return redis.Redis(
host=self.config.REDIS_HOST,
port=int(self.config.REDIS_PORT),
db=int(self.config.REDIS_DB),
decode_responses=True,
)
def _list_sketch_ids_from_redis(self) -> set[str]:
sketch_ids: set[str] = set()
try:
client = self._redis_client()
raw_keys = cast(Any, client.keys("sketch_pad:*:*"))
for key in cast(List[str], raw_keys):
parts = key.split(":", 2)
if len(parts) >= 3 and parts[0] == "sketch_pad" and parts[1]:
sketch_ids.add(parts[1])
except Exception as e:
print(f"Warning: Failed to list sketch ids from Redis: {e}")
return sketch_ids
def _delete_sketch_redis_keys(self, sketch_id: str) -> bool:
try:
client = self._redis_client()
raw_keys = cast(Any, client.keys(f"sketch_pad:{sketch_id}:*"))
keys = cast(List[str], raw_keys)
if not keys:
return False
deleted = cast(Any, client.delete(*keys))
return int(deleted) > 0
except Exception as e:
print(f"Warning: Failed to delete Redis sketch keys for {sketch_id}: {e}")
return False
def create_sketch_pad(
self,
sketch_id: Optional[str] = None,
**backend_kwargs,
) -> SketchPadBackend:
"""
Create a new SketchPad object.
Args:
sketch_id: SketchPad ID; generated automatically if None
**backend_kwargs: Extra parameters passed to the backend
Returns:
SketchPadBackend: Created SketchPad object
"""
with self._lock:
if sketch_id is None:
sketch_id = str(uuid.uuid4())
# Check whether it already exists.
if sketch_id in self._active_sketches:
return self._active_sketches[sketch_id]
# Generate the file path if the backend needs one.
if "file_path" not in backend_kwargs:
sketch_file = os.path.join(self.sketch_dir, f"skt_{sketch_id}.json")
backend_kwargs["file_path"] = sketch_file
# Create the SketchPad object.
sketch_pad = self.backend_class(
sketch_pad_id=sketch_id,
**backend_kwargs,
)
# Add it to the active SketchPad list.
self._active_sketches[sketch_id] = sketch_pad
return sketch_pad
def get_sketch_pad(self, sketch_id: str) -> Optional[SketchPadBackend]:
"""
Get the SketchPad object with the specified ID.
Args:
sketch_id: SketchPad ID
Returns:
SketchPadBackend: SketchPad object, or None if it does not exist
"""
with self._lock:
# First check active SketchPads.
if sketch_id in self._active_sketches:
return self._active_sketches[sketch_id]
# Try to load from file if the backend supports it.
sketch_file = os.path.join(self.sketch_dir, f"skt_{sketch_id}.json")
if os.path.exists(sketch_file):
try:
sketch_pad = self.backend_class(
sketch_pad_id=sketch_id,
file_path=sketch_file,
)
self._active_sketches[sketch_id] = sketch_pad
return sketch_pad
except Exception as e:
print(f"Warning: Failed to load sketch {sketch_id}: {e}")
return None
def delete_sketch_pad(self, sketch_id: str) -> bool:
"""
Delete the SketchPad object with the specified ID.
Args:
sketch_id: SketchPad ID
Returns:
bool: Whether deletion succeeded
"""
with self._lock:
success = False
# Remove it from active SketchPads.
if sketch_id in self._active_sketches:
del self._active_sketches[sketch_id]
success = True
# Delete sketch keys from Redis.
if self._delete_sketch_redis_keys(sketch_id):
success = True
# Delete the file if it exists.
sketch_file = os.path.join(self.sketch_dir, f"skt_{sketch_id}.json")
if os.path.exists(sketch_file):
try:
os.remove(sketch_file)
success = True
except Exception as e:
print(f"Warning: Failed to delete sketch file {sketch_file}: {e}")
return success
def list_sketch_ids(self) -> List[str]:
"""List all known SketchPad IDs, including Redis and the file system."""
sketch_ids = set(self._active_sketches.keys())
sketch_ids.update(self._list_sketch_ids_from_redis())
try:
for filename in os.listdir(self.sketch_dir):
if filename.startswith("skt_") and filename.endswith(".json"):
sketch_ids.add(filename[4:-5])
except Exception as e:
print(f"Warning: Failed to scan sketch dir for ids: {e}")
return sorted(sketch_ids)
def list_sketch_pads(self) -> List[Dict[str, Any]]:
"""
List all available SketchPads.
Returns:
List[Dict]: SketchPad information list
"""
sketches = []
# Scan SketchPad files in the file system.
try:
for filename in os.listdir(self.sketch_dir):
if filename.startswith("skt_") and filename.endswith(".json"):
sketch_id = filename[4:-5] # Remove the "skt_" prefix and ".json" suffix.
sketch_info = {
"sketch_id": sketch_id,
"file_path": os.path.join(self.sketch_dir, filename),
"is_active": sketch_id in self._active_sketches,
}
# Try to read basic information.
try:
file_path = sketch_info["file_path"]
if isinstance(file_path, str):
with open(file_path, "r", encoding="utf-8") as f:
data = json.load(f)
sketch_info.update(
{
"total_items": len(data.get("items", {})),
"last_saved": data.get(
"serialization_timestamp"
),
"sketch_pad_id": data.get("sketch_pad_id"),
}
)
except Exception:
pass # Ignore read errors.
sketches.append(sketch_info)
except Exception as e:
print(f"Warning: Failed to list sketches: {e}")
return sketches
def save_sketch_pad(self, sketch_id: str) -> bool:
"""
Manually save the specified SketchPad to file.
Args:
sketch_id: SketchPad ID
Returns:
bool: Whether saving succeeded
"""
with self._lock:
if sketch_id in self._active_sketches:
try:
sketch_pad = self._active_sketches[sketch_id]
sketch_pad.persist()
return True
except Exception as e:
print(f"Warning: Failed to save sketch {sketch_id}: {e}")
return False
async def save_all_sketch_pads(self) -> int:
"""
Save all active SketchPads to files.
Returns:
int: Number of SketchPads successfully saved
"""
saved_count = 0
with self._lock:
for sketch_id in list(self._active_sketches.keys()):
if self.save_sketch_pad(sketch_id):
saved_count += 1
return saved_count
async def cleanup_inactive_sketches(self, max_inactive_count: int = 10) -> int:
"""
Clean up inactive SketchPads based on usage frequency.
Args:
max_inactive_count: Maximum number of SketchPads to keep active
Returns:
int: Number of cleaned SketchPads
"""
cleaned_count = 0
with self._lock:
if len(self._active_sketches) <= max_inactive_count:
return 0
# Sort by access statistics and keep the most-used items.
sketches_by_usage = []
for sketch_id, sketch_pad in self._active_sketches.items():
try:
stats = sketch_pad.get_statistics()
total_accesses = stats.total_accesses
sketches_by_usage.append((sketch_id, sketch_pad, total_accesses))
except Exception:
sketches_by_usage.append((sketch_id, sketch_pad, 0))
# Sort by access count.
sketches_by_usage.sort(key=lambda x: x[2], reverse=True)
# Save and remove low-usage SketchPads.
sketches_to_remove = sketches_by_usage[max_inactive_count:]
for sketch_id, sketch_pad, _ in sketches_to_remove:
try:
# Save to file.
sketch_pad.persist()
# Remove from the active list.
del self._active_sketches[sketch_id]
cleaned_count += 1
except Exception as e:
print(f"Warning: Error cleaning sketch {sketch_id}: {e}")
return cleaned_count
# ===== Convenience interfaces =====
async def set_item(
self,
sketch_id: str,
key: str,
value: Any,
ttl: Optional[int] = None,
summary: Optional[str] = None,
tags: Optional[set] = None,
) -> Optional[str]:
"""
Convenience method for setting an item.
Args:
sketch_id: SketchPad ID
key: Key name
value: Value
ttl: Expiration time in seconds
summary: Summary
tags: Tags
Returns:
Optional[str]: Set key name, or None on failure
"""
sketch_pad = self.get_sketch_pad(sketch_id)
if not sketch_pad:
return None
try:
return await sketch_pad.set_item(key, value, ttl, summary, tags)
except Exception as e:
print(f"Warning: Failed to set item: {e}")
return None
def get_item(self, sketch_id: str, key: str) -> Optional[Any]:
"""
Get an item.
Args:
sketch_id: SketchPad ID
key: Key name
Returns:
Optional[Any]: Item value
"""
sketch_pad = self.get_sketch_pad(sketch_id)
if not sketch_pad:
return None
return sketch_pad.get_item(key)
def get_value(self, sketch_id: str, key: str) -> Optional[Any]:
"""
Get a value.
Args:
sketch_id: SketchPad ID
key: Key name
Returns:
Optional[Any]: Value
"""
sketch_pad = self.get_sketch_pad(sketch_id)
if not sketch_pad:
return None
return sketch_pad.get_value(key)
def search_by_tags(
self, sketch_id: str, tags: set, match_all: bool = False
) -> List[tuple]:
"""
Search by tags.
Args:
sketch_id: SketchPad ID
tags: Tag set
match_all: Whether to match all tags
Returns:
List[tuple]: Search results
"""
sketch_pad = self.get_sketch_pad(sketch_id)
if not sketch_pad:
return []
return sketch_pad.search_by_tags(tags, match_all)
def search_by_content(
self, sketch_id: str, query: str, limit: int = 5
) -> List[tuple]:
"""
Search by content.
Args:
sketch_id: SketchPad ID
query: Search query
limit: Result count limit
Returns:
List[tuple]: Search results
"""
sketch_pad = self.get_sketch_pad(sketch_id)
if not sketch_pad:
return []
return sketch_pad.search_by_content(query, limit)
def delete_item(self, sketch_id: str, key: str) -> bool:
"""
Delete an item.
Args:
sketch_id: SketchPad ID
key: Key name
Returns:
bool: Whether deletion succeeded
"""
sketch_pad = self.get_sketch_pad(sketch_id)
if not sketch_pad:
return False
return sketch_pad.delete(key)
def get_statistics(self, sketch_id: str) -> Optional[Any]:
"""
Get statistics.
Args:
sketch_id: SketchPad ID
Returns:
Optional[Any]: Statistics
"""
sketch_pad = self.get_sketch_pad(sketch_id)
if not sketch_pad:
return None
try:
return sketch_pad.get_statistics()
except Exception as e:
print(f"Warning: Failed to get statistics: {e}")
return None
def list_items(self, sketch_id: str, include_value: bool = False) -> List[Any]:
"""
List all items.
Args:
sketch_id: SketchPad ID
include_value: Whether to include values
Returns:
List[Any]: Item list
"""
sketch_pad = self.get_sketch_pad(sketch_id)
if not sketch_pad:
return []
try:
return sketch_pad.list_items(include_value)
except Exception as e:
print(f"Warning: Failed to list items: {e}")
return []
# Global instance.
_global_sketch_manager: Optional[SketchManager] = None
def get_sketch_manager() -> SketchManager:
"""Get the global SketchManager instance."""
global _global_sketch_manager
if _global_sketch_manager is None:
# Get Redis configuration from config.
config = get_config()
# redis config
redis_host = config.REDIS_HOST
redis_port = int(config.REDIS_PORT)
redis_db = int(config.REDIS_DB)
# Create a custom backend class with preconfigured Redis parameters.
class ConfiguredRedisFileSketchPadBackend(RedisFileSketchPadBackend):
def __init__(
self,
sketch_pad_id: str,
redis_host: str = redis_host,
redis_port: int = redis_port,
redis_db: int = redis_db,
file_path: Optional[str] = None,
):
super().__init__(
sketch_pad_id=sketch_pad_id,
redis_host=redis_host,
redis_port=redis_port,
redis_db=redis_db,
file_path=file_path,
)
# Create SketchManager using the configured backend class.
_global_sketch_manager = SketchManager(
backend_class=ConfiguredRedisFileSketchPadBackend
)
return _global_sketch_manager
+586
View File
@@ -0,0 +1,586 @@
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