Files
MemRelay/backend/tests/test_concurrency.py
T

200 lines
7.0 KiB
Python

from __future__ import annotations
import asyncio
from typing import Any
import httpx
import pytest
from fastapi.testclient import TestClient
from fastmcp import Client
from fastmcp.client.transports import StreamableHttpTransport
from sqlalchemy import func, select
from memrelay.database import database_health
from memrelay.models import (
CurationDependencyState,
CurationJob,
CurationSchedule,
CurationSettings,
FileRecord,
MemoryReference,
)
from tests.conftest import csrf_headers
from tests.fakes import FakeBasicMemory
def _mcp_transport(app: Any, token: str) -> StreamableHttpTransport:
def factory(**kwargs: Any) -> httpx.AsyncClient:
headers = dict(kwargs.get("headers") or {})
headers["Authorization"] = f"Bearer {token}"
return httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
headers=headers,
timeout=kwargs.get("timeout"),
follow_redirects=True,
)
return StreamableHttpTransport(
"http://testserver/mcp",
auth=token,
httpx_client_factory=factory,
)
def _result_data(result: Any) -> dict[str, Any]:
data = result.data
return data.model_dump(mode="json") if hasattr(data, "model_dump") else data
def test_twenty_concurrent_web_and_mcp_writes_keep_sqlite_consistent(
initialized_client: TestClient,
monkeypatch: pytest.MonkeyPatch,
) -> None:
client = initialized_client
fake = FakeBasicMemory()
client.app.state.basic_memory = fake
client.app.state.curation.basic_memory = fake
client.app.state.git_manager.basic_memory = fake
project_response = client.post(
"/api/v1/projects",
json={"name": "Concurrent workspace"},
headers=csrf_headers(client),
)
assert project_response.status_code == 201, project_response.text
project_id = project_response.json()["id"]
token_response = client.post(
"/api/v1/tokens",
json={"name": "Concurrent MCP clients", "access_mode": "read_write"},
headers=csrf_headers(client),
)
assert token_response.status_code == 201, token_response.text
token = token_response.json()["token"]
csrf = csrf_headers(client)
with client.app.state.session_factory() as db:
settings = db.get(CurationSettings, 1)
dependency = db.get(CurationDependencyState, 1)
assert settings is not None and dependency is not None
settings.enabled = True
dependency.status = "available"
db.add(
CurationSchedule(
name="Concurrent quiet merge",
trigger_type="workspace_quiet",
scope="project",
enabled=True,
inheritance="instance",
recurrence_json='{"mode":"quiet","seconds":0}',
)
)
db.commit()
async def complete_curation_job(job_id: str, claim_owner: str, _lease: int) -> None:
client.app.state.curation._complete_job(
job_id,
claim_owner,
{"document_count": 0, "no_changes": True},
)
monkeypatch.setattr(client.app.state.curation, "_run_job", complete_curation_job)
async def web_memory(index: int) -> dict[str, Any]:
response = await asyncio.to_thread(
client.post,
"/api/v1/memories",
json={
"request_id": f"concurrent-web-memory-{index}",
"scope": "global",
"memory_type": "fact",
"title": f"Concurrent web memory {index}",
"content": f"Web write {index}",
},
headers=csrf,
)
assert response.status_code == 201, response.text
return response.json()
async def web_directory(index: int) -> dict[str, Any]:
response = await asyncio.to_thread(
client.post,
"/api/v1/files/directories",
json={"path": f"concurrency/web-{index}"},
headers=csrf,
)
assert response.status_code == 201, response.text
return response.json()
async def mcp_memory(index: int) -> dict[str, Any]:
async with Client(_mcp_transport(client.app, token)) as writer:
result = await writer.call_tool(
"memory_save",
{
"request_id": f"concurrent-mcp-memory-{index}",
"scope": "project",
"project_id": project_id,
"memory_type": "experience",
"title": f"Concurrent MCP memory {index}",
"content": f"MCP write {index}",
},
)
return _result_data(result)
async def mcp_checkpoint(index: int) -> dict[str, Any]:
async with Client(_mcp_transport(client.app, token)) as writer:
result = await writer.call_tool(
"checkpoint_save",
{
"project_id": project_id,
"request_id": f"concurrent-checkpoint-{index}",
"content": f"Completed concurrent checkpoint {index}",
},
)
return _result_data(result)
async def background_curation() -> int:
processed = 0
for _ in range(200):
processed += int(await client.app.state.curation.process_once())
await asyncio.sleep(0.002)
return processed
async def exercise() -> tuple[list[dict[str, Any]], int]:
operations = [web_memory(index) for index in range(5)]
operations.extend(web_directory(index) for index in range(5))
operations.extend(mcp_memory(index) for index in range(5))
operations.extend(mcp_checkpoint(index) for index in range(5))
worker = asyncio.create_task(background_curation())
results = list(await asyncio.gather(*operations))
processed = await worker
while await client.app.state.curation.process_once():
processed += 1
return results, processed
results, processed_jobs = asyncio.run(exercise())
assert len(results) == 20
assert len({item["id"] for item in results}) == 20
assert processed_jobs >= 1
with client.app.state.session_factory() as db:
memories = int(db.scalar(select(func.count(MemoryReference.id))) or 0)
directories = int(
db.scalar(select(func.count(FileRecord.id)).where(FileRecord.is_directory.is_(True)))
or 0
)
checkpoints = int(
db.scalar(
select(func.count(MemoryReference.id)).where(
MemoryReference.memory_type == "checkpoint"
)
)
or 0
)
curation_jobs = db.scalars(select(CurationJob)).all()
health = database_health(client.app.state.engine, client.app.state.settings.database_path)
assert memories == 15
assert directories == 6
assert checkpoints == 5
assert curation_jobs and all(item.status == "completed" for item in curation_jobs)
assert health["integrity"] == "ok"
assert health["journal_mode"] == "wal"