252 lines
7.1 KiB
Python
252 lines
7.1 KiB
Python
# pyright: reportAssignmentType=false, reportArgumentType=false, reportIndexIssue=false
|
|
|
|
import os
|
|
import sys
|
|
from typing import Any, cast
|
|
|
|
import pytest
|
|
|
|
|
|
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
|
if PROJECT_ROOT not in sys.path:
|
|
sys.path.insert(0, PROJECT_ROOT)
|
|
|
|
|
|
from context.context import RedisFileContextBackend
|
|
from context.schemas import Message
|
|
from web_interface.models import (
|
|
ChatMessageContentImage,
|
|
ChatMessageContentImageUrl,
|
|
ChatMessageContentText,
|
|
)
|
|
from web_interface.routers.conversation_router import get_conversation_history
|
|
|
|
|
|
class _FakeRedis:
|
|
_store: dict[str, object] = {}
|
|
|
|
def __init__(self, *args, **kwargs):
|
|
pass
|
|
|
|
@classmethod
|
|
def reset(cls) -> None:
|
|
cls._store = {}
|
|
|
|
def lpush(self, key: str, value: str) -> None:
|
|
values = cast(list[str], self._store.setdefault(key, []))
|
|
values.insert(0, value)
|
|
|
|
def rpush(self, key: str, value: str) -> None:
|
|
values = cast(list[str], self._store.setdefault(key, []))
|
|
values.append(value)
|
|
|
|
def lrange(self, key: str, start: int, end: int):
|
|
values = list(cast(list[str], self._store.get(key, [])))
|
|
if end == -1:
|
|
end = len(values) - 1
|
|
return values[start : end + 1]
|
|
|
|
def llen(self, key: str) -> int:
|
|
values = self._store.get(key, [])
|
|
assert isinstance(values, list)
|
|
return len(values)
|
|
|
|
def delete(self, key: str) -> None:
|
|
self._store.pop(key, None)
|
|
|
|
def set(self, key: str, value: str) -> None:
|
|
self._store[key] = value
|
|
|
|
def get(self, key: str):
|
|
return self._store.get(key)
|
|
|
|
def ltrim(self, key: str, start: int, end: int) -> None:
|
|
values = list(cast(list[str], self._store.get(key, [])))
|
|
if end == -1:
|
|
trimmed = values[start:]
|
|
else:
|
|
trimmed = values[start : end + 1]
|
|
self._store[key] = trimmed
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_full_history_persists_even_when_working_memory_summarizes(
|
|
tmp_path, monkeypatch
|
|
):
|
|
import context.context as context_module
|
|
|
|
_FakeRedis.reset()
|
|
monkeypatch.setattr(context_module.redis, "Redis", _FakeRedis)
|
|
|
|
context_file = tmp_path / "ctx_conv-1.json"
|
|
backend = RedisFileContextBackend(
|
|
context_id="conv-1",
|
|
llm_interface=None,
|
|
max_history_length=2,
|
|
auto_summarize_trigger=2,
|
|
file_path=str(context_file),
|
|
)
|
|
cast(Any, backend).llm_interface = object()
|
|
|
|
async def fake_summarize(messages):
|
|
return "summary"
|
|
|
|
backend._summarize_func = fake_summarize
|
|
|
|
await backend.store_message(Message(role="user", content="first"))
|
|
await backend.store_message(Message(role="assistant", content="second"))
|
|
await backend.store_message(Message(role="user", content="third"))
|
|
|
|
assert [message.content for message in backend.retrieve_messages()] == ["third"]
|
|
assert [message.content for message in backend.retrieve_full_messages()] == [
|
|
"first",
|
|
"second",
|
|
"third",
|
|
]
|
|
|
|
await backend.persist()
|
|
|
|
_FakeRedis.reset()
|
|
restored = RedisFileContextBackend(
|
|
context_id="conv-1",
|
|
llm_interface=None,
|
|
max_history_length=2,
|
|
auto_summarize_trigger=2,
|
|
file_path=str(context_file),
|
|
)
|
|
await restored.restore()
|
|
|
|
assert [message.content for message in restored.retrieve_messages()] == ["third"]
|
|
assert [message.content for message in restored.retrieve_full_messages()] == [
|
|
"first",
|
|
"second",
|
|
"third",
|
|
]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_conversation_history_endpoint_returns_full_persisted_history():
|
|
archived_messages = [
|
|
Message(role="user", content="first"),
|
|
Message(role="assistant", content="second"),
|
|
Message(role="user", content="third"),
|
|
]
|
|
working_messages = [archived_messages[-1]]
|
|
|
|
class _FakeContext:
|
|
def retrieve_messages(self):
|
|
return working_messages
|
|
|
|
def retrieve_full_messages(self):
|
|
return archived_messages
|
|
|
|
class _FakeConversation:
|
|
uuid = "conversation-1"
|
|
context = _FakeContext()
|
|
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, tb):
|
|
return False
|
|
|
|
class _FakeConversationManager:
|
|
def get_conversation(self, conversation_id):
|
|
return _FakeConversation() if conversation_id == "conversation-1" else None
|
|
|
|
class _FakeState:
|
|
conversation_manager = _FakeConversationManager()
|
|
|
|
response = cast(
|
|
Any,
|
|
await get_conversation_history(
|
|
conversation_id="conversation-1",
|
|
limit=None,
|
|
state=cast(Any, _FakeState()),
|
|
),
|
|
)
|
|
|
|
assert response["total_messages"] == 3
|
|
assert [message["content"] for message in response["messages"]] == [
|
|
"first",
|
|
"second",
|
|
"third",
|
|
]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_large_auto_summarize_trigger_effectively_disables_summary(
|
|
tmp_path, monkeypatch
|
|
):
|
|
import context.context as context_module
|
|
|
|
_FakeRedis.reset()
|
|
monkeypatch.setattr(context_module.redis, "Redis", _FakeRedis)
|
|
|
|
context_file = tmp_path / "ctx_conv-2.json"
|
|
backend = RedisFileContextBackend(
|
|
context_id="conv-2",
|
|
llm_interface=None,
|
|
max_history_length=2,
|
|
auto_summarize_trigger=999999,
|
|
file_path=str(context_file),
|
|
)
|
|
cast(Any, backend).llm_interface = object()
|
|
|
|
summarize_call_count = 0
|
|
|
|
async def fake_summarize(messages):
|
|
nonlocal summarize_call_count
|
|
summarize_call_count += 1
|
|
return "summary"
|
|
|
|
backend._summarize_func = fake_summarize
|
|
|
|
await backend.store_message(Message(role="user", content="first"))
|
|
await backend.store_message(Message(role="assistant", content="second"))
|
|
await backend.store_message(Message(role="user", content="third"))
|
|
|
|
assert summarize_call_count == 0
|
|
assert backend.get_summary() is None
|
|
assert [message.content for message in backend.retrieve_messages()] == [
|
|
"second",
|
|
"third",
|
|
]
|
|
|
|
|
|
def test_context_message_accepts_web_multimodal_models():
|
|
message = Message(
|
|
role="user",
|
|
content=[
|
|
ChatMessageContentText.model_validate(
|
|
{"type": "text", "text": "Treat this as an assembly"}
|
|
),
|
|
ChatMessageContentImage.model_validate(
|
|
{
|
|
"type": "image_url",
|
|
"image_url": ChatMessageContentImageUrl.model_validate(
|
|
{
|
|
"url": "data:image/png;base64,abcd",
|
|
"local_path": "/tmp/query_image_001.png",
|
|
}
|
|
),
|
|
}
|
|
),
|
|
],
|
|
)
|
|
|
|
assert isinstance(message.content, list)
|
|
|
|
payload = message.model_dump()
|
|
assert payload["content"] == [
|
|
{"type": "text", "text": "Treat this as an assembly"},
|
|
{
|
|
"type": "image_url",
|
|
"image_url": {
|
|
"url": "data:image/png;base64,abcd",
|
|
"detail": None,
|
|
"local_path": "/tmp/query_image_001.png",
|
|
},
|
|
},
|
|
]
|