feat: 导入 MemRelay 初始源码
This commit is contained in:
@@ -0,0 +1,199 @@
|
||||
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"
|
||||
Reference in New Issue
Block a user