first commit
This commit is contained in:
@@ -0,0 +1,251 @@
|
||||
# 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",
|
||||
},
|
||||
},
|
||||
]
|
||||
Reference in New Issue
Block a user