430 lines
17 KiB
Python
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
|