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"