test: use SQLite sessions in core memory (#39066)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
This commit is contained in:
Asuka Minato 2026-07-27 19:15:08 +09:00 committed by GitHub
parent 9fd1fea58c
commit 609ddf9e01
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -1,11 +1,19 @@
"""Comprehensive unit tests for core/memory/token_buffer_memory.py"""
"""Comprehensive SQLite-backed tests for token-buffer memory."""
from collections.abc import Iterator
from dataclasses import dataclass
from datetime import UTC, datetime, timedelta
from decimal import Decimal
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from sqlalchemy import Engine, event
from sqlalchemy.orm import Session
from core.memory import token_buffer_memory as memory_module
from core.memory.token_buffer_memory import TokenBufferMemory
from graphon.file import FileTransferMethod, FileType
from graphon.model_runtime.entities import (
AssistantPromptMessage,
ImagePromptMessageContent,
@ -13,13 +21,44 @@ from graphon.model_runtime.entities import (
TextPromptMessageContent,
UserPromptMessage,
)
from models.model import AppMode
from models.base import TypeBase
from models.enums import ConversationFromSource, CreatorUserRole, MessageFileBelongsTo
from models.model import AppMode, Message, MessageFile
from models.workflow import Workflow, WorkflowType
# ---------------------------------------------------------------------------
# Helpers / shared fixtures
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class Database:
"""Typed SQLite binding plus executed SQL for query-count assertions."""
engine: Engine
session: Session
statements: list[tuple[str, object]]
@pytest.fixture
def database(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Database]:
TypeBase.metadata.create_all(
sqlite_engine,
tables=[Message.__table__, MessageFile.__table__, Workflow.__table__],
)
statements: list[tuple[str, object]] = []
def record_statement(_connection, _cursor, statement, parameters, _context, _executemany) -> None:
statements.append((statement, parameters))
event.listen(sqlite_engine, "before_cursor_execute", record_statement)
with Session(sqlite_engine, expire_on_commit=False) as session:
database = Database(engine=sqlite_engine, session=session, statements=statements)
monkeypatch.setattr(memory_module, "db", database)
yield database
event.remove(sqlite_engine, "before_cursor_execute", record_statement)
def _make_conversation(mode: AppMode = AppMode.CHAT) -> MagicMock:
"""Return a minimal Conversation mock."""
conv = MagicMock()
@ -47,6 +86,73 @@ def _make_message(answer: str = "hello", answer_tokens: int = 5) -> MagicMock:
return msg
def _persist_message(
database: Database,
conversation_id: str,
*,
query: str = "user query",
answer: str = "hello",
answer_tokens: int = 5,
created_at: datetime | None = None,
workflow_run_id: str | None = None,
) -> Message:
message = Message(
id=str(uuid4()),
app_id="app-1",
conversation_id=conversation_id,
_inputs={},
query=query,
message={},
message_unit_price=Decimal(0),
answer=answer,
answer_tokens=answer_tokens,
answer_unit_price=Decimal(0),
currency="USD",
from_source=ConversationFromSource.API,
workflow_run_id=workflow_run_id,
created_at=created_at or datetime.now(UTC).replace(tzinfo=None),
)
database.session.add(message)
database.session.commit()
return message
def _persist_message_file(
database: Database,
message: Message,
*,
belongs_to: MessageFileBelongsTo | None,
) -> MessageFile:
message_file = MessageFile(
message_id=message.id,
type=FileType.IMAGE,
transfer_method=FileTransferMethod.REMOTE_URL,
created_by_role=CreatorUserRole.ACCOUNT,
created_by="account-1",
belongs_to=belongs_to,
url="https://example.com/image.png",
)
database.session.add(message_file)
database.session.commit()
return message_file
def _persist_workflow(database: Database, *, workflow_id: str) -> Workflow:
workflow = Workflow(
id=workflow_id,
tenant_id="tenant-1",
app_id="app-1",
type=WorkflowType.CHAT,
version="1",
graph="{}",
features="{}",
created_by="account-1",
)
database.session.add(workflow)
database.session.commit()
return workflow
# ===========================================================================
# Tests for __init__ and workflow_run_repo property
# ===========================================================================
@ -61,25 +167,25 @@ class TestInit:
assert mem.model_instance is mi
assert mem._workflow_run_repo is None
def test_workflow_run_repo_is_created_lazily(self):
def test_workflow_run_repo_is_created_lazily(self, database: Database):
conv = _make_conversation()
mi = _make_model_instance()
mem = TokenBufferMemory(conversation=conv, model_instance=mi)
mock_repo = MagicMock()
with (
patch("core.memory.token_buffer_memory.sessionmaker") as mock_sm,
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.DifyAPIRepositoryFactory.create_api_workflow_run_repository",
return_value=mock_repo,
),
):
mock_db.engine = MagicMock()
with patch(
"core.memory.token_buffer_memory.DifyAPIRepositoryFactory.create_api_workflow_run_repository",
return_value=mock_repo,
) as repository_factory:
repo = mem.workflow_run_repo
assert repo is mock_repo
assert mem._workflow_run_repo is mock_repo
session_factory = repository_factory.call_args.args[0]
with session_factory() as session:
assert isinstance(session, Session)
assert session.get_bind() is database.engine
def test_workflow_run_repo_cached_after_first_access(self):
conv = _make_conversation()
mi = _make_model_instance()
@ -410,7 +516,7 @@ class TestBuildPromptMessageWithFiles:
)
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW])
def test_workflow_mode_workflow_not_found_raises(self, mode):
def test_workflow_mode_workflow_not_found_raises(self, mode, database: Database):
"""Raises ValueError when Workflow lookup returns None."""
conv = _make_conversation(mode)
conv.app = MagicMock()
@ -422,22 +528,17 @@ class TestBuildPromptMessageWithFiles:
mem._workflow_run_repo = MagicMock()
mem._workflow_run_repo.get_workflow_run_by_id.return_value = mock_workflow_run
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
):
mock_db.session.scalar.return_value = None # workflow not found
with pytest.raises(ValueError, match="Workflow not found"):
mem._build_prompt_message_with_files(
message_files=[],
text_content="text",
message=_make_message(),
app_record=MagicMock(),
is_user_message=True,
)
with pytest.raises(ValueError, match="Workflow not found"):
mem._build_prompt_message_with_files(
message_files=[],
text_content="text",
message=_make_message(),
app_record=MagicMock(),
is_user_message=True,
)
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW])
def test_workflow_mode_success_no_files_user(self, mode):
def test_workflow_mode_success_no_files_user(self, mode, database: Database):
"""Happy path: workflow mode, no message files → plain UserPromptMessage."""
conv = _make_conversation(mode)
conv.app = MagicMock()
@ -445,22 +546,16 @@ class TestBuildPromptMessageWithFiles:
mock_workflow_run = MagicMock()
mock_workflow_run.workflow_id = str(uuid4())
mock_workflow = MagicMock()
mock_workflow.features_dict = {}
workflow = _persist_workflow(database, workflow_id=mock_workflow_run.workflow_id)
mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance())
mem._workflow_run_repo = MagicMock()
mem._workflow_run_repo.get_workflow_run_by_id.return_value = mock_workflow_run
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
with patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
):
mock_db.session.scalar.return_value = mock_workflow
result = mem._build_prompt_message_with_files(
message_files=[],
text_content="wf text",
@ -471,6 +566,7 @@ class TestBuildPromptMessageWithFiles:
assert isinstance(result, UserPromptMessage)
assert result.content == "wf text"
assert database.session.get(Workflow, workflow.id) is workflow
# ------------------------------------------------------------------
# Invalid mode
@ -498,417 +594,140 @@ class TestBuildPromptMessageWithFiles:
class TestGetHistoryPromptMessages:
"""Tests for get_history_prompt_messages."""
"""Tests for persisted history retrieval, file batching, and pruning."""
def _make_memory(self, mode: AppMode = AppMode.CHAT) -> TokenBufferMemory:
conv = _make_conversation(mode)
conv.app = MagicMock()
return TokenBufferMemory(conversation=conv, model_instance=_make_model_instance())
def test_returns_empty_when_no_messages(self):
def test_returns_empty_when_no_messages(self, database: Database) -> None:
assert self._make_memory().get_history_prompt_messages() == []
def test_skips_newest_message_without_answer(self, database: Database) -> None:
mem = self._make_memory()
with patch("core.memory.token_buffer_memory.db") as mock_db:
mock_db.session.scalars.return_value.all.return_value = []
result = mem.get_history_prompt_messages()
assert result == []
message = _persist_message(database, mem.conversation.id, answer="", answer_tokens=0)
def test_skips_first_message_without_answer(self):
"""The newest message (index 0 after extraction) without answer and tokens==0 is skipped."""
assert mem.get_history_prompt_messages() == []
assert database.session.get(Message, message.id) is message
def test_message_with_answer_returns_user_and_assistant_prompts(self, database: Database) -> None:
mem = self._make_memory()
_persist_message(database, mem.conversation.id, query="My query", answer="My answer", answer_tokens=10)
msg_no_answer = _make_message(answer="", answer_tokens=0)
msg_no_answer.parent_message_id = None # ensures extract_thread_messages returns it
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg_no_answer],
),
):
mock_db.session.scalars.return_value.all.side_effect = [
[msg_no_answer], # first call: messages query
[], # second call: user files query (never hit, but safe)
]
result = mem.get_history_prompt_messages()
assert result == []
def test_message_with_answer_not_skipped(self):
"""A message with a non-empty answer is NOT popped."""
mem = self._make_memory()
msg = _make_message(answer="some answer", answer_tokens=10)
msg.parent_message_id = None
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
# user files query → empty; assistant files query → empty
mock_db.session.scalars.return_value.all.return_value = []
result = mem.get_history_prompt_messages()
assert len(result) == 2 # one user + one assistant
def test_message_limit_default_is_500(self):
"""When message_limit is None the stmt is limited to 500."""
mem = self._make_memory()
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.select") as mock_select,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]),
):
mock_stmt = MagicMock()
mock_select.return_value.where.return_value.order_by.return_value = mock_stmt
mock_stmt.limit.return_value = mock_stmt
mock_db.session.scalars.return_value.all.return_value = []
mem.get_history_prompt_messages(message_limit=None)
mock_stmt.limit.assert_called_with(500)
def test_message_limit_clipped_to_500(self):
"""A message_limit > 500 is clamped to 500."""
mem = self._make_memory()
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.select") as mock_select,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]),
):
mock_stmt = MagicMock()
mock_select.return_value.where.return_value.order_by.return_value = mock_stmt
mock_stmt.limit.return_value = mock_stmt
mock_db.session.scalars.return_value.all.return_value = []
mem.get_history_prompt_messages(message_limit=9999)
mock_stmt.limit.assert_called_with(500)
def test_message_limit_positive_used(self):
"""A positive message_limit < 500 is used as-is."""
mem = self._make_memory()
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.select") as mock_select,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]),
):
mock_stmt = MagicMock()
mock_select.return_value.where.return_value.order_by.return_value = mock_stmt
mock_stmt.limit.return_value = mock_stmt
mock_db.session.scalars.return_value.all.return_value = []
mem.get_history_prompt_messages(message_limit=10)
mock_stmt.limit.assert_called_with(10)
def test_message_limit_zero_uses_default(self):
"""message_limit=0 triggers the else branch → default 500."""
mem = self._make_memory()
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.select") as mock_select,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]),
):
mock_stmt = MagicMock()
mock_select.return_value.where.return_value.order_by.return_value = mock_stmt
mock_stmt.limit.return_value = mock_stmt
mock_db.session.scalars.return_value.all.return_value = []
mem.get_history_prompt_messages(message_limit=0)
mock_stmt.limit.assert_called_with(500)
def test_user_files_cause_build_with_files_call(self):
"""When user_files is non-empty _build_prompt_message_with_files is invoked."""
mem = self._make_memory()
msg = _make_message()
msg.parent_message_id = None
mock_user_file = MagicMock()
mock_user_file.message_id = msg.id # must match so batched grouping keys it to this message
mock_user_prompt = UserPromptMessage(content="from build")
mock_assistant_prompt = AssistantPromptMessage(content="answer")
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
# messages query
r.all.return_value = [msg]
elif call_count["n"] == 1:
# user files
r.all.return_value = [mock_user_file]
else:
# assistant files
r.all.return_value = []
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch.object(
mem,
"_build_prompt_message_with_files",
side_effect=[mock_user_prompt, mock_assistant_prompt],
) as mock_build,
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages()
assert mock_build.call_count >= 1
# First call should be user message
first_call_kwargs = mock_build.call_args_list[0][1]
assert first_call_kwargs["is_user_message"] is True
def test_assistant_files_cause_build_with_files_call(self):
"""When assistant_files is non-empty, build is called with is_user_message=False."""
mem = self._make_memory()
msg = _make_message()
msg.parent_message_id = None
mock_assistant_file = MagicMock()
mock_assistant_file.message_id = msg.id # must match so batched grouping keys it to this message
mock_user_prompt = UserPromptMessage(content="query")
mock_assistant_prompt = AssistantPromptMessage(content="built")
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
r.all.return_value = [msg]
elif call_count["n"] == 1:
r.all.return_value = [] # no user files
else:
r.all.return_value = [mock_assistant_file]
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch.object(
mem,
"_build_prompt_message_with_files",
return_value=mock_assistant_prompt,
) as mock_build,
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages()
mock_build.assert_called_once()
call_kwargs = mock_build.call_args[1]
assert call_kwargs["is_user_message"] is False
def test_message_files_loaded_with_constant_query_count(self):
"""Regression guard against N+1: message files must be batch-loaded.
Regardless of the number of messages in the thread, file loading must use a
constant number of queries (1 messages query + 2 batched file queries),
never 2 queries per message.
"""
mem = self._make_memory()
messages = [_make_message() for _ in range(5)]
for m in messages:
m.parent_message_id = None
scalars_calls = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
# First call returns the thread messages; the batched file queries return none.
r.all.return_value = messages if scalars_calls["n"] == 0 else []
scalars_calls["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=messages),
patch("core.memory.token_buffer_memory.FileUploadConfigManager.convert", return_value=None),
):
mock_db.session.scalars.side_effect = scalars_side_effect
mem.get_history_prompt_messages()
# 1 (messages) + 2 (batched user/assistant files) = 3, independent of message count.
# Before this fix it would have been 1 + 2 * 5 = 11 (an N+1 pattern).
assert scalars_calls["n"] == 3
def test_token_pruning_removes_oldest_messages(self):
"""If tokens exceed limit, oldest messages are removed until within limit."""
conv = _make_conversation()
conv.app = MagicMock()
# Model returns tokens that decrease only after removing pairs
token_values = [3000, 1500] # first call over limit, second within
mi = MagicMock()
mi.get_llm_num_tokens.side_effect = token_values
mem = TokenBufferMemory(conversation=conv, model_instance=mi)
msg = _make_message()
msg.parent_message_id = None
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
r.all.return_value = [msg]
else:
r.all.return_value = []
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages(max_token_limit=2000)
# After pruning, we should have fewer than the 2 initial messages
assert len(result) <= 1
def test_token_pruning_stops_at_single_message(self):
"""Pruning stops when only 1 message remains (to prevent empty list)."""
conv = _make_conversation()
conv.app = MagicMock()
# Always over limit
mi = MagicMock()
mi.get_llm_num_tokens.return_value = 99999
mem = TokenBufferMemory(conversation=conv, model_instance=mi)
msg = _make_message()
msg.parent_message_id = None
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
r.all.return_value = [msg]
else:
r.all.return_value = []
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages(max_token_limit=1)
# At least 1 message should remain
assert len(result) >= 1
def test_no_pruning_when_within_limit(self):
"""When tokens ≤ limit, no pruning occurs."""
mem = self._make_memory()
mem.model_instance.get_llm_num_tokens.return_value = 50 # well under default 2000
msg = _make_message()
msg.parent_message_id = None
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
r.all.return_value = [msg]
else:
r.all.return_value = []
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages(max_token_limit=2000)
assert len(result) == 2 # user + assistant
def test_plain_user_and_assistant_messages_returned(self):
"""Without files, plain UserPromptMessage and AssistantPromptMessage appear."""
mem = self._make_memory()
msg = _make_message(answer="My answer")
msg.query = "My query"
msg.parent_message_id = None
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
r.all.return_value = [msg]
else:
r.all.return_value = []
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages()
result = mem.get_history_prompt_messages()
assert len(result) == 2
user_msg, ai_msg = result
assert isinstance(user_msg, UserPromptMessage)
assert user_msg.content == "My query"
assert isinstance(ai_msg, AssistantPromptMessage)
assert ai_msg.content == "My answer"
assert isinstance(result[0], UserPromptMessage)
assert result[0].content == "My query"
assert isinstance(result[1], AssistantPromptMessage)
assert result[1].content == "My answer"
def test_history_is_conversation_scoped(self, database: Database) -> None:
mem = self._make_memory()
_persist_message(database, mem.conversation.id, answer="visible")
_persist_message(database, "other-conversation", answer="hidden")
result = mem.get_history_prompt_messages()
assert [prompt.content for prompt in result] == ["user query", "visible"]
@pytest.mark.parametrize(
("message_limit", "expected_limit"),
[(None, 500), (9999, 500), (10, 10), (0, 500)],
)
def test_message_limit_is_applied_to_executable_query(
self,
database: Database,
message_limit: int | None,
expected_limit: int,
) -> None:
mem = self._make_memory()
before = len(database.statements)
mem.get_history_prompt_messages(message_limit=message_limit)
statements = database.statements[before:]
assert len(statements) == 1
sql, parameters = statements[0]
assert "LIMIT" in sql
assert expected_limit in parameters
@pytest.mark.parametrize(
("belongs_to", "is_user_message"),
[
(MessageFileBelongsTo.USER, True),
(None, True),
(MessageFileBelongsTo.ASSISTANT, False),
],
)
def test_message_files_use_persisted_ownership(
self,
database: Database,
belongs_to: MessageFileBelongsTo | None,
is_user_message: bool,
) -> None:
mem = self._make_memory()
message = _persist_message(database, mem.conversation.id)
message_file = _persist_message_file(database, message, belongs_to=belongs_to)
built_prompt = (
UserPromptMessage(content="built user")
if is_user_message
else AssistantPromptMessage(content="built assistant")
)
with patch.object(mem, "_build_prompt_message_with_files", return_value=built_prompt) as build_prompt:
result = mem.get_history_prompt_messages()
build_prompt.assert_called_once()
assert build_prompt.call_args.kwargs["message_files"] == [message_file]
assert build_prompt.call_args.kwargs["is_user_message"] is is_user_message
assert built_prompt in result
def test_message_files_are_batch_loaded_with_constant_query_count(self, database: Database) -> None:
mem = self._make_memory()
base_time = datetime.now(UTC).replace(tzinfo=None)
messages = [
_persist_message(
database,
mem.conversation.id,
query=f"query-{index}",
answer=f"answer-{index}",
created_at=base_time + timedelta(seconds=index),
)
for index in range(5)
]
before = len(database.statements)
with patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=messages):
result = mem.get_history_prompt_messages()
selects = [sql for sql, _ in database.statements[before:] if sql.lstrip().upper().startswith("SELECT")]
assert len(selects) == 3
assert len(result) == 10
@pytest.mark.parametrize(
("token_values", "max_token_limit", "expected_length"),
[
([3000, 1500], 2000, 1),
([99999, 99999], 1, 1),
([50], 2000, 2),
],
)
def test_token_pruning_uses_persisted_history(
self,
database: Database,
token_values: list[int],
max_token_limit: int,
expected_length: int,
) -> None:
mem = self._make_memory()
mem.model_instance.get_llm_num_tokens.side_effect = token_values
_persist_message(database, mem.conversation.id)
result = mem.get_history_prompt_messages(max_token_limit=max_token_limit)
assert len(result) == expected_length
# ===========================================================================