first commit
This commit is contained in:
@@ -0,0 +1,473 @@
|
||||
"""
|
||||
BaseAgent is the base class for all agents. It defines the basic agent interface and common functionality.
|
||||
All concrete agent implementations should inherit from this class and implement the abstract methods.
|
||||
|
||||
BaseAgent provides the following features:
|
||||
- Singleton pattern
|
||||
- Conversation history management
|
||||
- SketchPad management
|
||||
- Toolkit management
|
||||
"""
|
||||
|
||||
from bootstrap_env import load_project_env
|
||||
|
||||
load_project_env()
|
||||
|
||||
from typing import (
|
||||
Dict,
|
||||
List,
|
||||
Optional,
|
||||
Generator,
|
||||
Sequence,
|
||||
Tuple,
|
||||
AsyncGenerator,
|
||||
Any,
|
||||
)
|
||||
from abc import ABC, abstractmethod
|
||||
from SimpleLLMFunc import llm_chat, OpenAICompatible # type: ignore
|
||||
import threading
|
||||
from context.conversation_manager import get_current_context, get_current_sketch_pad
|
||||
from context.schemas import Message
|
||||
from react_stream import extract_output_text, is_response_yield
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
|
||||
|
||||
class BaseAgent(ABC):
|
||||
"""
|
||||
Agent base class, defining the basic agent interface and common functionality.
|
||||
|
||||
All concrete agent implementations should inherit from this class and implement the abstract methods.
|
||||
"""
|
||||
|
||||
# Class-level instance cache, ensuring a singleton for each Agent subclass.
|
||||
_class_instances: Dict[str, "BaseAgent"] = {}
|
||||
_class_lock = threading.Lock()
|
||||
|
||||
@classmethod
|
||||
def get_instance(
|
||||
cls,
|
||||
model_name: str,
|
||||
name: Optional[str] = None,
|
||||
description: Optional[str] = None,
|
||||
llm_interface: Optional[OpenAICompatible] = None,
|
||||
**kwargs,
|
||||
) -> "BaseAgent":
|
||||
"""
|
||||
Class method for obtaining an Agent instance (singleton pattern).
|
||||
|
||||
Args:
|
||||
model_name: Model name
|
||||
name: Agent name
|
||||
description: Agent description
|
||||
llm_interface: LLM interface
|
||||
**kwargs: Other parameters
|
||||
|
||||
Returns:
|
||||
Agent instance
|
||||
"""
|
||||
with cls._class_lock:
|
||||
# Use the class name and model_name as the unique identifier.
|
||||
instance_key = f"{cls.__name__}:{model_name}"
|
||||
|
||||
if instance_key not in cls._class_instances:
|
||||
if not llm_interface:
|
||||
# If llm_interface is not provided, try to obtain it from the configuration.
|
||||
from config.config import get_config
|
||||
|
||||
config = get_config()
|
||||
llm_interface = config.BASIC_INTERFACE
|
||||
|
||||
instance_name = name or f"{model_name}-agent"
|
||||
instance_description = description or f"Agent instance for {model_name}"
|
||||
|
||||
cls._class_instances[instance_key] = cls(
|
||||
name=instance_name,
|
||||
description=instance_description,
|
||||
llm_interface=llm_interface,
|
||||
model_name=model_name,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
return cls._class_instances[instance_key]
|
||||
|
||||
@classmethod
|
||||
def clear_instances(cls):
|
||||
"""Clear all instance caches."""
|
||||
with cls._class_lock:
|
||||
cls._class_instances.clear()
|
||||
|
||||
@classmethod
|
||||
def get_all_instances(cls) -> Dict[str, "BaseAgent"]:
|
||||
"""Get all instances."""
|
||||
return cls._class_instances.copy()
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
name: str,
|
||||
description: str,
|
||||
llm_interface: Optional[OpenAICompatible] = None,
|
||||
model_name: Optional[str] = None, # Add the model_name parameter.
|
||||
**kwargs, # Extra parameters that subclasses can handle.
|
||||
):
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.model_name = model_name # Store model_name.
|
||||
self.llm_interface = llm_interface
|
||||
|
||||
if not self.llm_interface:
|
||||
raise ValueError("llm_interface must be provided")
|
||||
|
||||
# Subclasses need to define their own toolkit.
|
||||
self.toolkit = self.get_toolkit()
|
||||
|
||||
# Initialize the chat function.
|
||||
self.chat = llm_chat(
|
||||
llm_interface=self.llm_interface,
|
||||
toolkit=self.toolkit, # type: ignore
|
||||
stream=True,
|
||||
return_mode="raw",
|
||||
enable_event=True,
|
||||
max_tool_calls=2000,
|
||||
timeout=600,
|
||||
temperature=1.0,
|
||||
)(self.chat_impl)
|
||||
|
||||
@abstractmethod
|
||||
def get_toolkit(self) -> Sequence[Any]:
|
||||
"""
|
||||
Get the agent-specific toolkit (abstract method).
|
||||
|
||||
Subclasses must implement this method to define their own toolkit.
|
||||
|
||||
Returns:
|
||||
List of tool functions
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def chat_impl(
|
||||
self,
|
||||
history: List[Dict[str, Any]],
|
||||
query: Any,
|
||||
sketch_pad_summary: str,
|
||||
) -> Generator[Tuple[str, List[Dict[str, Any]]], None, None]:
|
||||
"""
|
||||
Agent conversation implementation logic (abstract method).
|
||||
|
||||
Subclasses must implement this method to define the concrete conversation behavior.
|
||||
|
||||
Args:
|
||||
history: Conversation history
|
||||
query: User query
|
||||
sketch_pad_summary: SketchPad summary
|
||||
|
||||
Returns:
|
||||
Generator yielding (response_chunk, updated_history)
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def run(
|
||||
self, query: Any, raw_user_content: Any = None
|
||||
) -> AsyncGenerator[Any, None]:
|
||||
"""
|
||||
Run the agent to process the user query (abstract method).
|
||||
|
||||
Args:
|
||||
query: User query
|
||||
raw_user_content: Raw user message content (optional), used for persisting multimodal messages
|
||||
|
||||
Returns:
|
||||
AsyncGenerator yielding response chunks
|
||||
"""
|
||||
pass
|
||||
|
||||
# Common helper methods.
|
||||
def get_sketch_pad_summary(self) -> str:
|
||||
"""Get SketchPad summary information, including all keys and truncated values."""
|
||||
try:
|
||||
sketch_pad = get_current_sketch_pad()
|
||||
if sketch_pad is None:
|
||||
return "SketchPad unavailable: no active conversation context"
|
||||
|
||||
# Get detailed information for all items (including values).
|
||||
all_items = sketch_pad.list_items(include_value=True)
|
||||
|
||||
if not all_items:
|
||||
return "SketchPad is empty: no stored content"
|
||||
|
||||
summary_lines = [f"Current SketchPad state ({len(all_items)} items total):"]
|
||||
|
||||
for item in all_items[:20]: # Limit display to the first 20 items.
|
||||
key = item.key
|
||||
tags = ", ".join(item.tags) if item.tags else "no tags"
|
||||
timestamp = item.timestamp
|
||||
content_type = item.content_type
|
||||
|
||||
# Use the value included in the list item for preview.
|
||||
value_obj = item.value
|
||||
value_str = str(value_obj) if value_obj is not None else ""
|
||||
if len(value_str) > 100:
|
||||
value_preview = value_str[:100] + "..."
|
||||
else:
|
||||
value_preview = value_str
|
||||
|
||||
value_preview = value_preview.replace("\n", "\\n")
|
||||
|
||||
summary_lines.append(
|
||||
f" - {key}: [{content_type}] {value_preview} "
|
||||
f"(tags: {tags}, time: {timestamp[:19]})"
|
||||
)
|
||||
|
||||
if len(all_items) > 20:
|
||||
summary_lines.append(f" ... {len(all_items) - 20} more items not shown")
|
||||
|
||||
return "\n".join(summary_lines)
|
||||
|
||||
except Exception as e:
|
||||
return f"Error while retrieving SketchPad summary: {str(e)}"
|
||||
|
||||
# Convenience methods for context management.
|
||||
def get_conversation_history(self, limit: Optional[int] = None):
|
||||
"""Get the conversation history for the current session."""
|
||||
context = get_current_context()
|
||||
if context is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
return context.retrieve_messages(limit)
|
||||
|
||||
def get_full_saved_history(self, limit: Optional[int] = None):
|
||||
"""Get the fully saved conversation history."""
|
||||
context = get_current_context()
|
||||
if context is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
return context.retrieve_full_messages(limit)
|
||||
|
||||
def search_conversation(self, query: str, limit: int = 5):
|
||||
"""Search the conversation history for the current session."""
|
||||
context = get_current_context()
|
||||
if context is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
# Use a simple search implementation.
|
||||
return context.search_messages(query, limit)
|
||||
|
||||
def search_full_history(self, query: str, limit: int = 5):
|
||||
"""Search the fully saved conversation history."""
|
||||
context = get_current_context()
|
||||
if context is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
return context.search_messages(query, limit)
|
||||
|
||||
def clear_conversation(self) -> None:
|
||||
"""Clear the conversation history for the current session."""
|
||||
context = get_current_context()
|
||||
if context is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
context.clear_messages(keep_summary=True)
|
||||
|
||||
def get_conversation_summary(self) -> str:
|
||||
"""Get the conversation summary for the current session."""
|
||||
context = get_current_context()
|
||||
if context is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
return context.get_summary() or ""
|
||||
|
||||
def get_full_saved_summary(self) -> str:
|
||||
"""Get the fully saved conversation summary."""
|
||||
context = get_current_context()
|
||||
if context is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
return context.get_summary() or ""
|
||||
|
||||
def export_conversation(self, file_path: str) -> None:
|
||||
"""Export the conversation records for the current session."""
|
||||
context = get_current_context()
|
||||
if context is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
data = context.serialize()
|
||||
dir_path = os.path.dirname(file_path)
|
||||
if dir_path:
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
with open(file_path, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, ensure_ascii=False, indent=2)
|
||||
|
||||
def import_conversation(self, file_path: str, merge: bool = False) -> None:
|
||||
"""Import conversation records."""
|
||||
context = get_current_context()
|
||||
if context is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
if not merge:
|
||||
# Clear existing messages while preserving the summary.
|
||||
context.clear_messages(keep_summary=True)
|
||||
context.deserialize(data)
|
||||
|
||||
# Convenience methods for SketchPad management.
|
||||
async def store_in_sketch_pad(
|
||||
self,
|
||||
value,
|
||||
key: Optional[str] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
ttl: Optional[int] = None,
|
||||
) -> str:
|
||||
"""Store data in SketchPad."""
|
||||
sketch_pad = get_current_sketch_pad()
|
||||
if sketch_pad is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
# Generate a key name if one is not provided.
|
||||
item_key = key or f"item_{uuid.uuid4().hex[:8]}"
|
||||
# Convert tags to a set.
|
||||
tags_set = set(tags) if tags else None
|
||||
await sketch_pad.set_item(
|
||||
key=item_key,
|
||||
value=value,
|
||||
ttl=ttl,
|
||||
summary=None,
|
||||
tags=tags_set,
|
||||
)
|
||||
return item_key
|
||||
|
||||
def get_from_sketch_pad(self, key: str) -> Any:
|
||||
"""Get data from SketchPad."""
|
||||
sketch_pad = get_current_sketch_pad()
|
||||
if sketch_pad is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
return sketch_pad.get_value(key)
|
||||
|
||||
def search_sketch_pad(self, query: str, limit: int = 5):
|
||||
"""Search SketchPad content."""
|
||||
sketch_pad = get_current_sketch_pad()
|
||||
if sketch_pad is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
return sketch_pad.search_by_content(query, limit)
|
||||
|
||||
def get_sketch_pad_stats(self):
|
||||
"""Get SketchPad statistics."""
|
||||
sketch_pad = get_current_sketch_pad()
|
||||
if sketch_pad is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
return sketch_pad.get_statistics()
|
||||
|
||||
def clear_sketch_pad(self):
|
||||
"""Clear SketchPad."""
|
||||
sketch_pad = get_current_sketch_pad()
|
||||
if sketch_pad is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
sketch_pad.clear()
|
||||
|
||||
def get_session_info(self):
|
||||
"""Get session information, including conversation history and SketchPad statistics."""
|
||||
try:
|
||||
conversation_count = len(self.get_conversation_history())
|
||||
sketch_pad_stats = self.get_sketch_pad_stats()
|
||||
conversation_summary = self.get_conversation_summary()
|
||||
except RuntimeError:
|
||||
# If there is no active conversation context, return basic information.
|
||||
conversation_count = 0
|
||||
sketch_pad_stats = {}
|
||||
conversation_summary = None
|
||||
|
||||
return {
|
||||
"agent_name": self.name,
|
||||
"model_name": self.model_name,
|
||||
"agent_class": self.__class__.__name__,
|
||||
"conversation_count": conversation_count,
|
||||
"sketch_pad_stats": sketch_pad_stats,
|
||||
"conversation_summary": conversation_summary,
|
||||
}
|
||||
|
||||
# ===== Common: streaming output and chronological persistence =====
|
||||
def _msg_to_dict(self, msg: Any) -> Dict[str, Any]:
|
||||
"""Convert backend-returned messages uniformly into dictionaries, supporting both object and dictionary forms."""
|
||||
if isinstance(msg, dict):
|
||||
return msg
|
||||
return {
|
||||
"role": getattr(msg, "role", None),
|
||||
"content": getattr(msg, "content", None),
|
||||
"tool_calls": getattr(msg, "tool_calls", None),
|
||||
"tool_call_id": getattr(msg, "tool_call_id", None),
|
||||
}
|
||||
|
||||
async def _stream_and_persist(
|
||||
self, response_packages: AsyncGenerator[Any, None]
|
||||
) -> AsyncGenerator[Any, None]:
|
||||
"""
|
||||
Unified streaming processing and history persistence logic:
|
||||
- Continuously accumulate assistant text; when encountering tooluse/tool results, persist the accumulated text first, then write the tool message.
|
||||
- Ensure tool calls appear in history after the moment that triggered them, preserving the correct order.
|
||||
"""
|
||||
context = get_current_context()
|
||||
if context is None:
|
||||
raise RuntimeError("No active conversation context")
|
||||
|
||||
assistant_buffer: str = ""
|
||||
baseline_len: Optional[int] = None
|
||||
|
||||
async for output in response_packages:
|
||||
yield output
|
||||
|
||||
if not is_response_yield(output):
|
||||
continue
|
||||
|
||||
current_messages = output.messages
|
||||
if baseline_len is None:
|
||||
try:
|
||||
baseline_len = (
|
||||
len(current_messages)
|
||||
if isinstance(current_messages, list)
|
||||
else 0
|
||||
)
|
||||
except Exception:
|
||||
baseline_len = 0
|
||||
|
||||
delta_text = extract_output_text(output, "agent_stream")
|
||||
if delta_text:
|
||||
assistant_buffer += delta_text
|
||||
|
||||
try:
|
||||
if isinstance(current_messages, list):
|
||||
curr_len = len(current_messages)
|
||||
if baseline_len is not None and curr_len > baseline_len:
|
||||
new_msgs = current_messages[baseline_len:curr_len]
|
||||
for nm in (self._msg_to_dict(x) for x in new_msgs):
|
||||
role = nm.get("role")
|
||||
content = nm.get("content")
|
||||
tool_calls = nm.get("tool_calls")
|
||||
tool_call_id = nm.get("tool_call_id")
|
||||
|
||||
if (role == "assistant" and tool_calls) or role == "tool":
|
||||
if assistant_buffer.strip():
|
||||
await context.store_message(
|
||||
Message(
|
||||
role="assistant", content=assistant_buffer
|
||||
)
|
||||
)
|
||||
assistant_buffer = ""
|
||||
|
||||
if role == "assistant" and tool_calls:
|
||||
await context.store_message(
|
||||
Message(
|
||||
role="assistant",
|
||||
content=None,
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
)
|
||||
elif role == "tool":
|
||||
await context.store_message(
|
||||
Message(
|
||||
role="tool",
|
||||
content=content,
|
||||
tool_call_id=tool_call_id,
|
||||
)
|
||||
)
|
||||
baseline_len = curr_len
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# At the end of the stream, write the remaining assistant text.
|
||||
if assistant_buffer.strip():
|
||||
await context.store_message(
|
||||
Message(role="assistant", content=assistant_buffer)
|
||||
)
|
||||
Reference in New Issue
Block a user