635 lines
21 KiB
Python
635 lines
21 KiB
Python
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
|
|
"""
|