Files
cad-Integration/CADDesigner-Code-main/context/context.py
T
2026-07-22 13:48:46 +08:00

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
"""