fix(backend): 修复全面审阅发现的核心问题

修复索引首次失败回滚、Vault 扫描边界、Markdown 代码围栏和 UTF-16 Citation 偏移。

收紧 Plugin 权限与 JSON Schema 校验,补齐 Note Move、Task、Attachment 和 Transcript Tool。

接入 OpenAI SSE 与 Ollama JSONL 真流式输出,修正 Provider PATCH 语义并限制运行时内存保留。

新增对应回归测试,后端测试增至 62 项。
This commit is contained in:
2026-08-28 09:56:07 +08:00
parent 167ae24796
commit 36bc1022f1
26 changed files with 1531 additions and 99 deletions
+62 -2
View File
@@ -2,8 +2,25 @@ import asyncio
from app.main import health, service_status
from app.routes import get_index_status, list_notes, list_plugins, list_providers, list_skills
from app.routes import create_provider, delete_provider, get_provider, update_provider
from app.contracts import ProviderCreateRequest, ProviderType, ProviderUpdateRequest
from app.routes import (
create_provider,
create_task,
delete_provider,
delete_task,
get_provider,
get_task,
list_tasks,
update_provider,
update_task,
)
from app.contracts import (
ProviderCreateRequest,
ProviderType,
ProviderUpdateRequest,
TaskCreateRequest,
TaskStatus,
TaskUpdateRequest,
)
def test_health() -> None:
@@ -79,3 +96,46 @@ def test_provider_configuration_lifecycle() -> None:
assert fetched.provider_type == ProviderType.ollama
assert disabled.enabled is False
assert deleted.resource_id == created.provider_id
def test_provider_patch_can_clear_nullable_fields() -> None:
created = asyncio.run(
create_provider(
ProviderCreateRequest(
provider_type=ProviderType.ollama,
name="Clearable",
base_url="http://127.0.0.1:11434",
default_model="qwen",
credential_id="unused",
)
)
)
try:
cleared = asyncio.run(
update_provider(
created.provider_id,
ProviderUpdateRequest(
base_url=None, default_model=None, credential_id=None
),
)
)
assert cleared.base_url is None
assert cleared.default_model is None
assert cleared.credential_id is None
finally:
asyncio.run(delete_provider(created.provider_id))
def test_task_lifecycle_is_persistent() -> None:
created = asyncio.run(create_task(TaskCreateRequest(title="审阅修复")))
fetched = asyncio.run(get_task(created.task_id))
updated = asyncio.run(
update_task(created.task_id, TaskUpdateRequest(status=TaskStatus.done))
)
listed = asyncio.run(list_tasks(limit=20, offset=0))
deleted = asyncio.run(delete_task(created.task_id))
assert fetched.title == "审阅修复"
assert updated.status == TaskStatus.done
assert any(item.task_id == created.task_id for item in listed.items)
assert deleted.resource_id == created.task_id
+125
View File
@@ -13,6 +13,7 @@ from app.contracts import (
)
from app.extensions import ExtensionError
from app.services import note_service
from app.config import get_settings
def run(coroutine):
@@ -172,3 +173,127 @@ tools: [plugin.not-installed]
with pytest.raises(ExtensionError) as exc:
container.skills.enable("missing-tool")
assert exc.value.code == "SKILL_DEPENDENCY_MISSING"
def test_plugin_permissions_must_be_known_and_granted(tmp_path) -> None:
package = tmp_path / "write-plugin"
package.mkdir()
(package / "plugin.yaml").write_text(
"""
id: write-plugin
name: Write Plugin
version: 1.0.0
permissions: [notes.write]
contributes:
tools: [plugin.write]
backend:
type: internal_rpc
transport: none
""".strip(),
encoding="utf-8",
)
(package / "tools.yaml").write_text(
"""
tools:
- name: plugin.write
description: permission test
permission: notes.write
handler: echo
parameters:
type: object
properties: {text: {type: string}}
required: [text]
""".strip(),
encoding="utf-8",
)
container = build_container()
installed = container.plugins.install(package)
assert installed.status == "permission_required"
with pytest.raises(ExtensionError) as exc:
container.plugins.enable("write-plugin")
assert exc.value.code == "PLUGIN_PERMISSION_REQUIRED"
granted = container.plugins.set_permissions("write-plugin", ["notes.write"])
enabled = container.plugins.enable("write-plugin")
assert granted.granted_permissions == ["notes.write"]
assert enabled.status == "ready"
def test_plugin_rejects_unknown_permissions_and_invalid_schema(tmp_path) -> None:
unknown = tmp_path / "unknown-permission"
unknown.mkdir()
(unknown / "plugin.yaml").write_text(
"""
id: unknown-permission
name: Unknown
version: 1.0.0
permissions: [notes.wirte]
""".strip(),
encoding="utf-8",
)
container = build_container()
with pytest.raises(ExtensionError) as exc:
container.plugins.install(unknown)
assert exc.value.code == "EXTENSION_PERMISSION_INVALID"
malformed = tmp_path / "malformed-schema"
malformed.mkdir()
(malformed / "plugin.yaml").write_text(
"""
id: malformed-schema
name: Malformed
version: 1.0.0
contributes:
tools: [bad.schema]
""".strip(),
encoding="utf-8",
)
(malformed / "tools.yaml").write_text(
"""
tools:
- name: bad.schema
description: invalid schema
handler: echo
parameters:
type: object
properties: []
""".strip(),
encoding="utf-8",
)
with pytest.raises(ExtensionError) as exc:
container.plugins.install(malformed)
assert exc.value.code == "PLUGIN_TOOL_SCHEMA_INVALID"
def test_attachment_and_transcription_tools_use_host_storage() -> None:
async def scenario() -> None:
root = get_settings().attachments_path
root.mkdir(parents=True, exist_ok=True)
(root / "meeting.txt").write_text("会议转写内容", encoding="utf-8")
container = build_container()
attachment = await container.tools.execute(
ToolCall(
tool_call_id="call_attachment",
name="attachments.read",
arguments={"attachment_id": "meeting.txt"},
),
ToolExecutionContext(run_id="run_attachment"),
)
transcription = await container.tools.execute(
ToolCall(
tool_call_id="call_transcription",
name="audio.transcribe",
arguments={"attachment_id": "meeting.txt"},
),
ToolExecutionContext(run_id="run_transcription"),
)
assert attachment.success is True
assert attachment.output["content"] == "会议转写内容"
assert transcription.success is True
assert transcription.output["status"] == "completed"
assert transcription.output["text"] == "会议转写内容"
run(scenario())
+74 -1
View File
@@ -3,7 +3,14 @@ import json
import httpx
from app.contracts import Message, MessageRole, ModelRequest, ToolCall, ToolDefinition
from app.contracts import (
Message,
MessageRole,
ModelEventType,
ModelRequest,
ToolCall,
ToolDefinition,
)
from app.providers.ollama import OllamaProvider
from app.providers.openai_compatible import OpenAICompatibleProvider
@@ -17,6 +24,10 @@ def run(coroutine):
return asyncio.run(coroutine)
async def collect(stream):
return [event async for event in stream]
def test_openai_compatible_maps_tool_call_and_credentials() -> None:
captured: dict = {}
@@ -152,3 +163,65 @@ def test_ollama_maps_models_and_completion() -> None:
assert turn.text == "local answer"
assert turn.input_tokens == 5
assert turn.output_tokens == 2
def test_openai_compatible_streams_incremental_sse() -> None:
def handler(request: httpx.Request) -> httpx.Response:
payload = json.loads(request.content)
assert payload["stream"] is True
body = "\n".join(
[
'data: {"choices":[{"delta":{"content":"hel"}}]}',
'data: {"choices":[{"delta":{"content":"lo"},"finish_reason":"stop"}]}',
'data: {"choices":[],"usage":{"prompt_tokens":2,"completion_tokens":1}}',
"data: [DONE]",
"",
]
)
return httpx.Response(200, text=body)
provider = OpenAICompatibleProvider(
base_url="https://provider.test/v1",
credential_id=None,
credentials=StaticCredentials(),
transport=httpx.MockTransport(handler),
)
request = ModelRequest(
provider_id="test", model="model",
messages=[Message(role=MessageRole.user, content="hello")],
)
events = run(collect(provider.stream(request)))
assert [item.data["text"] for item in events if item.event == ModelEventType.text_delta] == [
"hel", "lo"
]
assert events[-1].event == ModelEventType.done
assert [item.sequence for item in events] == list(range(len(events)))
def test_ollama_streams_incremental_jsonl() -> None:
def handler(request: httpx.Request) -> httpx.Response:
payload = json.loads(request.content)
assert payload["stream"] is True
body = "\n".join(
[
'{"message":{"content":""},"done":false}',
'{"message":{"content":""},"done":false}',
'{"message":{"content":""},"done":true,"prompt_eval_count":3,"eval_count":2}',
"",
]
)
return httpx.Response(200, text=body)
provider = OllamaProvider(transport=httpx.MockTransport(handler))
request = ModelRequest(
provider_id="ollama", model="qwen",
messages=[Message(role=MessageRole.user, content="hello")],
)
events = run(collect(provider.stream(request)))
assert [item.data["text"] for item in events if item.event == ModelEventType.text_delta] == [
"", ""
]
assert events[-2].event == ModelEventType.usage
assert events[-1].event == ModelEventType.done
+77 -1
View File
@@ -77,6 +77,28 @@ def test_block_ids_are_stable() -> None:
assert all(b.block_id.startswith("blk_") for b in p1.blocks)
def test_code_fence_does_not_create_fake_headings() -> None:
markdown = '# Real Heading\n\n```python\n# code comment\nprint("x")\n```\n'
parsed = parse_note(
markdown=markdown, file_path="code.md", folder="",
tags=None, created_at=_dt(), updated_at=_dt(),
)
assert all(block.heading_path == ["Real Heading"] for block in parsed.blocks)
assert any("# code comment" in block.content for block in parsed.blocks)
def test_offsets_use_browser_compatible_utf16_units() -> None:
markdown = "# \U0001f600\n\nbody"
parsed = parse_note(
markdown=markdown, file_path="emoji.md", folder="",
tags=None, created_at=_dt(), updated_at=_dt(),
)
body = next(block for block in parsed.blocks if block.content == "body")
assert body.start_offset == len(markdown[: markdown.index("body")].encode("utf-16-le")) // 2
def test_tokens_split_cjk_bigrams_and_match_query() -> None:
toks = tokens("向量检索")
assert "" in toks and "" in toks
@@ -454,7 +476,6 @@ def test_patch_partial_content_no_orphan_vectors(vault) -> None:
)
)
assert ids("vec_blocks") == ids("blocks")
asyncio.run(
note_service.update_note(note.note_id, markdown="# 标题\n\n段落一改了。\n\n新增段落。")
)
@@ -462,6 +483,22 @@ def test_patch_partial_content_no_orphan_vectors(vault) -> None:
assert ids("vec_blocks") == ids("blocks")
def test_move_note_preserves_id_and_updates_indexed_path(vault) -> None:
from app.services import note_service
created = asyncio.run(
note_service.create_note(
title="可移动", markdown="# 标题\n\n正文", folder="原目录", tags=[]
)
)
moved = asyncio.run(note_service.move_note(created.note_id, folder="新目录"))
assert moved.note_id == created.note_id
assert moved.file_path == "新目录/可移动.md"
assert not (vault / "原目录" / "可移动.md").exists()
assert (vault / "新目录" / "可移动.md").exists()
assert asyncio.run(note_service.get_note(created.note_id)).file_path == moved.file_path
def test_fts_metadata_filter_recalls_beyond_candidate_pool(vault) -> None:
"""metadata 过滤不能受候选池截断影响:目标块排在 50 名之外也应被召回(审阅 #4)。"""
from app.retrieval.engine import engine
@@ -524,3 +561,42 @@ def test_rebuild_failure_restores_old_index(vault, monkeypatch) -> None:
asyncio.run(index_service.rebuild(IndexRebuildRequest(scope="all")))
assert repository.stats() == before # 旧索引已恢复,无半成品
def test_first_rebuild_failure_removes_partial_database(vault, monkeypatch) -> None:
"""首次启动没有旧库时,失败也不能留下已经写入的部分索引。"""
from app.services import index_service
_write_vault(
vault,
{"a.md": "# A\n\nfirst", "b.md": "# B\n\nsecond"},
)
real_index = index_service.index_note
calls = {"count": 0}
async def fail_on_second(parsed):
calls["count"] += 1
if calls["count"] == 2:
raise RuntimeError("injected first-rebuild failure")
await real_index(parsed)
monkeypatch.setattr(index_service, "index_note", fail_on_second)
with pytest.raises(RuntimeError):
asyncio.run(index_service.rebuild(IndexRebuildRequest(scope="all")))
assert not get_settings().db_path.exists()
def test_rebuild_preserves_task_note_links(vault) -> None:
from app.services import index_service, note_service, task_service
note = asyncio.run(
note_service.create_note(
title="任务关联", markdown="# 任务关联", folder="", tags=[]
)
)
task = task_service.create_task(title="跟进", note_id=note.note_id)
asyncio.run(index_service.rebuild(IndexRebuildRequest(scope="all")))
assert task_service.get_task(task.task_id).note_id == note.note_id