forked from HAL9000/cleveragents-core
358 lines
12 KiB
Python
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"
|