first commit
This commit is contained in:
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user