mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-28 01:31:46 +00:00
feat(history): add history chunking and retrieval features
- Implemented history chunking logic with hybrid turn/token-based splitting - Added RetrieveHistory tool for semantic search on history chunks - Updated AddHistory to store parent nodes with multiple chunk nodes - Modified read_history to accept optional single history_id parameter - Integrated RetrieveHistory into main retrieval workflow with multiple agents - Updated personal_retriever.yaml to prioritize chunk retrieval over full reads - Added comprehensive test coverage for chunking and retrieval functionality - Enhanced logging and metadata tracking for history operations
This commit is contained in:
parent
d27ca21386
commit
6e9e4e43d7
8 changed files with 476 additions and 16 deletions
|
|
@ -53,18 +53,20 @@ user_message_s2: |
|
|||
- Refine Phase 1 queries with 3-5 diverse appropriate time filters
|
||||
|
||||
### Phase 3: Deep Dive into History
|
||||
**Tool**: `read_history`
|
||||
**When to use**: After exhausting retrieval attempts OR when specific conversation context is needed
|
||||
**Primary Tool**: `retrieve_history`
|
||||
**Fallback Tool**: `read_history`
|
||||
**When to use**: After retrieving memory references with `history_id`, or when specific conversation context is needed
|
||||
**Important Constraints**:
|
||||
- Each history is very long and resource-intensive to read
|
||||
- **Maximum limit: Read no more than 3 histories total**
|
||||
- Only use this phase when absolutely necessary for answering the question
|
||||
- `retrieve_history` searches concise history chunks and should be preferred before reading full histories
|
||||
- `read_history` loads the entire original history and is resource-intensive
|
||||
- **Maximum limit: Read no more than 3 full histories total with `read_history`**
|
||||
- Only use full-history reading when chunk retrieval is insufficient for answering the question
|
||||
**Approach**:
|
||||
- Extract `history_id` from retrieved memory references
|
||||
- Prioritize the most relevant or recent histories
|
||||
- Can read multiple histories at once by passing multiple history_ids
|
||||
- Be selective: choose only the top 1-3 most promising histories
|
||||
- Use this to understand the full conversation surrounding a memory
|
||||
- Use `retrieve_history` with focused queries; pass `history_id` to search within a specific history group
|
||||
- If no `history_id` is available, use `retrieve_history` without `history_id` to search across all history chunks
|
||||
- Prioritize the most relevant or recent histories and retrieve only the chunks needed
|
||||
- Use `read_history` only for the top 1-3 most promising histories when full surrounding conversation is necessary
|
||||
|
||||
## Response Guidelines
|
||||
- Base your answer EXCLUSIVELY on the profile search results, retrieved memories, and history data
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from .delegate_task import DelegateTask
|
|||
from .history.add_history import AddHistory
|
||||
from .history.read_history import ReadHistory
|
||||
from .history.read_history_v2 import ReadHistoryV2
|
||||
from .history.retrieve_history import RetrieveHistory
|
||||
|
||||
# profiles tools
|
||||
from .profiles.add_draft_and_read_all_profiles import AddDraftAndReadAllProfiles
|
||||
|
|
@ -41,6 +42,7 @@ __all__ = [
|
|||
"AddHistory",
|
||||
"ReadHistory",
|
||||
"ReadHistoryV2",
|
||||
"RetrieveHistory",
|
||||
# profiles tools
|
||||
"AddDraftAndReadAllProfiles",
|
||||
"AddProfile",
|
||||
|
|
|
|||
|
|
@ -8,14 +8,24 @@ from ..base_memory_tool import BaseMemoryTool
|
|||
from ....core.enumeration import MemoryType
|
||||
from ....core.schema import ToolCall, MemoryNode, Message
|
||||
from ....core.utils import format_messages
|
||||
from .history_chunking import split_history_messages
|
||||
|
||||
|
||||
class AddHistory(BaseMemoryTool):
|
||||
"""Tool to add historical dialogue to vector store"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
def __init__(
|
||||
self,
|
||||
chunk_strategy: str = "hybrid",
|
||||
turn_block_size: int = 3,
|
||||
max_chunk_tokens: int = 800,
|
||||
**kwargs,
|
||||
):
|
||||
kwargs["enable_multiple"] = False
|
||||
super().__init__(**kwargs)
|
||||
self.chunk_strategy = chunk_strategy
|
||||
self.turn_block_size = turn_block_size
|
||||
self.max_chunk_tokens = max_chunk_tokens
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
"""Build and return the tool call schema"""
|
||||
|
|
@ -35,6 +45,12 @@ class AddHistory(BaseMemoryTool):
|
|||
self.context.messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages]
|
||||
history_content: str = self.context.description + "\n" + format_messages(self.context.messages)
|
||||
history_content = history_content.strip()
|
||||
history_chunks = split_history_messages(
|
||||
self.context.messages,
|
||||
chunk_strategy=self.chunk_strategy,
|
||||
turn_block_size=self.turn_block_size,
|
||||
max_chunk_tokens=self.max_chunk_tokens,
|
||||
)
|
||||
history_node = MemoryNode(
|
||||
memory_type=MemoryType.HISTORY,
|
||||
when_to_use=history_content[:1024],
|
||||
|
|
@ -45,13 +61,48 @@ class AddHistory(BaseMemoryTool):
|
|||
[m.model_dump(exclude_none=True) for m in self.context.messages],
|
||||
ensure_ascii=False,
|
||||
),
|
||||
"node_kind": "history",
|
||||
"chunk_count": len(history_chunks),
|
||||
"chunk_strategy": self.chunk_strategy,
|
||||
"turn_block_size": self.turn_block_size,
|
||||
"max_chunk_tokens": self.max_chunk_tokens,
|
||||
},
|
||||
)
|
||||
self.context.history_node = history_node
|
||||
logger.info(f"Adding history node: {history_node.model_dump_json(indent=2)}")
|
||||
|
||||
vector_node = history_node.to_vector_node()
|
||||
await self.vector_store.delete(vector_node.vector_id)
|
||||
await self.vector_store.insert([vector_node])
|
||||
chunk_nodes = [
|
||||
MemoryNode(
|
||||
memory_id=f"{history_node.memory_id}:chunk:{chunk.index:04d}",
|
||||
memory_type=MemoryType.HISTORY,
|
||||
memory_target=history_node.memory_target,
|
||||
when_to_use=chunk.content,
|
||||
content=chunk.content,
|
||||
author=self.author,
|
||||
ref_memory_id=history_node.memory_id,
|
||||
metadata={
|
||||
"node_kind": "history_chunk",
|
||||
"history_id": history_node.memory_id,
|
||||
"chunk_index": chunk.index,
|
||||
"start_message_index": chunk.start_message_index,
|
||||
"end_message_index": chunk.end_message_index,
|
||||
"token_count": chunk.token_count,
|
||||
"chunk_strategy": self.chunk_strategy,
|
||||
},
|
||||
)
|
||||
for chunk in history_chunks
|
||||
]
|
||||
|
||||
return f"Successfully added history: {history_node.memory_id}"
|
||||
vector_nodes = [history_node.to_vector_node(), *[node.to_vector_node() for node in chunk_nodes]]
|
||||
existing_chunks = await self.vector_store.list(
|
||||
filters={
|
||||
"node_kind": "history_chunk",
|
||||
"history_id": history_node.memory_id,
|
||||
},
|
||||
)
|
||||
delete_ids = [node.vector_id for node in vector_nodes]
|
||||
delete_ids.extend(node.vector_id for node in existing_chunks if node.vector_id not in delete_ids)
|
||||
await self.vector_store.delete(delete_ids)
|
||||
await self.vector_store.insert(vector_nodes)
|
||||
|
||||
return f"Successfully added history: {history_node.memory_id} ({len(chunk_nodes)} chunks)"
|
||||
|
|
|
|||
139
reme/memory/vector_tools/history/history_chunking.py
Normal file
139
reme/memory/vector_tools/history/history_chunking.py
Normal file
|
|
@ -0,0 +1,139 @@
|
|||
"""Helpers for chunking history messages before vector insertion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
from ....core.enumeration import Role
|
||||
from ....core.schema import Message
|
||||
from ....core.utils import format_messages
|
||||
|
||||
|
||||
_TOKEN_PATTERN = re.compile(
|
||||
r"[\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff]|[A-Za-z0-9]+(?:[-_'][A-Za-z0-9]+)*|[^\s]",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class HistoryChunk:
|
||||
"""A chunk of conversation history ready to be stored as one vector node."""
|
||||
|
||||
index: int
|
||||
start_message_index: int
|
||||
end_message_index: int
|
||||
content: str
|
||||
token_count: int
|
||||
|
||||
|
||||
def normalize_messages(messages: list[Message | dict]) -> list[Message]:
|
||||
"""Convert raw dict messages to Message objects."""
|
||||
return [Message(**message) if isinstance(message, dict) else message for message in messages]
|
||||
|
||||
|
||||
def estimate_mixed_tokens(text: str) -> int:
|
||||
"""Estimate tokens while treating CJK chars and English words differently."""
|
||||
if not text:
|
||||
return 0
|
||||
return len(_TOKEN_PATTERN.findall(text))
|
||||
|
||||
|
||||
def split_history_messages(
|
||||
messages: list[Message],
|
||||
*,
|
||||
chunk_strategy: str = "hybrid",
|
||||
turn_block_size: int = 3,
|
||||
max_chunk_tokens: int = 800,
|
||||
) -> list[HistoryChunk]:
|
||||
"""Split messages by dialogue turns and/or approximate token count.
|
||||
|
||||
``turn_block_size`` counts user-started dialogue rounds. ``max_chunk_tokens``
|
||||
uses a lightweight mixed Chinese/English tokenizer to keep chunks retrieval-sized
|
||||
without requiring an external tokenizer dependency.
|
||||
"""
|
||||
if not messages:
|
||||
return []
|
||||
|
||||
chunk_strategy = chunk_strategy.lower()
|
||||
if chunk_strategy not in {"turn", "token", "hybrid"}:
|
||||
raise ValueError("chunk_strategy must be one of: turn, token, hybrid")
|
||||
|
||||
turn_block_size = max(1, turn_block_size)
|
||||
max_chunk_tokens = max(1, max_chunk_tokens)
|
||||
|
||||
units = _build_turn_units(messages) if chunk_strategy in {"turn", "hybrid"} else _build_message_units(messages)
|
||||
chunks: list[HistoryChunk] = []
|
||||
current_units: list[tuple[int, int, list[Message]]] = []
|
||||
current_turns = 0
|
||||
|
||||
def current_text() -> str:
|
||||
current_messages = [message for _, _, unit_messages in current_units for message in unit_messages]
|
||||
return format_messages(current_messages)
|
||||
|
||||
def flush() -> None:
|
||||
nonlocal current_units, current_turns
|
||||
if not current_units:
|
||||
return
|
||||
text = current_text().strip()
|
||||
if text:
|
||||
chunks.append(
|
||||
HistoryChunk(
|
||||
index=len(chunks),
|
||||
start_message_index=current_units[0][0],
|
||||
end_message_index=current_units[-1][1],
|
||||
content=text,
|
||||
token_count=estimate_mixed_tokens(text),
|
||||
),
|
||||
)
|
||||
current_units = []
|
||||
current_turns = 0
|
||||
|
||||
for unit in units:
|
||||
unit_text = format_messages(unit[2])
|
||||
unit_tokens = estimate_mixed_tokens(unit_text)
|
||||
next_turns = current_turns + 1
|
||||
should_split_by_turn = chunk_strategy in {"turn", "hybrid"} and next_turns > turn_block_size
|
||||
should_split_by_token = (
|
||||
chunk_strategy in {"token", "hybrid"}
|
||||
and current_units
|
||||
and estimate_mixed_tokens(current_text()) + unit_tokens > max_chunk_tokens
|
||||
)
|
||||
|
||||
if should_split_by_turn or should_split_by_token:
|
||||
flush()
|
||||
|
||||
current_units.append(unit)
|
||||
current_turns += 1
|
||||
|
||||
if chunk_strategy == "token" and unit_tokens >= max_chunk_tokens:
|
||||
flush()
|
||||
|
||||
flush()
|
||||
return chunks
|
||||
|
||||
|
||||
def _build_message_units(messages: list[Message]) -> list[tuple[int, int, list[Message]]]:
|
||||
"""Treat each message as a split unit for token-only chunking."""
|
||||
return [(index, index, [message]) for index, message in enumerate(messages)]
|
||||
|
||||
|
||||
def _build_turn_units(messages: list[Message]) -> list[tuple[int, int, list[Message]]]:
|
||||
"""Group messages into user-started dialogue turns."""
|
||||
units: list[tuple[int, int, list[Message]]] = []
|
||||
current_start = 0
|
||||
current_messages: list[Message] = []
|
||||
|
||||
for index, message in enumerate(messages):
|
||||
role = message.role.value if isinstance(message.role, Role) else str(message.role)
|
||||
if role == Role.USER.value and current_messages:
|
||||
units.append((current_start, index - 1, current_messages))
|
||||
current_start = index
|
||||
current_messages = []
|
||||
elif not current_messages:
|
||||
current_start = index
|
||||
current_messages.append(message)
|
||||
|
||||
if current_messages:
|
||||
units.append((current_start, len(messages) - 1, current_messages))
|
||||
|
||||
return units
|
||||
|
|
@ -39,20 +39,38 @@ class ReadHistory(BaseMemoryTool):
|
|||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"history_id": {
|
||||
"type": "string",
|
||||
"description": "Single history ID to read",
|
||||
},
|
||||
"history_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "List of history IDs to read",
|
||||
},
|
||||
},
|
||||
"required": ["history_ids"],
|
||||
"required": [],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
"""Execute the tool call"""
|
||||
history_ids = self.context.history_ids if self.enable_multiple else [self.context.history_id]
|
||||
if self.enable_multiple:
|
||||
if "history_ids" in self.context:
|
||||
raw_history_ids = self.context.history_ids
|
||||
if isinstance(raw_history_ids, str):
|
||||
history_ids = [raw_history_ids]
|
||||
else:
|
||||
history_ids = list(raw_history_ids)
|
||||
elif "history_id" in self.context:
|
||||
history_ids = [self.context.history_id]
|
||||
else:
|
||||
history_ids = []
|
||||
else:
|
||||
history_ids = [self.context.history_id]
|
||||
|
||||
history_ids = [history_id for history_id in history_ids if history_id]
|
||||
|
||||
if not history_ids or (len(history_ids) == 1 and not history_ids[0]):
|
||||
output = "No history_ids provided."
|
||||
|
|
|
|||
114
reme/memory/vector_tools/history/retrieve_history.py
Normal file
114
reme/memory/vector_tools/history/retrieve_history.py
Normal file
|
|
@ -0,0 +1,114 @@
|
|||
"""Retrieve relevant history chunks from vector store."""
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.schema import MemoryNode, ToolCall
|
||||
|
||||
|
||||
class RetrieveHistory(BaseMemoryTool):
|
||||
"""Retrieve history chunks by semantic similarity."""
|
||||
|
||||
def __init__(self, top_k: int = 5, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.top_k = top_k
|
||||
|
||||
def _build_query_parameters(self) -> dict:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "Query used to retrieve relevant history chunks.",
|
||||
},
|
||||
"history_id": {
|
||||
"type": "string",
|
||||
"description": "Optional history_id. If provided, search only within this history group.",
|
||||
},
|
||||
},
|
||||
"required": ["query"],
|
||||
}
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
"""Build and return the tool call schema."""
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": "Retrieve relevant history chunks. Optionally restrict search by history_id.",
|
||||
"parameters": self._build_query_parameters(),
|
||||
},
|
||||
)
|
||||
|
||||
def _build_multiple_tool_call(self) -> ToolCall:
|
||||
"""Build and return the tool call schema for multiple retrieval queries."""
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": "Retrieve relevant history chunks for multiple queries.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query_items": {
|
||||
"type": "array",
|
||||
"description": "List of query items.",
|
||||
"items": self._build_query_parameters(),
|
||||
},
|
||||
},
|
||||
"required": ["query_items"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
"""Execute semantic retrieval over stored history chunks."""
|
||||
query_items = self.context.get("query_items", []) if self.enable_multiple else [self.context]
|
||||
if not query_items:
|
||||
output = "No history retrieval queries provided."
|
||||
logger.warning(output)
|
||||
return output
|
||||
|
||||
memory_nodes: list[MemoryNode] = []
|
||||
for item in query_items:
|
||||
query = item["query"]
|
||||
filters = {"node_kind": "history_chunk"}
|
||||
history_id = item.get("history_id")
|
||||
if history_id:
|
||||
filters["history_id"] = history_id
|
||||
|
||||
vector_nodes = await self.vector_store.search(
|
||||
query=query,
|
||||
limit=self.top_k,
|
||||
filters=filters,
|
||||
)
|
||||
memory_nodes.extend(MemoryNode.from_vector_node(node) for node in vector_nodes)
|
||||
|
||||
retrieved_ids = {node.memory_id for node in self.retrieved_nodes if node.memory_id}
|
||||
new_nodes = []
|
||||
for node in memory_nodes:
|
||||
if node.memory_id not in retrieved_ids:
|
||||
retrieved_ids.add(node.memory_id)
|
||||
new_nodes.append(node)
|
||||
|
||||
self.retrieved_nodes.extend(new_nodes)
|
||||
|
||||
if not new_nodes:
|
||||
output = "No relevant history found."
|
||||
else:
|
||||
output = "\n\n".join(self._format_history_chunk(node) for node in new_nodes)
|
||||
|
||||
logger.info(f"Retrieved {len(memory_nodes)} history chunks, {len(new_nodes)} new after deduplication")
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def _format_history_chunk(node: MemoryNode) -> str:
|
||||
history_id = node.metadata.get("history_id") or node.ref_memory_id
|
||||
chunk_index = node.metadata.get("chunk_index")
|
||||
score = node.metadata.get("score") or node.score
|
||||
|
||||
header = f"Historical Dialogue Chunk[{node.memory_id}]"
|
||||
if history_id:
|
||||
header += f" history_id={history_id}"
|
||||
if chunk_index is not None and chunk_index != "":
|
||||
header += f" chunk_index={chunk_index}"
|
||||
if score:
|
||||
header += f" score={float(score):.4f}"
|
||||
|
||||
return f"{header}\n{node.content}"
|
||||
20
reme/reme.py
20
reme/reme.py
|
|
@ -14,6 +14,7 @@ from .memory.vector_tools import (
|
|||
DelegateTask,
|
||||
ReadAllProfiles,
|
||||
ReadHistory,
|
||||
RetrieveHistory,
|
||||
RetrieveProfile,
|
||||
RetrieveMemory,
|
||||
UpdateProfilesV1,
|
||||
|
|
@ -495,6 +496,12 @@ class ReMe(Application):
|
|||
enable_multiple=True,
|
||||
raise_exception=raise_exception,
|
||||
),
|
||||
RetrieveHistory(
|
||||
top_k=retrieve_top_k,
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_multiple=True,
|
||||
raise_exception=raise_exception,
|
||||
),
|
||||
],
|
||||
)
|
||||
personal_retriever: BaseMemoryAgent = PersonalRetriever(
|
||||
|
|
@ -520,6 +527,12 @@ class ReMe(Application):
|
|||
enable_multiple=True,
|
||||
raise_exception=raise_exception,
|
||||
),
|
||||
RetrieveHistory(
|
||||
top_k=retrieve_top_k,
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_multiple=True,
|
||||
raise_exception=raise_exception,
|
||||
),
|
||||
],
|
||||
raise_exception=raise_exception,
|
||||
)
|
||||
|
|
@ -538,6 +551,12 @@ class ReMe(Application):
|
|||
enable_multiple=True,
|
||||
raise_exception=raise_exception,
|
||||
),
|
||||
RetrieveHistory(
|
||||
top_k=retrieve_top_k,
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_multiple=True,
|
||||
raise_exception=raise_exception,
|
||||
),
|
||||
],
|
||||
raise_exception=raise_exception,
|
||||
)
|
||||
|
|
@ -585,6 +604,7 @@ class ReMe(Application):
|
|||
|
||||
reme_retriever: BaseMemoryAgent = ReMeRetriever(
|
||||
tools=[DelegateTask(memory_agents=memory_agents, raise_exception=raise_exception)],
|
||||
prompt_path=Path(__file__).parent / "memory" / "vector_based" / "reme_retriever.yaml",
|
||||
raise_exception=raise_exception,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,9 @@ import pytest
|
|||
import reme.reme as reme_module
|
||||
from reme.core.runtime_context import RuntimeContext
|
||||
from reme.core.schema import MemoryNode
|
||||
from reme.memory.vector_tools.history.add_history import AddHistory
|
||||
from reme.memory.vector_tools.history.read_history import ReadHistory
|
||||
from reme.memory.vector_tools.history.retrieve_history import RetrieveHistory
|
||||
from reme.reme import ReMe
|
||||
|
||||
|
||||
|
|
@ -85,6 +87,22 @@ async def test_summarize_memory_propagates_raise_exception(
|
|||
assert all(instance.kwargs.get("raise_exception") is raise_exception for instance in Recorder.instances)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve_memory_binds_top_level_retriever_prompt_path(monkeypatch: pytest.MonkeyPatch):
|
||||
"""Verify top-level ReMeRetriever does not accidentally inherit child retriever prompts."""
|
||||
_patch_retrieve_dependencies(monkeypatch)
|
||||
reme = _make_reme()
|
||||
|
||||
result = await reme.retrieve_memory(
|
||||
query="hello",
|
||||
user_name="alice",
|
||||
)
|
||||
|
||||
assert result == "ok"
|
||||
top_level = next(instance for instance in Recorder.instances if isinstance(instance, TopLevelAgent))
|
||||
assert str(top_level.kwargs["prompt_path"]).endswith("reme/memory/vector_based/reme_retriever.yaml")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("raise_exception", [False, True])
|
||||
async def test_retrieve_memory_propagates_raise_exception(
|
||||
|
|
@ -166,3 +184,99 @@ async def test_read_history_accepts_single_history_id_in_multiple_mode():
|
|||
result = await tool.execute()
|
||||
|
||||
assert "Historical Dialogue[history_123]" in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_history_stores_parent_and_chunk_nodes():
|
||||
"""Verify AddHistory stores a full parent node plus retrievable chunk nodes."""
|
||||
|
||||
class FakeVectorStore:
|
||||
"""Minimal vector store stub for AddHistory tests."""
|
||||
|
||||
def __init__(self):
|
||||
self.deleted_ids = None
|
||||
self.inserted_nodes = []
|
||||
|
||||
async def list(self, filters=None, limit=None, sort_key=None, reverse=True):
|
||||
return []
|
||||
|
||||
async def delete(self, vector_ids):
|
||||
self.deleted_ids = vector_ids
|
||||
|
||||
async def insert(self, nodes):
|
||||
self.inserted_nodes = nodes
|
||||
|
||||
vector_store = FakeVectorStore()
|
||||
tool = AddHistory(turn_block_size=1, max_chunk_tokens=1000)
|
||||
tool._vector_store = vector_store # pylint: disable=protected-access
|
||||
tool.context = RuntimeContext(
|
||||
description="A coffee chat.",
|
||||
messages=[
|
||||
{"role": "user", "content": "我想喝咖啡"},
|
||||
{"role": "assistant", "content": "好的,您喜欢什么类型的咖啡?"},
|
||||
{"role": "user", "content": "Latte with oat milk please."},
|
||||
{"role": "assistant", "content": "没问题。"},
|
||||
],
|
||||
author="tester",
|
||||
service_context=SimpleNamespace(memory_target_type_mapping={"alice": "personal"}),
|
||||
)
|
||||
|
||||
result = await tool.execute()
|
||||
|
||||
assert "Successfully added history:" in result
|
||||
assert len(vector_store.inserted_nodes) == 3
|
||||
parent_node = vector_store.inserted_nodes[0]
|
||||
chunk_nodes = vector_store.inserted_nodes[1:]
|
||||
assert parent_node.metadata["node_kind"] == "history"
|
||||
assert parent_node.metadata["chunk_count"] == 2
|
||||
assert all(node.metadata["node_kind"] == "history_chunk" for node in chunk_nodes)
|
||||
assert all(node.metadata["history_id"] == parent_node.vector_id for node in chunk_nodes)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve_history_filters_by_optional_history_id():
|
||||
"""Verify RetrieveHistory can search within a history group."""
|
||||
|
||||
class FakeVectorStore:
|
||||
"""Minimal vector store stub for RetrieveHistory tests."""
|
||||
|
||||
def __init__(self):
|
||||
self.search_calls = []
|
||||
|
||||
async def search(self, query, limit, filters):
|
||||
self.search_calls.append({"query": query, "limit": limit, "filters": filters})
|
||||
return [
|
||||
MemoryNode(
|
||||
memory_id="history_123:chunk:0000",
|
||||
memory_type="history",
|
||||
content="user: 我想喝咖啡\nassistant: 好的",
|
||||
ref_memory_id="history_123",
|
||||
metadata={
|
||||
"node_kind": "history_chunk",
|
||||
"history_id": "history_123",
|
||||
"chunk_index": 0,
|
||||
},
|
||||
).to_vector_node(),
|
||||
]
|
||||
|
||||
vector_store = FakeVectorStore()
|
||||
tool = RetrieveHistory(top_k=2, enable_multiple=False)
|
||||
tool._vector_store = vector_store # pylint: disable=protected-access
|
||||
tool.context = RuntimeContext(
|
||||
query="咖啡",
|
||||
history_id="history_123",
|
||||
retrieved_nodes=[],
|
||||
service_context=SimpleNamespace(memory_target_type_mapping={"alice": "personal"}),
|
||||
)
|
||||
|
||||
result = await tool.execute()
|
||||
|
||||
assert vector_store.search_calls == [
|
||||
{
|
||||
"query": "咖啡",
|
||||
"limit": 2,
|
||||
"filters": {"node_kind": "history_chunk", "history_id": "history_123"},
|
||||
},
|
||||
]
|
||||
assert "Historical Dialogue Chunk[history_123:chunk:0000]" in result
|
||||
assert "history_id=history_123" in result
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue