deerflow-code/offline-backend-20260512/backend/tests/test_wiki_package_roundtrip.py
2026-09-07 18:24:55 +08:00

430 lines
17 KiB
Python

import base64
import json
import sys
from io import BytesIO
from types import SimpleNamespace
import pytest
from fastapi import FastAPI, UploadFile
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
from app.gateway.routers import wiki_packages
from deerflow.assistant_knowledge.archive import PackageArchive
from deerflow.assistant_knowledge.package_io import write_package
from deerflow.config.llmwiki_config import LocalWikiEmbeddingConfig
from deerflow.integrations.weknora.local_index import embedding as embedding_module
from deerflow.integrations.weknora.local_index.chunking import chunk_wiki_page
from deerflow.integrations.weknora.local_index.embedding import StrictWikiEmbeddingClient
from deerflow.integrations.weknora.local_index.normalize import normalize_wiki_page
from deerflow.integrations.weknora.local_index.vector_codec import encode_vector
from deerflow.persistence.assistant_knowledge import AssistantKnowledgeRepository
from deerflow.persistence.base import Base
from deerflow.persistence.llmwiki_index.memory import MemoryLlmWikiIndexStore
@pytest.mark.asyncio
async def test_metadata_vectors_and_links_survive_roundtrip_without_encoding(tmp_path, monkeypatch):
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
store = AssistantKnowledgeRepository(async_sessionmaker(engine, expire_on_commit=False))
blob, _ = encode_vector([0.3, 0.5, 0.7])
pages = [
{
"slug": "concept/cooling",
"title": "散热",
"page_type": "concept",
"content": "散热管理用于降低设备温度。参见 [[entity/fan]]",
"summary": "设备散热",
"tags": ["设备", "维护"],
"category_path": ["运维", "硬件"],
"aliases": ["降温"],
"out_links": ["entity/fan"],
"in_links": [],
},
{"slug": "entity/fan", "title": "风扇", "page_type": "entity", "content": "风扇通过气流帮助散热,参见 [[concept/cooling]]。", "aliases": [], "out_links": ["concept/cooling"], "in_links": ["concept/cooling"]},
]
vectors = [
{
"wiki_slug": p["slug"],
"section_index": 0,
"section_content": p["content"],
"content_hash": "test",
"embedding_model": "test-model",
"embedding_fingerprint": "test-fp",
"embedding_dimensions": 3,
"vector_blob_base64": base64.b64encode(blob).decode(),
}
for p in pages
]
original = tmp_path / "original.jsonl"
entities = [{"id": p["slug"], "name": p["title"], "type": p["page_type"], "source_wiki_slug": p["slug"]} for p in pages]
relations = [{"source_entity_id": pages[0]["slug"], "target_entity_id": pages[1]["slug"], "source_wiki_slug": pages[0]["slug"], "predicate": "references"}]
write_package(original, [
("source", {"provider": "weknora", "package_version": 1}),
*[("wiki_page", p) for p in pages],
*[("vector", v) for v in vectors],
*[("entity", e) for e in entities],
*[("relation", r) for r in relations],
("chunk", {"id": "raw-chunk", "content": "原始文档分段不是 Wiki 页面"}),
])
archive = PackageArchive(original)
base = await store.initialize(created_by="admin")
job = await store.import_package(package=archive, source_type="weknora", source_key="source", source_name="来源", trigger="test", created_by="admin")
assert job["counts"]["vectors"] == 2
assert job["counts"]["wiki_pages"] == 2
assert job["counts"]["chunk_pages"] == 1
imported = await store.get_page(pages[0]["slug"])
for field in ("category_path", "tags", "aliases", "out_links"):
assert imported[field] == pages[0][field]
assert imported["content_markdown"] == pages[0]["content"]
# The global base may receive the same Wiki through multiple libraries.
# Only the current article's import supplies vectors; older contributions
# remain auditable but must not duplicate/stale the restored index.
await store.import_package(package=archive, source_type="weknora", source_key="second-source", source_name="第二来源", trigger="test", created_by="admin")
exported = tmp_path / "export.jsonl"
write_package(exported, [record async for record in store.iter_export_records(base["id"])])
restored_archive = PackageArchive(exported)
assert {v["vector_blob_base64"] for v in restored_archive["vectors"]} == {base64.b64encode(blob).decode()}
assert sorted((v["wiki_slug"], v["section_index"]) for v in restored_archive["vectors"]) == sorted((v["wiki_slug"], v["section_index"]) for v in vectors)
assert all(r["source_wiki_slug"] in {p["slug"] for p in pages} for r in restored_archive["relations"])
index = MemoryLlmWikiIndexStore()
received = []
class Remote:
async def create_wiki_folder(self, kb, *, name, parent_id=""):
return {"id": f"{parent_id}/{name}"}
async def create_wiki_page(self, kb, page):
received.append(page)
return page
async def rebuild_wiki_links(self, kb):
pass
class Embedding:
dimensions = 3
fingerprint = "test-fp"
async def embed_query(self, query):
raise AssertionError("Existing package must not be encoded")
app = FastAPI()
app.state.llmwiki_index_store = index
app.state.llmwiki_embedding = Embedding()
app.state.llmwiki_vector_cache = None
app.state.config = None
monkeypatch.setattr(wiki_packages, "get_resolved_llmwiki_runtime", lambda c: None)
monkeypatch.setattr(wiki_packages, "build_weknora_client", lambda c: Remote())
monkeypatch.setattr(wiki_packages, "_package_root", lambda: tmp_path)
from app.gateway import knowledge_transfer
async def sync(app, mapping):
return {"status": "completed"}
monkeypatch.setattr(knowledge_transfer, "sync_to_global", sync)
await wiki_packages.restore_package(app, {"id": "target", "weknora_id": "remote"}, "job-test", exported)
for raw in received:
assert raw["content"] == next(p["content"] for p in pages if p["slug"] == raw["slug"])
local = await index.get_page("target", raw["slug"])
assert local["content_hash"] == normalize_wiki_page("target", raw)["content_hash"]
snapshot = await index.load_vector_snapshot("target", "test-fp")
assert len(snapshot["rows"]) == 2
assert all(r["vector_blob"] == blob for r in snapshot["rows"])
class SearchEmbedding:
dimensions = 3
fingerprint = "test-fp"
async def embed_query(self, query):
return [1, 0, 0]
found = await store.search_pages(base_ids=[base["id"]], query="没有关键词重合的查询", embedding_client=SearchEmbedding())
assert found and all(r["retrieval_mode"] == "vector" for r in found)
SearchEmbedding.fingerprint = "different-model-same-dimension"
assert await store.search_pages(base_ids=[base["id"]], query="没有关键词重合的查询", embedding_client=SearchEmbedding()) == []
archive.close()
restored_archive.close()
await engine.dispose()
@pytest.mark.asyncio
async def test_local_encoder_is_bounded_cpu_and_does_not_need_http(monkeypatch, tmp_path):
observed = {}
class LocalModel:
@staticmethod
def list_supported_models():
return [{"model": "BAAI/bge-m3"}]
def __init__(self, **kwargs):
observed.update(kwargs)
def embed(self, texts, *, batch_size):
observed["batch_size"] = batch_size
return [[1.0, *([0.0] * 1023)] for _ in texts]
monkeypatch.setitem(sys.modules, "fastembed", SimpleNamespace(TextEmbedding=LocalModel))
(tmp_path / "onnx").mkdir()
(tmp_path / "onnx" / "model.onnx").touch()
config = LocalWikiEmbeddingConfig(
provider="local",
model="bge-embedding-m3",
dimensions=1024,
batch_size=2,
local_threads=2,
local_model_path=str(tmp_path),
local_process_isolation=False,
)
client = StrictWikiEmbeddingClient(config)
assert len(await client.embed_query("本地查询")) == 1024
assert observed["model_name"] == "BAAI/bge-m3"
assert observed["providers"] == ["CPUExecutionProvider"]
assert observed["local_files_only"] is True
assert observed["specific_model_path"] == str(tmp_path)
assert observed["batch_size"] == observed["threads"] == 2
assert client._client is None
assert client.local_encoder_status()["state"] == "ready"
assert client.local_encoder_status()["model_files_ready"] is True
cold_client = StrictWikiEmbeddingClient(config)
status = cold_client.start_local_encoder()
assert status["state"] == "starting"
assert status["running"] is True
await cold_client._local_warmup_task
assert cold_client.local_encoder_status()["state"] == "ready"
monkeypatch.delenv("DEERFLOW_BGE_M3_MODEL_PATH", raising=False)
missing_path = StrictWikiEmbeddingClient(
LocalWikiEmbeddingConfig(provider="local", model="bge-embedding-m3", dimensions=1024),
)
missing_path.start_local_encoder()
await missing_path._local_warmup_task
assert missing_path.local_encoder_status()["state"] == "failed"
assert "目录未配置" in str(missing_path.local_encoder_status()["last_error"])
@pytest.mark.asyncio
async def test_local_encoder_uses_offline_isolated_worker(monkeypatch, tmp_path):
(tmp_path / "onnx").mkdir()
(tmp_path / "onnx" / "model.onnx").touch()
observed: dict = {}
class FakeStdout:
def __init__(self):
self.lines = [json.dumps({"status": "ready"}) + "\n"]
def readline(self):
return self.lines.pop(0) if self.lines else ""
class FakeProcess:
def __init__(self):
self.stdout = FakeStdout()
self.returncode = None
self.stdin = SimpleNamespace(write=self.write, flush=lambda: None)
def write(self, line):
payload = json.loads(line)
if payload.get("command") == "shutdown":
self.returncode = 0
return
self.stdout.lines.append(
json.dumps(
{
"id": payload["id"],
"status": "ok",
"vectors": [[1.0, *([0.0] * 1023)] for _ in payload["texts"]],
}
)
+ "\n"
)
def poll(self):
return self.returncode
def wait(self, timeout=None):
return self.returncode
def terminate(self):
self.returncode = -15
def kill(self):
self.returncode = -9
process = FakeProcess()
def fake_popen(command, **kwargs):
observed.update({"command": command, **kwargs})
return process
monkeypatch.setattr(embedding_module.subprocess, "Popen", fake_popen)
client = StrictWikiEmbeddingClient(
LocalWikiEmbeddingConfig(
provider="local",
model="bge-embedding-m3",
dimensions=1024,
local_model_path=str(tmp_path),
)
)
assert len(await client.embed_query("隔离编码")) == 1024
assert "deerflow.integrations.weknora.local_index.embedding_worker" in observed["command"]
assert observed["env"]["HF_HUB_OFFLINE"] == "1"
assert observed["env"]["PYTHONIOENCODING"] == "utf-8"
assert client.local_encoder_status()["process_isolation"] is True
await client.close()
assert process.poll() == 0
@pytest.mark.asyncio
async def test_builtin_skill_without_ownership_record_can_be_distilled(monkeypatch, tmp_path):
from starlette.requests import Request
from app.gateway.routers import skill_knowledge
async def admin(_):
return "admin", True
class Mappings:
async def get_authorized(self, *args, **kwargs):
return {"name": "测试库", "weknora_id": "remote"}
class Jobs:
async def create_job(self, **kwargs):
return kwargs
app = FastAPI()
app.state.config = None
app.state.llmwiki_store = Mappings()
app.state.skill_knowledge_store = Jobs()
# Deliberately no skill_store ownership row: built-ins live on disk.
monkeypatch.setattr(skill_knowledge, "_actor", admin)
monkeypatch.setattr(skill_knowledge, "_skill_directory", lambda *args: tmp_path)
body = skill_knowledge.CreateSyncJobRequest(skill_names=["builtin"], targets=[{"target_type": "weknora", "target_mode": "wiki", "target_id": "mapping"}])
job = await skill_knowledge.create_sync_job(Request({"type": "http", "app": app}), body)
assert job["skill_names"] == ["builtin"]
@pytest.mark.asyncio
async def test_regular_user_can_distill_visible_skill_only_to_writable_wiki(monkeypatch, tmp_path):
from starlette.requests import Request
from app.gateway.routers import skill_knowledge
async def regular_user(_):
return "user-1", False
class Skills:
async def get_visible(self, name, user_id):
assert user_id == "user-1"
return {"name": name} if name == "visible-skill" else None
async def list_favorite_skill_names(self, user_id):
assert user_id == "user-1"
return []
class Mappings:
async def get_authorized(self, mapping_id, user_id, *, write, is_admin):
assert (mapping_id, user_id, write, is_admin) == ("mapping", "user-1", True, False)
return {"name": "我的 Wiki", "weknora_id": "remote"}
class Jobs:
async def create_job(self, **kwargs):
return kwargs
app = FastAPI()
app.state.config = None
app.state.skill_store = Skills()
app.state.llmwiki_store = Mappings()
app.state.skill_knowledge_store = Jobs()
monkeypatch.setattr(skill_knowledge, "_actor", regular_user)
monkeypatch.setattr(skill_knowledge, "_skill_directory", lambda *args: tmp_path)
monkeypatch.setattr(
skill_knowledge,
"get_or_new_skill_storage",
lambda **_: SimpleNamespace(get_skills_root_path=lambda: tmp_path),
)
body = skill_knowledge.CreateSyncJobRequest(
skill_names=["visible-skill"],
targets=[{"target_type": "weknora", "target_mode": "wiki", "target_id": "mapping"}],
review_mode="required",
)
job = await skill_knowledge.create_sync_job(Request({"type": "http", "app": app}), body)
assert job["created_by"] == "user-1"
assert job["review_mode"] == "off"
def test_short_wiki_summary_and_body_still_have_a_vector_section():
sections = chunk_wiki_page({"title": "风扇", "summary": "通过气流帮助散热。", "content_md": "# 风扇\n用于降低设备温度。"}, max_chars=1000, overlap=100, min_chars=80, title_prefix_enabled=True)
assert len(sections) == 1
assert "降低设备温度" in sections[0].embedding_text
@pytest.mark.asyncio
async def test_admin_can_upload_and_activate_offline_encoder(monkeypatch, tmp_path):
from starlette.requests import Request
from app.gateway.routers import llmwiki_index
from deerflow.config.system_settings import SystemSettings
async def admin(_):
return "admin", True
class Embedding:
def __init__(self):
self.model_path = ""
async def reset_local_encoder(self, *, model_path=None):
if model_path is not None:
self.model_path = model_path
def local_encoder_status(self):
return {
"provider": "local",
"model": "bge-embedding-m3",
"state": "idle",
"model_path": self.model_path,
"model_files_ready": bool(self.model_path),
}
saved = []
monkeypatch.setattr(llmwiki_index, "_actor", admin)
monkeypatch.setattr(llmwiki_index, "runtime_home", lambda: tmp_path)
monkeypatch.setattr(llmwiki_index, "load_system_settings", SystemSettings)
monkeypatch.setattr(llmwiki_index, "save_system_settings", lambda value: saved.append(value) or value)
embedding = Embedding()
app = FastAPI()
app.state.config = SimpleNamespace(
llmwiki=SimpleNamespace(
local_wiki_index=SimpleNamespace(embedding=SimpleNamespace(provider="local"))
)
)
app.state.llmwiki_embedding = embedding
request = Request({"type": "http", "headers": [], "app": app})
files = [
UploadFile(filename="model.onnx", file=BytesIO(b"onnx")),
UploadFile(filename="config.json", file=BytesIO(b"{}")),
]
result = await llmwiki_index.upload_local_encoder(
request,
files=files,
paths=["onnx/model.onnx", "config.json"],
)
active = tmp_path / "models" / "bge-m3"
assert (active / "onnx" / "model.onnx").read_bytes() == b"onnx"
assert result["uploaded_files"] == 2
assert saved[-1].local_wiki_encoder.model_path == str(active.resolve())
assert embedding.model_path == str(active.resolve())
def test_encoder_upload_rejects_parent_traversal():
from app.gateway.routers import llmwiki_index
with pytest.raises(Exception) as exc_info:
llmwiki_index._safe_relative_upload_path("../model.onnx")
assert getattr(exc_info.value, "status_code", None) == 422