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
|
||||
"""
|
||||
Reference in New Issue
Block a user