Files
temp/features/steps/memory_service_coverage_steps.py
T

358 lines
12 KiB
Python

"""Behave steps covering memory_service module."""
from __future__ import annotations
import asyncio
import shutil
import tempfile
from pathlib import Path
from typing import Any
from unittest.mock import patch
from behave import given, then, when
from langchain_core.chat_history import InMemoryChatMessageHistory
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage
from cleveragents.application.services.memory_service import (
ConversationBufferMemoryAdapter,
MemoryService,
)
@given("a conversation adapter configured to return text values")
def step_adapter_returning_text(context: Any) -> None:
"""Create an adapter that returns buffer strings."""
history = InMemoryChatMessageHistory()
history.add_user_message("Hello from user")
history.add_ai_message("Greetings from AI")
context.adapter = ConversationBufferMemoryAdapter(
history,
return_messages=False,
)
context.adapter_keys = context.adapter.memory_variables
@when("I load memory variables from the adapter")
def step_load_memory_variables(context: Any) -> None:
"""Load memory payload from the adapter."""
context.loaded_memory = context.adapter.load_memory_variables()
@then("the adapter provides a string buffer representation")
def step_validate_string_payload(context: Any) -> None:
"""Ensure adapter returns buffer string when configured."""
assert context.adapter.memory_key in context.loaded_memory
payload = context.loaded_memory[context.adapter.memory_key]
assert isinstance(payload, str)
assert payload
assert "Hello from user" in payload or "Greetings from AI" in payload
assert context.adapter_keys == [context.adapter.memory_key]
@given("a conversation adapter with prune tracking")
def step_adapter_with_prune(context: Any) -> None:
"""Create adapter that records prune callback invocations."""
history = InMemoryChatMessageHistory()
context.prune_calls = 0
def record_prune() -> None:
context.prune_calls += 1
context.adapter = ConversationBufferMemoryAdapter(
history, prune_callback=record_prune
)
context.adapter_keys = context.adapter.memory_variables
@when("I save nested user and ai context through the adapter")
def step_save_nested_context(context: Any) -> None:
"""Store nested values to exercise adapter recursion."""
inputs = {
"input": [
HumanMessage(content="first human"),
"plain user",
["nested user", HumanMessage(content="second human")],
None,
]
}
outputs = {
"output": [
AIMessage(content="first ai"),
"plain ai",
[AIMessage(content="second ai"), "another ai"],
None,
]
}
context.adapter.save_context(inputs, outputs)
context.saved_snapshot = [
msg.content for msg in context.adapter._message_history.messages
]
@when("I invoke the adapter async operations")
def step_async_adapter_ops(context: Any) -> None:
"""Exercise async adapter APIs."""
async def run_ops() -> None:
context.async_loaded = await context.adapter.aload_memory_variables({})
await context.adapter.asave_context(
{"input": [HumanMessage(content="async human"), ["async text"]]},
{"output": [AIMessage(content="async ai"), ["async reply"]]},
)
context.post_async_snapshot = [
msg.content for msg in context.adapter._message_history.messages
]
await context.adapter.aclear()
asyncio.run(run_ops())
@then("the adapter normalizes all values and prunes history three times")
def step_verify_prune_and_storage(context: Any) -> None:
"""Validate adapter processed values and triggered pruning."""
assert context.prune_calls == 3
assert len(context.saved_snapshot) >= 4
assert any("plain user" in content for content in context.saved_snapshot)
async_payload = context.async_loaded[context.adapter.memory_key]
assert isinstance(async_payload, list)
assert any(isinstance(msg, BaseMessage) for msg in async_payload)
assert len(context.post_async_snapshot) >= 2
assert not context.adapter._message_history.messages
@given("a memory service limited to {count:d} messages")
def step_memory_service_with_limit(context: Any, count: int) -> None:
"""Create memory service with max message window."""
context.memory_service = MemoryService(
session_id="limited-session", max_messages=count
)
@when("I add five interactions via save context")
def step_add_multiple_interactions(context: Any) -> None:
"""Add several interactions to trigger pruning."""
for idx in range(5):
context.memory_service.save_context(
{"input": f"user {idx}"},
{"output": f"ai {idx}"},
)
context.final_messages = context.memory_service.get_messages()
@when("I ask for the two most recent messages")
def step_recent_two(context: Any) -> None:
"""Fetch limited recent messages."""
context.recent_two = context.memory_service.get_recent_messages(2)
@when("I ask for more messages than available")
def step_recent_all(context: Any) -> None:
"""Fetch more messages than stored to exercise branch."""
context.recent_all = context.memory_service.get_recent_messages(10)
@then("the service keeps only the three latest messages")
def step_verify_pruned_messages(context: Any) -> None:
"""Ensure pruning retained the configured window."""
assert len(context.final_messages) == 3
contents = [msg.content for msg in context.final_messages]
assert contents[-1] == "ai 4"
assert "user 4" in contents[-2]
@then("the recent message requests respect their sizes")
def step_verify_recent_requests(context: Any) -> None:
"""Validate recent message helpers."""
assert len(context.recent_two) == 2
assert len(context.recent_all) == len(context.final_messages)
assert [msg.content for msg in context.recent_two] == ["user 4", "ai 4"]
@then("the service summary honors a 20 character limit")
def step_verify_summary_limit(context: Any) -> None:
"""Confirm summary respects the provided character cap."""
summary = context.memory_service.get_summary(max_chars=20)
context.summary = summary
lines = summary.splitlines()
assert 1 <= len(lines) <= 2
assert lines[-1].endswith("ai 4")
@then("memory variables expose history and counts")
def step_verify_memory_variables(context: Any) -> None:
"""Check memory variables include history, summary, and count."""
variables = context.memory_service.get_memory_variables()
assert variables["message_count"] == len(context.final_messages)
assert variables["chat_history"] == context.final_messages
assert isinstance(variables["chat_summary"], str)
assert context.final_messages[-1].content in variables["chat_summary"]
if context.summary:
assert (
context.summary.splitlines()[-1].split(":", 1)[-1].strip()
in variables["chat_summary"]
)
@given("a memory service with default settings")
def step_memory_service_default(context: Any) -> None:
"""Create default memory service without max limit."""
context.memory_service = MemoryService(session_id="default-session")
@when("I add direct message objects to the service")
def step_add_direct_messages(context: Any) -> None:
"""Add BaseMessage instances directly and capture summary/token count."""
context.memory_service.add_message(
HumanMessage(content="Alpha conversation segment")
)
context.memory_service.add_message(AIMessage(content="Beta response text"))
context.summary_before_clear = context.memory_service.get_summary()
context.token_count_before_clear = context.memory_service.get_token_count()
context.messages_before_clear = context.memory_service.get_messages().copy()
@when("I inspect the default conversation adapter")
def step_inspect_default_adapter(context: Any) -> None:
"""Grab default adapter payload."""
context.default_adapter = context.memory_service.conversation_memory
context.default_payload = context.default_adapter.load_memory_variables()
@when("I reset the max message limit to 1")
def step_reset_max_messages(context: Any) -> None:
"""Reduce max message count and trigger immediate pruning."""
context.memory_service.set_max_messages(1)
context.post_set_messages = list(context.memory_service.get_messages())
@when("I save a new interaction to enforce pruning")
def step_save_interaction_after_limit(context: Any) -> None:
"""Save interaction that should prune to configured limit."""
context.memory_service.save_context(
{"input": "latest user"},
{"output": "latest ai"},
)
context.pruned_messages = context.memory_service.get_messages()
@when("I clear the service memory")
def step_clear_service(context: Any) -> None:
"""Clear the memory service history."""
context.memory_service.clear()
@then("the service computed token counts and summaries before clearing")
def step_validate_pre_clear_metrics(context: Any) -> None:
"""Ensure summary and token counts were recorded before clearing."""
assert context.summary_before_clear
assert context.token_count_before_clear > 0
payload = context.default_payload[context.default_adapter.memory_key]
assert isinstance(payload, list)
assert len(context.post_set_messages) == 1
@then("the service history is empty after clearing")
def step_validate_cleared_history(context: Any) -> None:
"""History should be empty after clear."""
assert not context.memory_service.get_messages()
@then("a custom conversation adapter returns text values")
def step_custom_adapter_returns_text(context: Any) -> None:
"""Create custom adapter and ensure it returns string payloads."""
custom_adapter = context.memory_service.create_conversation_memory(
memory_key="custom_history",
input_key="prompt",
output_key="reply",
return_messages=False,
)
payload = custom_adapter.load_memory_variables()
assert payload["custom_history"] == ""
@given("a patched SQL chat history for memory service")
def step_patch_sql_history(context: Any) -> None:
"""Patch SQLChatMessageHistory to observe construction."""
context.sql_history_calls = []
context.sql_history_instance = InMemoryChatMessageHistory()
context.sql_temp_dir = tempfile.mkdtemp()
def factory(*args: Any, **kwargs: Any) -> InMemoryChatMessageHistory:
context.sql_history_calls.append((args, kwargs))
return context.sql_history_instance
sql_patch = patch(
"cleveragents.application.services.memory_service.SQLChatMessageHistory",
side_effect=factory,
)
context.sql_patch = sql_patch
sql_patch.start()
if hasattr(context, "_cleanup_handlers"):
context._cleanup_handlers.append(sql_patch.stop)
context._cleanup_handlers.append(
lambda: shutil.rmtree(context.sql_temp_dir, ignore_errors=True)
)
@when("I create a memory service with a connection string")
def step_create_sql_memory_service(context: Any) -> None:
"""Instantiate memory service pointing at SQL backend."""
db_path = Path(context.sql_temp_dir) / "memory.db"
context.connection_string = f"sqlite:///{db_path}"
context.sql_service = MemoryService(
session_id="sql-session",
connection_string=context.connection_string,
)
@then("the SQL history constructor receives the session and connection")
def step_validate_sql_constructor(context: Any) -> None:
"""Verify the patched SQL history constructor was invoked correctly."""
assert context.sql_history_calls
_, kwargs = context.sql_history_calls[0]
assert kwargs["session_id"] == "sql-session"
assert kwargs["connection_string"] == context.connection_string
assert kwargs.get("table_name") == "message_history"
@then("stored messages flow through the patched SQL history")
def step_validate_sql_storage(context: Any) -> None:
"""Ensure messages saved through the service reach the patched history."""
context.sql_service.save_context(
{"input": "persisted human"},
{"output": "persisted ai"},
)
messages = context.sql_history_instance.messages
assert len(messages) == 2
assert isinstance(messages[0], HumanMessage)
assert isinstance(messages[1], AIMessage)
assert messages[0].content == "persisted human"
assert messages[1].content == "persisted ai"