feat: 导入 MemRelay 初始源码
This commit is contained in:
@@ -0,0 +1,270 @@
|
||||
import time
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import memrelay.app as app_module
|
||||
from memrelay.app import create_app
|
||||
from memrelay.config import Settings
|
||||
from memrelay.prompting import build_agent_prompt
|
||||
from memrelay.schemas import FileScanResponse
|
||||
from tests.conftest import csrf_headers
|
||||
from tests.fakes import FakeBasicMemory
|
||||
|
||||
|
||||
def test_dashboard_settings_prompt_and_client_configs(initialized_client: TestClient) -> None:
|
||||
client = initialized_client
|
||||
client.app.state.basic_memory = FakeBasicMemory()
|
||||
dashboard = client.get("/api/v1/system/dashboard")
|
||||
assert dashboard.status_code == 200
|
||||
assert dashboard.json()["projects"] == 0
|
||||
assert len(dashboard.json()["activity_trend"]) == 14
|
||||
assert all(item["count"] == 0 for item in dashboard.json()["activity_trend"])
|
||||
assert dashboard.json()["basic_memory"]["semantic_search"]["status"] == "available"
|
||||
|
||||
updated = client.patch(
|
||||
"/api/v1/system/settings",
|
||||
json={
|
||||
"file_scan_interval_seconds": 0,
|
||||
"vault_sync_interval_seconds": 0,
|
||||
"context_max_chars": 12000,
|
||||
},
|
||||
headers=csrf_headers(client),
|
||||
)
|
||||
assert updated.status_code == 200
|
||||
assert updated.json()["context_max_chars"] == 12000
|
||||
|
||||
token = client.post(
|
||||
"/api/v1/tokens",
|
||||
json={"name": "Client config", "access_mode": "read_write"},
|
||||
headers=csrf_headers(client),
|
||||
).json()
|
||||
for name in ("codex", "cursor", "claude-code", "vscode", "generic"):
|
||||
config = client.get(
|
||||
f"/api/v1/system/client-config/{name}", params={"token_id": token["id"]}
|
||||
)
|
||||
assert config.status_code == 200
|
||||
assert "http://testserver/mcp" in config.json()["content"]
|
||||
assert token["token"] in config.json()["content"]
|
||||
|
||||
prompt = client.get("/api/v1/system/prompt").json()["content"]
|
||||
assert "context_get" in prompt
|
||||
assert "memrelay://guide" in prompt
|
||||
assert "memrelay://capabilities" in prompt
|
||||
assert "不把猜测" in prompt
|
||||
assert "memory_search` 默认使用 `current" in prompt
|
||||
assert "curated_document_get" in prompt
|
||||
assert "HTTP PUT" in prompt
|
||||
assert "file_upload_status" in prompt
|
||||
assert "download_url" in prompt
|
||||
assert "AI 自动整理" not in prompt
|
||||
assert "记忆版本历史" not in prompt
|
||||
english_prompt = client.get("/api/v1/system/prompt", params={"locale": "en-US"})
|
||||
assert "Do not save guesses" in english_prompt.json()["content"]
|
||||
assert "memrelay://guide" in english_prompt.json()["content"]
|
||||
assert "memrelay://capabilities" in english_prompt.json()["content"]
|
||||
assert "memory_search` defaults to `current" in english_prompt.json()["content"]
|
||||
assert "HTTP PUT" in english_prompt.json()["content"]
|
||||
assert "## AI curation" not in english_prompt.json()["content"]
|
||||
|
||||
|
||||
def test_agent_prompt_only_includes_enabled_optional_capabilities() -> None:
|
||||
base = build_agent_prompt(True, "zh-CN")
|
||||
assert "账号与凭证" in base
|
||||
assert "AI 自动整理" not in base
|
||||
assert "记忆版本历史" not in base
|
||||
|
||||
complete = build_agent_prompt(
|
||||
True,
|
||||
"zh-CN",
|
||||
curation_enabled=True,
|
||||
memory_git_enabled=True,
|
||||
)
|
||||
assert "AI 自动整理" in complete
|
||||
assert "curation_source_submit" in complete
|
||||
assert "记忆版本历史" in complete
|
||||
|
||||
disconnected = build_agent_prompt(False, "en-US")
|
||||
assert "The vault is not connected" in disconnected
|
||||
assert "Memory history" not in disconnected
|
||||
|
||||
|
||||
def test_password_change(initialized_client: TestClient) -> None:
|
||||
client = initialized_client
|
||||
wrong = client.post(
|
||||
"/api/v1/system/password",
|
||||
json={"current_password": "wrong", "new_password": "new-strong-password"},
|
||||
headers=csrf_headers(client),
|
||||
)
|
||||
assert wrong.status_code == 401
|
||||
changed = client.post(
|
||||
"/api/v1/system/password",
|
||||
json={
|
||||
"current_password": "strong-test-password",
|
||||
"new_password": "new-strong-password",
|
||||
},
|
||||
headers=csrf_headers(client),
|
||||
)
|
||||
assert changed.status_code == 204
|
||||
client.post("/api/v1/auth/logout", headers=csrf_headers(client))
|
||||
assert (
|
||||
client.post(
|
||||
"/api/v1/auth/login",
|
||||
json={"username": "admin", "password": "new-strong-password"},
|
||||
).status_code
|
||||
== 200
|
||||
)
|
||||
|
||||
|
||||
def test_dashboard_activity_trend_counts_new_memories(initialized_client: TestClient) -> None:
|
||||
client = initialized_client
|
||||
client.app.state.basic_memory = FakeBasicMemory()
|
||||
created = client.post(
|
||||
"/api/v1/memories",
|
||||
json={
|
||||
"request_id": "dashboard-trend",
|
||||
"scope": "global",
|
||||
"memory_type": "fact",
|
||||
"title": "Dashboard trend",
|
||||
"content": "A memory created for dashboard trend verification.",
|
||||
},
|
||||
headers=csrf_headers(client),
|
||||
)
|
||||
assert created.status_code == 201
|
||||
|
||||
dashboard = client.get("/api/v1/system/dashboard").json()
|
||||
assert dashboard["activity_trend"][-1]["count"] == 1
|
||||
assert sum(item["count"] for item in dashboard["activity_trend"]) == 1
|
||||
|
||||
|
||||
def test_dashboard_reports_unavailable_memory_dependency(initialized_client: TestClient) -> None:
|
||||
class UnavailableBasicMemory(FakeBasicMemory):
|
||||
async def health(self) -> dict[str, object]:
|
||||
return {
|
||||
"status": "unavailable",
|
||||
"error_code": "BASIC_MEMORY_TIMEOUT",
|
||||
"url": "memory://fake",
|
||||
}
|
||||
|
||||
client = initialized_client
|
||||
client.app.state.basic_memory = UnavailableBasicMemory()
|
||||
dashboard = client.get("/api/v1/system/dashboard")
|
||||
assert dashboard.status_code == 200
|
||||
capability = dashboard.json()["basic_memory"]["semantic_search"]
|
||||
assert capability == {
|
||||
"enabled": False,
|
||||
"status": "unavailable",
|
||||
"fallback": "text",
|
||||
"error_code": "BASIC_MEMORY_TIMEOUT",
|
||||
}
|
||||
|
||||
|
||||
def test_frontend_favicon_is_served_as_svg(tmp_path) -> None:
|
||||
frontend_dir = tmp_path / "frontend"
|
||||
frontend_dir.mkdir()
|
||||
(frontend_dir / "index.html").write_text("<html></html>", encoding="utf-8")
|
||||
(frontend_dir / "favicon.svg").write_text("<svg></svg>", encoding="utf-8")
|
||||
settings = Settings(
|
||||
data_dir=tmp_path / "data",
|
||||
frontend_dir=frontend_dir,
|
||||
public_url="http://testserver",
|
||||
initial_admin_username=None,
|
||||
initial_admin_password=None,
|
||||
)
|
||||
|
||||
with TestClient(create_app(settings)) as client:
|
||||
response = client.get("/favicon.svg")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.headers["content-type"].startswith("image/svg+xml")
|
||||
assert response.text == "<svg></svg>"
|
||||
|
||||
|
||||
def test_runtime_settings_are_restored_after_restart(tmp_path) -> None:
|
||||
data_dir = tmp_path / "persistent"
|
||||
settings = Settings(
|
||||
data_dir=data_dir,
|
||||
public_url="http://testserver",
|
||||
initial_admin_username=None,
|
||||
initial_admin_password=None,
|
||||
)
|
||||
with TestClient(create_app(settings)) as client:
|
||||
client.post(
|
||||
"/api/v1/auth/initialize",
|
||||
json={"username": "admin", "password": "strong-test-password"},
|
||||
)
|
||||
response = client.patch(
|
||||
"/api/v1/system/settings",
|
||||
json={"context_max_chars": 4321, "file_scan_interval_seconds": 0},
|
||||
headers=csrf_headers(client),
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
restarted_settings = Settings(
|
||||
data_dir=data_dir,
|
||||
public_url="http://testserver",
|
||||
initial_admin_username=None,
|
||||
initial_admin_password=None,
|
||||
)
|
||||
with TestClient(create_app(restarted_settings)) as restarted:
|
||||
login = restarted.post(
|
||||
"/api/v1/auth/login",
|
||||
json={"username": "admin", "password": "strong-test-password"},
|
||||
)
|
||||
assert login.status_code == 200
|
||||
restored = restarted.get("/api/v1/system/settings").json()
|
||||
assert restored["context_max_chars"] == 4321
|
||||
assert restored["file_scan_interval_seconds"] == 0
|
||||
|
||||
|
||||
def test_file_scan_interval_can_be_enabled_and_disabled_without_restart(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
calls: list[float] = []
|
||||
|
||||
def fake_scan(_db, _root, _on_change=None) -> FileScanResponse: # type: ignore[no-untyped-def]
|
||||
calls.append(time.monotonic())
|
||||
return FileScanResponse(added=0, updated=0, missing=0)
|
||||
|
||||
monkeypatch.setattr(app_module, "scan_files", fake_scan)
|
||||
settings = Settings(
|
||||
data_dir=tmp_path / "dynamic-interval",
|
||||
public_url="http://testserver",
|
||||
initial_admin_username=None,
|
||||
initial_admin_password=None,
|
||||
file_scan_interval_seconds=0,
|
||||
vault_sync_interval_seconds=0,
|
||||
)
|
||||
with TestClient(create_app(settings)):
|
||||
assert len(calls) == 1
|
||||
settings.file_scan_interval_seconds = 0.1 # type: ignore[assignment]
|
||||
deadline = time.monotonic() + 2
|
||||
while len(calls) < 2 and time.monotonic() < deadline:
|
||||
time.sleep(0.02)
|
||||
assert len(calls) >= 2
|
||||
|
||||
settings.file_scan_interval_seconds = 0
|
||||
time.sleep(0.2)
|
||||
stopped_count = len(calls)
|
||||
time.sleep(0.3)
|
||||
assert len(calls) == stopped_count
|
||||
|
||||
|
||||
def test_migration_prompt_variant(initialized_client: TestClient) -> None:
|
||||
client = initialized_client
|
||||
migration = client.get(
|
||||
"/api/v1/system/prompt", params={"variant": "migration"}
|
||||
).json()["content"]
|
||||
assert "MemRelay" in migration
|
||||
assert "project_resolve_or_create" in migration
|
||||
assert "checkpoint_save" in migration
|
||||
assert "memory_save" in migration
|
||||
|
||||
english = client.get(
|
||||
"/api/v1/system/prompt", params={"variant": "migration", "locale": "en-US"}
|
||||
).json()["content"]
|
||||
assert "Migrate an existing project" in english
|
||||
|
||||
workflow = client.get("/api/v1/system/prompt").json()["content"]
|
||||
assert "capabilities_get" in workflow
|
||||
assert "Migrate an existing project" not in workflow
|
||||
Reference in New Issue
Block a user