Files
agent-desktop/tests/unit/agents/test_memory_middleware.py

556 lines
18 KiB
Python

# -*- coding: utf-8 -*-
"""Tests for MemoryMiddleware automation-source skip logic."""
# pylint: disable=protected-access
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from agentscope.message import Msg, TextBlock
from pineagents.agents.middlewares import MemoryMiddleware
from pineagents.constant import (
EXTERNAL_USER_QUERY_MESSAGE_TAG,
LOOP_CONTINUATION_MESSAGE_TAG,
QWENPAW_MESSAGE_TAG_KEY,
)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_agent(*, source: str | None = None):
"""Build a minimal fake agent with optional request_context source."""
agent = MagicMock()
agent.name = "TestAgent"
agent.state = SimpleNamespace(
context=[],
session_id="session-1",
reply_id="reply-1",
)
agent._context_manager = None
if source is not None:
agent._request_context = {"source": source, "session_id": "session-1"}
else:
agent._request_context = {"session_id": "session-1"}
return agent
def _user_msg(text: str = "hello", *, msg_id: str = "turn-1") -> Msg:
msg = Msg(
name="user",
role="user",
content=[TextBlock(type="text", text=text)],
metadata={
QWENPAW_MESSAGE_TAG_KEY: EXTERNAL_USER_QUERY_MESSAGE_TAG,
},
)
msg.id = msg_id
return msg
def _make_memory_manager(*, interval: int = 1):
mm = MagicMock()
mm.agent_id = "test-agent"
mm.get_auto_memory_interval.return_value = interval
mm.auto_memory = AsyncMock()
mm.auto_memory_search = AsyncMock(return_value=None)
mm.get_memory_prompt.return_value = ""
mm._auto_memory_turn_states = {}
def _get_auto_memory_turn_state(session_id: str):
return mm._auto_memory_turn_states.setdefault(
session_id or "__default__",
{
"pending": [],
"seen": {},
"touched_at": 0,
},
)
mm.get_auto_memory_turn_state.side_effect = _get_auto_memory_turn_state
return mm
def _auto_memory_turn_state(mm, session_id: str = "session-1"):
return mm.get_auto_memory_turn_state(session_id)
# ---------------------------------------------------------------------------
# _is_automation_request unit tests
# ---------------------------------------------------------------------------
class TestIsAutomationRequest:
def test_cron_source(self):
agent = _make_agent(source="cron")
assert MemoryMiddleware._is_automation_request(agent) is True
def test_heartbeat_source(self):
agent = _make_agent(source="heartbeat")
assert MemoryMiddleware._is_automation_request(agent) is True
def test_cron_uppercase(self):
agent = _make_agent(source="CRON")
assert MemoryMiddleware._is_automation_request(agent) is True
def test_heartbeat_mixed_case(self):
agent = _make_agent(source="HeartBeat")
assert MemoryMiddleware._is_automation_request(agent) is True
def test_user_source(self):
agent = _make_agent(source="user")
assert MemoryMiddleware._is_automation_request(agent) is False
def test_empty_source(self):
agent = _make_agent(source="")
assert MemoryMiddleware._is_automation_request(agent) is False
def test_no_source_key(self):
agent = _make_agent(source=None)
assert MemoryMiddleware._is_automation_request(agent) is False
def test_no_request_context_attr(self):
agent = MagicMock(spec=[])
assert MemoryMiddleware._is_automation_request(agent) is False
def test_request_context_not_dict(self):
agent = MagicMock()
agent._request_context = "not-a-dict"
assert MemoryMiddleware._is_automation_request(agent) is False
# ---------------------------------------------------------------------------
# on_model_call integration tests
# ---------------------------------------------------------------------------
class TestOnModelCallAutomationSkip:
@pytest.mark.asyncio
async def test_cron_skips_auto_memory_search(self):
"""Automation requests must skip auto_memory_search entirely."""
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="cron")
agent.state.context = [_user_msg()]
next_handler = AsyncMock(return_value="model_result")
result = await mw.on_model_call(agent, {"messages": []}, next_handler)
mm.auto_memory_search.assert_not_awaited()
next_handler.assert_awaited_once()
assert result == "model_result"
@pytest.mark.asyncio
async def test_user_calls_auto_memory_search(self):
"""Normal user requests should trigger auto_memory_search."""
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="user")
agent.state.context = [_user_msg()]
next_handler = AsyncMock(return_value="model_result")
await mw.on_model_call(agent, {"messages": []}, next_handler)
mm.auto_memory_search.assert_awaited_once()
assert mm.auto_memory_search.await_args.args[0].id == "turn-1"
@pytest.mark.asyncio
async def test_untagged_user_message_does_not_search(self):
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="user")
agent.state.context = [
Msg(
name="user",
role="user",
content=[TextBlock(text="internal prompt")],
),
]
await mw.on_model_call(
agent,
{"messages": []},
AsyncMock(return_value="model_result"),
)
mm.auto_memory_search.assert_not_awaited()
@pytest.mark.asyncio
async def test_loop_continuation_does_not_retrigger_search(self):
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="user")
real_query = _user_msg("real query")
agent.state.context = [real_query]
next_handler = AsyncMock(return_value="model_result")
await mw.on_model_call(agent, {"messages": []}, next_handler)
continuation = Msg(
name="user",
role="user",
content=[
TextBlock(text="[WARNING] Repetitive pattern detected."),
],
metadata={
QWENPAW_MESSAGE_TAG_KEY: LOOP_CONTINUATION_MESSAGE_TAG,
},
)
agent.state.context.append(continuation)
await mw.on_model_call(agent, {"messages": []}, next_handler)
mm.auto_memory_search.assert_awaited_once()
assert mm.auto_memory_search.await_args.args[0] is real_query
@pytest.mark.asyncio
async def test_model_call_search_state_survives_middleware_rebuild(self):
"""A rebuilt middleware must not search twice for the same turn."""
mm = _make_memory_manager()
agent = _make_agent(source="user")
agent.state.context = [_user_msg(msg_id="turn-1")]
next_handler = AsyncMock(return_value="model_result")
await MemoryMiddleware(memory_manager=mm).on_model_call(
agent,
{"messages": []},
next_handler,
)
await MemoryMiddleware(memory_manager=mm).on_model_call(
agent,
{"messages": []},
next_handler,
)
mm.auto_memory_search.assert_awaited_once()
assert _auto_memory_turn_state(mm)["searched_turn"] == "turn-1"
# ---------------------------------------------------------------------------
# on_reply integration tests
# ---------------------------------------------------------------------------
class TestOnReplyAutomationSkip:
@pytest.mark.asyncio
async def test_cron_skips_marker_tracking(self):
"""Automation requests must not append to pending markers."""
mm = _make_memory_manager(interval=1)
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="cron")
agent.state.context = [_user_msg()]
async def _next(**_kwargs):
yield "done"
gen = mw.on_reply(agent, {}, _next)
async for _ in gen:
pass
state = _auto_memory_turn_state(mm)
assert not state["pending"]
assert not state["seen"]
mm.auto_memory.assert_not_awaited()
@pytest.mark.asyncio
async def test_user_triggers_auto_memory(self):
"""Normal user requests should trigger auto_memory as usual."""
mm = _make_memory_manager(interval=1)
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="user")
agent.state.context = [_user_msg()]
async def _next(**_kwargs):
yield "done"
gen = mw.on_reply(agent, {}, _next)
async for _ in gen:
pass
mm.auto_memory.assert_awaited_once()
@pytest.mark.asyncio
async def test_internal_user_message_is_excluded_from_memory(self):
"""Internal user-role controls must not enter auto-memory."""
mm = _make_memory_manager(interval=1)
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="user")
query = _user_msg("real query")
reply = Msg(
name="agent",
role="assistant",
content=[TextBlock(text="reply")],
)
continuation = Msg(
name="user",
role="user",
content=[TextBlock(text="[WARNING] Repetitive pattern detected.")],
metadata={
QWENPAW_MESSAGE_TAG_KEY: LOOP_CONTINUATION_MESSAGE_TAG,
},
)
final_reply = Msg(
name="agent",
role="assistant",
content=[TextBlock(text="done")],
)
agent.state.context = [query, reply, continuation, final_reply]
async def _next(**_kwargs):
yield "done"
async for _ in mw.on_reply(agent, {}, _next):
pass
mm.auto_memory.assert_awaited_once()
assert mm.auto_memory.await_args.args[0] == [query, reply, final_reply]
@pytest.mark.asyncio
async def test_interval_state_survives_middleware_rebuild(self):
"""A rebuilt middleware must keep interval state on the manager."""
mm = _make_memory_manager(interval=2)
async def _next(**_kwargs):
yield "done"
agent1 = _make_agent(source="user")
agent1.state.context = [_user_msg(msg_id="turn-1")]
gen1 = MemoryMiddleware(memory_manager=mm).on_reply(
agent1,
{},
_next,
)
async for _ in gen1:
pass
mm.auto_memory.assert_not_awaited()
assert _auto_memory_turn_state(mm)["pending"] == ["turn-1"]
agent2 = _make_agent(source="user")
agent2.state.context = [
_user_msg(msg_id="turn-1"),
Msg(
name="agent",
role="assistant",
content=[TextBlock(text="reply 1")],
),
_user_msg(msg_id="turn-2"),
Msg(
name="agent",
role="assistant",
content=[TextBlock(text="reply 2")],
),
]
gen2 = MemoryMiddleware(memory_manager=mm).on_reply(
agent2,
{},
_next,
)
async for _ in gen2:
pass
mm.auto_memory.assert_awaited_once()
assert not _auto_memory_turn_state(mm)["pending"]
# ---------------------------------------------------------------------------
# on_compress_context integration tests
# ---------------------------------------------------------------------------
class TestOnCompressContextAutomationSkip:
@pytest.mark.asyncio
async def test_heartbeat_skips_memory_flush_but_compresses(self):
"""Automation skips memory flush; compression still runs."""
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="heartbeat")
next_handler = AsyncMock()
await mw.on_compress_context(agent, {}, next_handler)
next_handler.assert_awaited_once_with()
mm.auto_memory.assert_not_awaited()
@pytest.mark.asyncio
async def test_heartbeat_does_not_call_will_compress(self):
"""_will_compress_context must NOT be called for automation."""
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="heartbeat")
next_handler = AsyncMock()
with patch.object(
MemoryMiddleware,
"_will_compress_context",
) as mock_wc:
await mw.on_compress_context(agent, {}, next_handler)
mock_wc.assert_not_called()
@pytest.mark.asyncio
async def test_normal_request_may_flush_on_compress(self):
"""Non-automation requests follow the normal compress path."""
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="user")
_auto_memory_turn_state(mm)["pending"] = ["m1"]
next_handler = AsyncMock()
with patch.object(
MemoryMiddleware,
"_memory_config",
) as mock_cfg, patch.object(
MemoryMiddleware,
"_will_compress_context",
return_value=True,
) as mock_wc:
cfg = MagicMock()
cfg.summarize_when_compact = True
mock_cfg.return_value = cfg
agent.state.context = [_user_msg()]
await mw.on_compress_context(agent, {}, next_handler)
mock_wc.assert_awaited_once()
next_handler.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"failing_step",
[
"memory_config",
"turn_state",
"will_compress",
"flush",
],
)
async def test_memory_failure_does_not_block_compression(
self,
failing_step,
):
"""Memory failures must not disable the context safety valve."""
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="user")
next_handler = AsyncMock()
if failing_step == "memory_config":
mm.get_memory_config.side_effect = RuntimeError(
"memory config unavailable",
)
else:
_auto_memory_turn_state(mm)["pending"] = ["m1"]
if failing_step == "turn_state":
mm.get_auto_memory_turn_state.side_effect = RuntimeError(
"memory state unavailable",
)
will_compress = AsyncMock(return_value=True)
flush = AsyncMock()
if failing_step == "will_compress":
will_compress.side_effect = RuntimeError("token count failed")
if failing_step == "flush":
flush.side_effect = RuntimeError("memory flush failed")
with patch.object(
MemoryMiddleware,
"_will_compress_context",
will_compress,
), patch.object(
MemoryMiddleware,
"_flush_auto_memory",
flush,
):
await mw.on_compress_context(agent, {}, next_handler)
next_handler.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_compression_failure_is_not_swallowed(self):
"""Only memory failures are fail-open; compression still fails loud."""
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="user")
next_handler = AsyncMock(
side_effect=RuntimeError("scroll compression failed"),
)
with pytest.raises(RuntimeError, match="scroll compression failed"):
await mw.on_compress_context(agent, {}, next_handler)
class TestWillCompressContextBoundary:
@staticmethod
def _agent_at(tokens: int):
agent = _make_agent(source="user")
agent.context_config = SimpleNamespace(trigger_ratio=0.8)
agent.model = SimpleNamespace(
context_size=1000,
count_tokens=AsyncMock(return_value=tokens),
)
agent._prepare_model_input = AsyncMock(return_value={})
return agent
@pytest.mark.asyncio
async def test_native_compacts_at_exact_trigger(self):
agent = self._agent_at(800)
assert await MemoryMiddleware._will_compress_context(agent, {}) is True
@pytest.mark.asyncio
async def test_scroll_does_not_compact_at_exact_trigger(self):
agent = self._agent_at(800)
agent._context_manager = SimpleNamespace(
should_compress=lambda tokens, trigger: tokens > trigger,
)
assert (
await MemoryMiddleware._will_compress_context(agent, {}) is False
)
@pytest.mark.asyncio
async def test_scroll_compacts_above_trigger(self):
agent = self._agent_at(801)
agent._context_manager = SimpleNamespace(
should_compress=lambda tokens, trigger: tokens > trigger,
)
assert await MemoryMiddleware._will_compress_context(agent, {}) is True
# ---------------------------------------------------------------------------
# _flush_auto_memory defensive guard
# ---------------------------------------------------------------------------
class TestFlushAutoMemoryDefensiveGuard:
@pytest.mark.asyncio
async def test_automation_clears_pending_and_skips(self):
"""Defensive guard in _flush_auto_memory clears markers."""
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="cron")
_auto_memory_turn_state(mm)["pending"] = ["m1", "m2"]
await mw._flush_auto_memory(agent)
assert not _auto_memory_turn_state(mm)["pending"]
mm.auto_memory.assert_not_awaited()
@pytest.mark.asyncio
async def test_normal_request_flushes(self):
"""Non-automation requests proceed with auto_memory."""
mm = _make_memory_manager()
mw = MemoryMiddleware(memory_manager=mm)
agent = _make_agent(source="user")
_auto_memory_turn_state(mm)["pending"] = ["turn-1"]
agent.state.context = [_user_msg()]
await mw._flush_auto_memory(agent)
mm.auto_memory.assert_awaited_once()