167 lines
5.4 KiB
Python
167 lines
5.4 KiB
Python
"""Unit tests for the AI 改写 prompt builder (app.gateway.routers.writing).
|
|
|
|
Focus: the 仿写 (imitate) action injects the reference sample into the prompt,
|
|
and existing actions are unaffected.
|
|
"""
|
|
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
|
|
import pytest
|
|
|
|
from app.gateway.routers.writing import (
|
|
_REFERENCE_MAX_CHARS,
|
|
RewriteRequest,
|
|
SelectionInfo,
|
|
_release_document_rewrite,
|
|
_reserve_document_rewrite,
|
|
build_rewrite_messages,
|
|
)
|
|
|
|
|
|
def _req(**kwargs) -> RewriteRequest:
|
|
base = {
|
|
"selection": SelectionInfo(**{"from": 0, "to": 4, "text": "原始文字"}),
|
|
}
|
|
base.update(kwargs)
|
|
return RewriteRequest(**base)
|
|
|
|
|
|
def test_imitate_injects_reference_text():
|
|
body = _req(action="imitate", reference_text="这是一段范文,节奏明快,排比有力。")
|
|
action, messages = build_rewrite_messages(body)
|
|
|
|
assert action == "imitate"
|
|
human = messages[1].content
|
|
assert "范文" in human
|
|
assert "这是一段范文,节奏明快,排比有力。" in human
|
|
assert "原始文字" in human
|
|
|
|
|
|
def test_imitate_without_reference_has_no_reference_block():
|
|
body = _req(action="imitate", reference_text=" ")
|
|
action, messages = build_rewrite_messages(body)
|
|
|
|
assert action == "imitate"
|
|
# No reference → the 范文 block is not appended
|
|
assert "【范文】" not in messages[1].content
|
|
|
|
|
|
def test_imitate_reference_is_truncated():
|
|
long_ref = "甲" * (_REFERENCE_MAX_CHARS + 500)
|
|
body = _req(action="imitate", reference_text=long_ref)
|
|
_, messages = build_rewrite_messages(body)
|
|
|
|
human = messages[1].content
|
|
# The reference is capped — the 500 extra chars are dropped
|
|
assert human.count("甲") == _REFERENCE_MAX_CHARS
|
|
|
|
|
|
def test_polish_action_unaffected_by_reference_field():
|
|
# A stray reference_text on a non-imitate action must be ignored
|
|
body = _req(action="polish", reference_text="不应出现的范文")
|
|
action, messages = build_rewrite_messages(body)
|
|
|
|
assert action == "polish"
|
|
assert "范文" not in messages[1].content
|
|
assert "不应出现的范文" not in messages[1].content
|
|
|
|
|
|
def test_unknown_action_falls_back_to_polish():
|
|
body = _req(action="not-a-real-action")
|
|
action, _ = build_rewrite_messages(body)
|
|
assert action == "polish"
|
|
|
|
|
|
def test_imitate_follow_up_preserves_reference_and_previous_result():
|
|
body = _req(
|
|
action="imitate",
|
|
reference_text="范文样例",
|
|
previous_result="上一版仿写结果",
|
|
follow_up_instruction="再正式一点",
|
|
)
|
|
_, messages = build_rewrite_messages(body)
|
|
|
|
# system, human(original w/ reference), ai(previous), human(follow-up)
|
|
assert len(messages) == 4
|
|
assert "范文样例" in messages[1].content
|
|
assert messages[2].content == "上一版仿写结果"
|
|
assert messages[3].content == "再正式一点"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_full_document_rewrite_reservation_rejects_a_second_live_request(tmp_path: Path):
|
|
"""Two streams must never generate competing commits for one real file."""
|
|
path = tmp_path / "report.md"
|
|
path.write_text("# report", encoding="utf-8")
|
|
|
|
assert await _reserve_document_rewrite(path) is True
|
|
try:
|
|
assert await _reserve_document_rewrite(path) is False
|
|
finally:
|
|
await _release_document_rewrite(path)
|
|
|
|
assert await _reserve_document_rewrite(path) is True
|
|
await _release_document_rewrite(path)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_full_document_rewrite_keeps_committed_file_when_undo_snapshot_fails(
|
|
monkeypatch,
|
|
tmp_path: Path,
|
|
):
|
|
"""Undo persistence is best effort and must not restore a valid rewrite."""
|
|
from app.gateway.routers import writing
|
|
|
|
path = tmp_path / "report.md"
|
|
path.write_text("# 原文\n\n正文", encoding="utf-8")
|
|
|
|
class _Model:
|
|
async def astream(self, *_args, **_kwargs):
|
|
yield "# 改写后\n\n新正文"
|
|
|
|
class _BrokenVersionStore:
|
|
async def create(self, **_kwargs):
|
|
raise RuntimeError("version store unavailable")
|
|
|
|
class _Config:
|
|
models: list[object] = []
|
|
|
|
@staticmethod
|
|
def get_model_config(_name):
|
|
return None
|
|
|
|
async def resolve_path(*_args):
|
|
return path
|
|
|
|
async def current_user(*_args):
|
|
return "user-a"
|
|
|
|
async def requirements(*_args, **_kwargs):
|
|
yield "保留事实。"
|
|
|
|
monkeypatch.setattr(writing, "aresolve_thread_virtual_path", resolve_path)
|
|
monkeypatch.setattr(writing, "get_current_user", current_user)
|
|
monkeypatch.setattr(writing, "create_chat_model", lambda **_kwargs: _Model())
|
|
monkeypatch.setattr(writing, "_astream_rewrite_requirements", requirements)
|
|
monkeypatch.setattr(writing, "_extract_chunk_parts", lambda chunk: ("", chunk))
|
|
response = await writing.rewrite_document(
|
|
body=writing.DocumentRewriteRequest(
|
|
threadId="thread-a",
|
|
path="/mnt/user-data/report.md",
|
|
instruction="统一表达风格",
|
|
),
|
|
request=SimpleNamespace(
|
|
_deerflow_test_bypass_auth=True,
|
|
app=SimpleNamespace(state=SimpleNamespace(document_rewrite_version_store=_BrokenVersionStore())),
|
|
),
|
|
config=_Config(),
|
|
)
|
|
|
|
events = "".join([frame async for frame in response.body_iterator])
|
|
|
|
assert path.read_text(encoding="utf-8") == "# 改写后\n\n新正文"
|
|
assert "event: committed" in events
|
|
assert '"versionSnapshotFailed": true' in events
|
|
assert "event: error" not in events
|