"""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