271 lines
9.7 KiB
Python
271 lines
9.7 KiB
Python
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
|