200 lines
7.0 KiB
Python
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"
|