Files
2026-07-22 13:48:46 +08:00

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",
},
},
]