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:
@@ -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
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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,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
|
||||
|
||||
Reference in New Issue
Block a user