ReMe/reme2/file_based/summarizer.py
jinliyl 52f1a33b3a
Some checks are pending
Pre-commit / run (ubuntu-latest) (push) Waiting to run
feat: add file parser support and search filtering (#214)
- Add file_parser component with default implementation
- Introduce SearchFilter schema for path and tag filtering
- Implement filter functionality in BaseFileStore and LocalFileStore
- Update file watcher to use parser-based filtering instead of suffix filters
- Register new FILE_PARSER component enum
- Add test_data directory to gitignore

refactor: improve component imports and initialization

- Fix relative imports in application.py
- Add file_parser import to component init
- Initialize registry dict when component type doesn't exist
- Remove circular import in HttpService by using string annotation
- Update config yaml to use proper component names

refactor: enhance file watcher architecture

- Replace MdFileWatcher with more flexible FullFileWatcher and LightFileWatcher
- Remove suffix-based filtering in favor of parser-based approach
- Update BaseFileWatcher to resolve parsers from app context
- Remove unused watch_filter method

refactor: update ReMe core functionality

- Remove memory_path creation
- Simplify dream and proactive methods to return empty strings
- Update config defaults for HTTP service and component backends

docs: update component configuration in paw.yaml

- Change service backend from cmd to http
- Rename components to use correct singular forms
- Add default file parser and file watcher configurations
- Set up local file store with default settings
```

Co-authored-by: huangsen <huangsen.huang@alibaba-inc.com>
2026-04-21 16:44:46 +08:00

324 lines
12 KiB
Python

"""Summarizer module for memory summarization operations."""
import datetime
import json
import zoneinfo
from agentscope.agent import ReActAgent
from agentscope.message import Msg
from agentscope.token import HuggingFaceTokenCounter
from agentscope.tool import Toolkit
from loguru import logger
from ..component import BaseStep
from ..schema import AsMsgStat, AsBlockStat
class Summarizer(BaseStep):
"""Summarizer step for summarizing memory messages."""
def __init__(
self,
working_dir: str,
memory_dir: str,
memory_compact_threshold: int,
toolkit: Toolkit | None = None,
console_enabled: bool = False,
timezone: str | None = None,
add_thinking_block: bool = True,
as_token_counter: HuggingFaceTokenCounter | None = None,
**kwargs,
):
"""Initialize the summarizer step.
Args:
working_dir: Working directory path.
memory_dir: Memory directory path for storing summaries.
memory_compact_threshold: Token threshold for memory compaction.
toolkit: Optional toolkit for the agent.
console_enabled: Whether to enable console output.
timezone: Optional timezone string for date formatting.
add_thinking_block: Whether to include thinking blocks in output.
as_token_counter: Optional token counter instance.
**kwargs: Additional keyword arguments passed to BaseStep.
"""
super().__init__(**kwargs)
self.working_dir: str = working_dir
self.memory_dir: str = memory_dir
self.memory_compact_threshold: int = memory_compact_threshold
self.toolkit: Toolkit | None = toolkit
self.console_enabled: bool = console_enabled
self.timezone: str | None = timezone
self.add_thinking_block: bool = add_thinking_block
self._as_token_counter: HuggingFaceTokenCounter | None = as_token_counter
def _get_current_datetime(self) -> datetime.datetime:
"""Get current datetime with timezone, fallback to local time if timezone is invalid."""
if self.timezone:
try:
return datetime.datetime.now(zoneinfo.ZoneInfo(self.timezone))
except Exception as e:
self.logger.error(f"Invalid timezone: {self.timezone}, falling back to local time error={e}")
return datetime.datetime.now()
async def _count_str_token(self, text: str) -> int:
"""Count tokens in a string."""
return await self.as_token_counter.count(messages=[], text=text)
async def _format_tool_result_output(self, output: str | list[dict]) -> tuple[str, int]:
"""Convert tool result output to string."""
if isinstance(output, str):
return output, await self._count_str_token(output)
textual_parts = []
total_token_count = 0
for block in output:
try:
if not isinstance(block, dict) or "type" not in block:
logger.warning(
f"Invalid block: {block}, expected a dict with 'type' key, skipped.",
)
continue
block_type = block["type"]
if block_type == "text":
textual_parts.append(block.get("text", ""))
total_token_count += await self._count_str_token(textual_parts[-1])
elif block_type in ["image", "audio", "video"]:
source = block.get("source", {})
if source.get("type") == "base64":
data = source.get("data", "")
total_token_count += len(data) // 4 if data else 10
else:
url = source.get("url", "")
total_token_count += await self._count_str_token(url) if url else 10
textual_parts.append(f"[{block_type}] {url}")
elif block_type == "file":
file_path = block.get("path", "") or block.get("url", "")
file_name = block.get("name", file_path)
textual_parts.append(f"[file] {file_name}: {file_path}")
total_token_count += await self._count_str_token(file_path)
else:
logger.warning(
f"Unsupported block type '{block_type}' in tool result, skipped.",
)
except Exception as e:
logger.warning(
f"Failed to process block {block}: {e}, skipped.",
)
return "\n".join(textual_parts), total_token_count
async def _stat_message(self, message: Msg) -> AsMsgStat:
"""Analyze a message and generate block statistics."""
blocks = []
if isinstance(message.content, str):
blocks.append(
AsBlockStat(
block_type="text",
text=message.content,
token_count=await self._count_str_token(message.content),
),
)
return AsMsgStat(
name=message.name or message.role,
role=message.role,
content=blocks,
timestamp=message.timestamp or "",
metadata=message.metadata or {},
)
for block in message.content:
block_type = block.get("type", "unknown")
if block_type == "text":
text = block.get("text", "")
token_count = await self._count_str_token(text)
blocks.append(
AsBlockStat(
block_type=block_type,
text=text,
token_count=token_count,
),
)
elif block_type == "thinking":
thinking = block.get("thinking", "")
token_count = await self._count_str_token(thinking)
blocks.append(
AsBlockStat(
block_type=block_type,
text=thinking,
token_count=token_count,
),
)
elif block_type in ("image", "audio", "video"):
source = block.get("source", {})
url = source.get("url", "")
if source.get("type") == "base64":
data = source.get("data", "")
token_count = len(data) // 4 if data else 10
else:
token_count = await self._count_str_token(url) if url else 10
blocks.append(
AsBlockStat(
block_type=block_type,
text="",
token_count=token_count,
media_url=url,
),
)
elif block_type == "tool_use":
tool_name = block.get("name", "")
tool_input = block.get("input", "")
try:
input_str = json.dumps(tool_input, ensure_ascii=False)
except (TypeError, ValueError):
input_str = str(tool_input)
token_count = await self._count_str_token(tool_name + input_str)
blocks.append(
AsBlockStat(
block_type=block_type,
text="",
token_count=token_count,
tool_name=tool_name,
tool_input=input_str,
),
)
elif block_type == "tool_result":
tool_name = block.get("name", "")
output = block.get("output", "")
formatted_output, token_count = await self._format_tool_result_output(output)
blocks.append(
AsBlockStat(
block_type=block_type,
text="",
token_count=token_count,
tool_name=tool_name,
tool_output=formatted_output,
),
)
else:
logger.warning(f"Unsupported block type {block_type}, skipped.")
return AsMsgStat(
name=message.name or message.role,
role=message.role,
content=blocks,
timestamp=message.timestamp or "",
metadata=message.metadata or {},
)
async def _count_msgs_token(self, messages: list[Msg]) -> int:
"""Count total token count of a list of messages."""
total = 0
for msg in messages:
stat = await self._stat_message(msg)
total += stat.total_tokens
return total
async def _format_msgs_to_str(
self,
messages: list[Msg],
memory_compact_threshold: int,
include_thinking: bool = True,
) -> str:
"""Format list of messages to a single formatted string.
Messages are processed in reverse order (newest first) and older
messages are skipped when token count exceeds memory_compact_threshold.
Args:
messages: List of Msg objects to format.
memory_compact_threshold: Maximum token count before skipping older messages.
include_thinking: Whether to include thinking blocks in output.
"""
if not messages:
return ""
formatted_parts: list[str] = []
total_token_count = 0
for i in range(len(messages) - 1, -1, -1):
stat = await self._stat_message(messages[i])
formatted_content = stat.format(include_thinking=include_thinking)
content_token_count = await self._count_str_token(formatted_content)
is_latest = i == len(messages) - 1
if not is_latest and total_token_count + content_token_count > memory_compact_threshold:
logger.info(
f"Skipping older messages: adding {content_token_count} tokens would exceed threshold "
f"{memory_compact_threshold} (current: {total_token_count})",
)
break
if is_latest and content_token_count > memory_compact_threshold:
logger.warning(
f"Latest message alone ({content_token_count} tokens) exceeds threshold "
f"{memory_compact_threshold}, including it anyway.",
)
formatted_parts.append(formatted_content)
total_token_count += content_token_count
formatted_parts.reverse()
return "\n\n".join(formatted_parts)
async def execute(self):
"""Execute the summarization step."""
messages: list[Msg] = self.context.data.get("messages", [])
if not messages:
return ""
before_token_count = await self._count_msgs_token(messages)
history_formatted_str: str = await self._format_msgs_to_str(
messages=messages,
memory_compact_threshold=self.memory_compact_threshold,
include_thinking=self.add_thinking_block,
)
after_token_count = await self._count_str_token(history_formatted_str)
logger.info(f"Summarizer before_token_count={before_token_count} after_token_count={after_token_count}")
if not history_formatted_str:
logger.warning(f"No history to summarize. messages={messages}")
return ""
agent = ReActAgent(
name="reme_summarizer",
model=self.as_llm.model,
sys_prompt="You are a helpful assistant.",
formatter=self.as_llm_formatter.formatter,
toolkit=self.toolkit,
)
agent.set_console_output_enabled(self.console_enabled)
user_message: str = f"# conversation\n{history_formatted_str}\n\n" + self.prompt_format(
"user_message",
date=self._get_current_datetime().strftime("%Y-%m-%d"),
working_dir=self.working_dir,
memory_dir=self.memory_dir,
)
summary_msg: Msg = await agent.reply(
Msg(
name="reme",
role="user",
content=user_message,
),
)
for i, (msg, _) in enumerate(agent.memory.content):
logger.info(f"Summarizer memory[{i}]: {msg.content}")
history_summary: str = summary_msg.get_text_content()
logger.info(f"Summarizer Result:\n{history_summary}")
return history_summary