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

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