"""Writing rewrite API — SSE streaming endpoint. POST /api/writing/rewrite Returns text/event-stream with frames: event: thinking data: {"chunk": "..."} event: text data: {"chunk": "..."} event: done data: {} event: error data: {"message": "..."} """ from __future__ import annotations import asyncio import hashlib import json import logging import os import re import uuid from pathlib import Path from urllib.parse import quote from fastapi import APIRouter, Depends, HTTPException, Query, Request from fastapi.responses import Response, StreamingResponse from langchain_core.messages import AIMessage, HumanMessage, SystemMessage from pydantic import BaseModel, Field from app.gateway.authz import require_permission from app.gateway.deps import get_config, get_current_user from app.gateway.document_rewrite_pipeline import ( DOCUMENT_REWRITE_SYSTEM_PROMPT as _DOCUMENT_REWRITE_SYSTEM_PROMPT, ) from app.gateway.document_rewrite_pipeline import ( DOCUMENT_STYLE_INSTRUCTIONS as _DOCUMENT_STYLE_INSTRUCTIONS, ) from app.gateway.document_rewrite_pipeline import ( MAX_DOCUMENT_REWRITE_CHARS as _DOCUMENT_REWRITE_MAX_CHARS, ) from app.gateway.document_rewrite_pipeline import ( astream_rewrite_requirements as _astream_rewrite_requirements, ) from app.gateway.document_rewrite_pipeline import ( atomic_write_text as _atomic_write_text, ) from app.gateway.document_rewrite_pipeline import ( build_fallback_rewrite_requirements_plan as _build_fallback_rewrite_requirements_plan, ) from app.gateway.document_rewrite_pipeline import ( content_hash as _content_hash, ) from app.gateway.document_rewrite_pipeline import ( extract_stream_chunk_parts as _extract_chunk_parts, ) from app.gateway.document_rewrite_pipeline import ( validate_markdown_document as _markdown_validation_error, ) from app.gateway.document_rewrite_pipeline import ( visible_document_char_count as _count_document_chars, ) from app.gateway.document_rewrite_summary import build_document_rewrite_summary from app.gateway.path_utils import aresolve_thread_virtual_path from app.gateway.word_export import ( DOCX_MEDIA_TYPE, WordExportError, WordExportFontError, build_markdown_docx, ) from deerflow.config.app_config import AppConfig from deerflow.models import create_chat_model from deerflow.utils.stream_text import InlineThinkTagFilter logger = logging.getLogger(__name__) router = APIRouter(prefix="/api/writing", tags=["writing"]) # ── Prompts ─────────────────────────────────────────────────────────────────── _SYSTEM_PROMPT = """你是一位专业的中文写作助手,擅长润色、扩写、缩写和调整文风。 规则: - 只输出改写后的正文,不加任何前缀、解释或括号说明。 - 保持原文的语言(中文输入 → 中文输出,英文输入 → 英文输出)。 - 改写结果应与原文长度相近,除非任务要求扩写或缩写。 - 不要添加 markdown 标题或特殊格式,除非原文本身包含。 """ _ACTION_INSTRUCTIONS: dict[str, str] = { "polish": "请对以下文字进行润色,使其表达更流畅、自然,同时保留原意。", "expand": "请对以下文字进行扩写,补充细节和说明,使内容更丰富饱满。", "shorten": "请对以下文字进行缩写,保留核心意思,去除冗余表达,使其更简洁。", "tone": "请调整以下文字的语气风格。", "imitate": "请模仿给定范文的写作风格、句式与语气,对以下文字进行仿写改写,在保留原文核心信息的前提下尽量贴近范文的表达方式。", } # 仿写参考范文的最大长度(防止超长 prompt) _REFERENCE_MAX_CHARS = 2000 _TONE_SUFFIX: dict[str, str] = { "professional": "目标风格:专业、正式。", "casual": "目标风格:轻松、口语化。", "readable": "目标风格:通俗易懂、易于阅读。", "academic": "目标风格:学术严谨、逻辑清晰。", "expressive": "目标风格:生动有文采、富有感染力。", } _DOCUMENT_REWRITE_RESERVATION_LOCK = asyncio.Lock() _ACTIVE_DOCUMENT_REWRITES: set[str] = set() # ── Request model ───────────────────────────────────────────────────────────── class SelectionInfo(BaseModel): from_: int = Field(..., alias="from", description="ProseMirror start offset") to: int = Field(..., description="ProseMirror end offset") text: str = Field(..., description="Selected plain text") model_config = {"populate_by_name": True} class SurroundingContext(BaseModel): before: str = Field(default="", description="Text before the selection") after: str = Field(default="", description="Text after the selection") class RewriteRequest(BaseModel): thread_id: str | None = Field(default=None, alias="threadId") selection: SelectionInfo action: str = Field(..., description="polish | expand | shorten | tone | imitate") tone: str | None = Field(default=None) model_name: str | None = Field(default=None, alias="modelName") surrounding_context: SurroundingContext = Field(default_factory=SurroundingContext, alias="surroundingContext") custom_instruction: str | None = Field(default=None, alias="customInstruction") reference_text: str | None = Field(default=None, alias="referenceText", description="仿写参考范文(action=imitate)") follow_up_instruction: str | None = Field(default=None, alias="followUpInstruction") previous_result: str | None = Field(default=None, alias="previousResult") model_config = {"populate_by_name": True} class MarkdownDocxExportRequest(BaseModel): """Markdown source for the server-rendered formal Word download.""" title: str | None = Field(default=None, max_length=300) markdown: str = Field(..., min_length=1, max_length=2_000_000) embedded_images: dict[str, str] = Field(default_factory=dict) class DocumentRewriteRequest(BaseModel): """Request for a complete Markdown artifact rewrite. This endpoint deliberately streams a *temporary* draft. The artifact is written only after generation and Markdown validation succeed and the original hash still matches. It is therefore safe for the sandbox to visually clear and retype the document without risking a half-written file. """ thread_id: str = Field(..., alias="threadId") path: str instruction: str = Field(..., min_length=1, max_length=4_000) style: str | None = Field(default=None) model_name: str | None = Field(default=None, alias="modelName") report_outline: str | None = Field(default=None, alias="reportOutline", max_length=20_000) generate_images: bool = Field(default=False, alias="generateImages") max_generated_images: int = Field(default=2, alias="maxGeneratedImages", ge=1, le=4) expected_hash: str | None = Field(default=None, alias="expectedHash") model_config = {"populate_by_name": True} class DocumentRewriteUndoRequest(BaseModel): """Owner-scoped target for restoring one completed document rewrite.""" thread_id: str = Field(..., alias="threadId") path: str model_config = {"populate_by_name": True} # ── Helpers ─────────────────────────────────────────────────────────────────── def _sse(event: str, data: dict) -> str: return f"event: {event}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n" def _document_rewrite_key(path: Path) -> str: """Stable per-real-file key used to reject concurrent full rewrites.""" return os.path.normcase(str(path.resolve())) def _instruction_with_report_outline(instruction: str, report_outline: str | None) -> str: """Make an optional report template an explicit constraint for the rewrite model.""" normalized_instruction = instruction.strip() normalized_outline = (report_outline or "").strip() if not normalized_outline: return normalized_instruction return "\n\n".join( ( normalized_instruction, "【报告结构要求】", "请按以下结构输出全文:文首题名与「日期:」「类型:」「密级:」等版头必须保留并填入真实值(占位符 XXX/xxx/xxxxx 不得照抄或删行;日期用当天中文日期;类型按题名判断;密级无另行规定时写「内部」)。章节按大纲展开;仅当某小节完全没有材料时才可合并或省略该小节,不得为套用结构编造事实。", normalized_outline, ) ) async def _reserve_document_rewrite(path: Path) -> bool: """Reserve a physical document until its SSE generator reaches a terminal state.""" key = _document_rewrite_key(path) async with _DOCUMENT_REWRITE_RESERVATION_LOCK: if key in _ACTIVE_DOCUMENT_REWRITES: return False _ACTIVE_DOCUMENT_REWRITES.add(key) return True async def _release_document_rewrite(path: Path) -> None: key = _document_rewrite_key(path) async with _DOCUMENT_REWRITE_RESERVATION_LOCK: _ACTIVE_DOCUMENT_REWRITES.discard(key) def _artifact_rewrite_scope(thread_id: str, virtual_path: str) -> str: """A stable synthetic job session id for one owner-scoped sandbox file.""" material = f"artifact:{thread_id}:{virtual_path}".encode() return f"arw_{hashlib.sha256(material).hexdigest()[:56]}" def _artifact_job_matches(job: dict, *, thread_id: str, virtual_path: str) -> bool: snapshot = job.get("input_snapshot") return bool(isinstance(snapshot, dict) and snapshot.get("kind") == "artifact_document_rewrite" and snapshot.get("thread_id") == thread_id and snapshot.get("path") == virtual_path) def _artifact_job_response(job: dict) -> dict: """Limit durable job data to the rewrite fields the browser can hydrate.""" snapshot = job.get("input_snapshot") or {} return { "id": job.get("id"), "status": job.get("status"), "phase": job.get("phase"), "progress": job.get("progress"), "createdAt": job.get("created_at"), "result": job.get("result_snapshot") or {}, "request": { "instruction": str(snapshot.get("instruction") or ""), "style": str(snapshot.get("style") or "professional"), "modelName": snapshot.get("model_name"), "reportOutline": str(snapshot.get("report_outline") or ""), "generateImages": bool(snapshot.get("generate_images")), "maxGeneratedImages": int(snapshot.get("max_generated_images") or 2), "originalContent": snapshot.get("original_content") or "", }, } async def _stream_durable_artifact_rewrite_job( *, request: Request, job: dict, after: int = 0, ): """Translate the durable generic event log back to the artifact SSE contract.""" job_store = getattr(request.app.state, "deep_research_job_store", None) event_store = getattr(request.app.state, "deep_research_event_store", None) live_hub = getattr(request.app.state, "deep_research_live_hub", None) if job_store is None or event_store is None: raise HTTPException(status_code=503, detail="改写任务服务暂不可用。") job_id = str(job["id"]) async def generate(): cursor = after live_queue = live_hub.subscribe(job_id) if live_hub is not None else None loop = asyncio.get_running_loop() deadline = loop.time() + 30 * 60 next_durable_poll = 0.0 try: while loop.time() < deadline: if await request.is_disconnected(): return if loop.time() >= next_durable_poll: events = await event_store.list_after(job_id, after=cursor, limit=200) for event in events: cursor = max(cursor, int(event.get("seq") or 0)) if event.get("event_type") != "document_rewrite": continue payload = event.get("payload") or {} if isinstance(payload, dict) and payload.get("type"): yield _sse(str(payload["type"]), payload) current = await job_store.get_unscoped(job_id) if current and current.get("status") in {"completed", "failed", "cancelled"}: if current.get("status") == "cancelled": yield _sse("cancelled", {"message": "已停止全文改写,原文件未被修改。"}) elif current.get("status") == "failed": # A worker crash can happen before it had a chance # to persist a document_rewrite error milestone. # Never let the browser retain a spinning timeline # merely because the terminal job row already knows # the failure. yield _sse( "error", { "message": str(current.get("error_message") or "全文改写任务未完成,原文件未被修改。"), "stage": str(current.get("phase") or "generation"), }, ) return next_durable_poll = loop.time() + 0.35 if live_queue is None: await asyncio.sleep(max(0.01, next_durable_poll - loop.time())) continue try: live = await asyncio.wait_for( live_queue.get(), timeout=min(0.25, max(0.01, next_durable_poll - loop.time())), ) except TimeoutError: continue if str(live.get("jobId") or "") != job_id or live.get("type") != "document_rewrite": continue payload = live.get("payload") or {} if isinstance(payload, dict) and payload.get("type"): yield _sse(str(payload["type"]), payload) finally: if live_queue is not None: live_hub.unsubscribe(job_id, live_queue) return generate() # ── Prompt building (pure, testable) ────────────────────────────────────────── def build_rewrite_messages(body: RewriteRequest) -> tuple[str, list]: """Build (resolved_action, messages) for the rewrite request. Pure function (no model / IO) so the prompt assembly — including 仿写 (imitate) reference-sample injection — can be unit-tested directly. """ action = body.action if body.action in _ACTION_INSTRUCTIONS else "polish" instruction = _ACTION_INSTRUCTIONS[action] if action == "tone" and body.tone and body.tone in _TONE_SUFFIX: instruction = f"{instruction}\n{_TONE_SUFFIX[body.tone]}" context_parts: list[str] = [] before = body.surrounding_context.before.strip() after = body.surrounding_context.after.strip() if before or after: context_parts.append("文章上下文(仅供参考,不要改写):") if before: context_parts.append(f"[前文] {before[-300:]}") if after: context_parts.append(f"[后文] {after[:300]}") context_parts.append("") context_block = "\n".join(context_parts) base_instruction = instruction if body.custom_instruction: base_instruction = f"{instruction}\n用户额外要求:{body.custom_instruction}" if action == "imitate" and body.reference_text and body.reference_text.strip(): reference = body.reference_text.strip()[:_REFERENCE_MAX_CHARS] base_instruction = f"{base_instruction}\n\n【范文】(请模仿其文风、句式、语气与结构,但不要照抄其具体内容):\n{reference}" user_message = f"{base_instruction}\n\n{context_block}需要改写的文字:\n{body.selection.text}" if body.follow_up_instruction and body.previous_result: messages = [ SystemMessage(content=_SYSTEM_PROMPT), HumanMessage(content=user_message), AIMessage(content=body.previous_result), HumanMessage(content=body.follow_up_instruction), ] else: messages = [SystemMessage(content=_SYSTEM_PROMPT), HumanMessage(content=user_message)] return action, messages # ── Endpoint ────────────────────────────────────────────────────────────────── @router.post( "/export/docx", summary="Export Markdown as a formally formatted Word document", response_class=Response, ) async def export_markdown_docx(body: MarkdownDocxExportRequest) -> Response: """Generate DOCX on the server and embed the required title font. This route is intentionally usable by both authenticated workspaces and public scheduled-task pages. It performs no external fetches and does not persist the supplied Markdown. """ try: content = await asyncio.to_thread( build_markdown_docx, body.markdown, body.title, embedded_images=body.embedded_images or None, ) except WordExportFontError as exc: raise HTTPException(status_code=503, detail=str(exc)) from exc except WordExportError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc fallback_name = (body.title or "document").strip() or "document" fallback_name = re.sub(r'[<>:"/\\|?*\x00-\x1f]+', "_", fallback_name).strip(" .")[:120] or "document" filename = f"{fallback_name}.docx" return Response( content=content, media_type=DOCX_MEDIA_TYPE, headers={ "Content-Disposition": ( f'attachment; filename="document.docx"; filename*=UTF-8\'\'{quote(filename, safe="")}' ), "Cache-Control": "no-store", }, ) @router.post("/document-rewrite/jobs") @require_permission("threads", "write", owner_check=True) async def start_durable_document_rewrite_job( body: DocumentRewriteRequest, request: Request, ) -> dict: """Freeze an artifact rewrite into the durable job/event infrastructure.""" requested_path = body.path.strip() user_id = await get_current_user(request) suffix = Path(requested_path).suffix.lower() if suffix not in {".md", ".markdown"}: raise HTTPException(status_code=400, detail="AI 全文改写仅支持 Markdown 文件。") if ".skill/" in requested_path: raise HTTPException(status_code=400, detail="技能归档中的文件不支持全文改写。") actual_path = await aresolve_thread_virtual_path(body.thread_id, requested_path) if not actual_path.exists() or not actual_path.is_file(): raise HTTPException(status_code=404, detail=f"Artifact not found: {requested_path}") try: original = await asyncio.to_thread(actual_path.read_text, encoding="utf-8") except UnicodeDecodeError as exc: raise HTTPException(status_code=400, detail="该文件不是 UTF-8 文本,不能进行全文改写。") from exc if not original.strip(): raise HTTPException(status_code=400, detail="空 Markdown 文件不能进行全文改写。") if len(original) > _DOCUMENT_REWRITE_MAX_CHARS: raise HTTPException( status_code=400, detail=f"文件超过 {_DOCUMENT_REWRITE_MAX_CHARS} 字符,请先选择章节改写。", ) base_hash = _content_hash(original) if body.expected_hash and body.expected_hash != base_hash: raise HTTPException(status_code=409, detail="文件已变化,请刷新后重新发起全文改写。") job_store = getattr(request.app.state, "deep_research_job_store", None) dispatcher = getattr(request.app.state, "deep_research_dispatcher", None) if job_store is None or dispatcher is None: raise HTTPException(status_code=503, detail="改写任务服务暂不可用。") scope = _artifact_rewrite_scope(body.thread_id, requested_path) report_outline = body.report_outline.strip() if body.report_outline else "" instruction = _instruction_with_report_outline(body.instruction, report_outline) snapshot = { "kind": "artifact_document_rewrite", "session_id": scope, "user_id": user_id, "thread_id": body.thread_id, "path": requested_path, "original_content": original, "base_hash": base_hash, "instruction": instruction, "report_outline": report_outline, "style": body.style or "professional", "model_name": body.model_name, "generate_images": body.generate_images, "max_generated_images": body.max_generated_images, } request_hash = hashlib.sha256( json.dumps( { "base_hash": base_hash, "instruction": instruction, "report_outline": report_outline, "style": body.style or "professional", "model_name": body.model_name, "generate_images": body.generate_images, "max_generated_images": body.max_generated_images, }, ensure_ascii=False, sort_keys=True, ).encode("utf-8") ).hexdigest()[:32] job, created = await job_store.try_create_or_get_active( id=f"drj_{uuid.uuid4().hex}", session_id=scope, user_id=user_id, request_id=str(uuid.uuid4()), request_hash=request_hash, input_snapshot=snapshot, phase="initializing", ) if not _artifact_job_matches(job, thread_id=body.thread_id, virtual_path=requested_path): raise HTTPException(status_code=409, detail="该文件已有不兼容的改写任务,请先等待或停止它。") if created: dispatcher.nudge() return { "jobId": job["id"], "status": job["status"], "reused": not created, "baseHash": base_hash, } @router.get("/document-rewrite/jobs/latest") @require_permission("threads", "read") async def get_latest_durable_document_rewrite_job( request: Request, thread_id: str = Query(..., min_length=1), path: str = Query(..., min_length=1), ) -> dict: """Hydrate the current user's latest artifact rewrite after a page refresh. The durable-job store is already filtered by ``user_id`` below. Do not additionally perform a thread-owner lookup here: a historical thread can be displayed under a guest/SSO alias even though the current user has no matching ``threads_meta`` owner row, and the redundant check would turn an otherwise safe empty history into a 404. """ user_id = await get_current_user(request) job_store = getattr(request.app.state, "deep_research_job_store", None) if job_store is None: raise HTTPException(status_code=503, detail="改写任务服务暂不可用。") scope = _artifact_rewrite_scope(thread_id, path.strip()) rows = await job_store.list_by_session(scope, user_id=user_id, limit=1) job = rows[0] if rows else None if job is None or not _artifact_job_matches(job, thread_id=thread_id, virtual_path=path.strip()): return {"job": None, "result": {}, "request": {}} response = _artifact_job_response(job) return {"job": response, "result": response["result"], "request": response["request"]} @router.get("/document-rewrite/jobs") @require_permission("threads", "read") async def list_durable_document_rewrite_jobs( request: Request, thread_id: str = Query(..., min_length=1), path: str = Query(..., min_length=1), limit: int = Query(default=20, ge=1, le=20), ) -> dict: """Return the current user's persisted rewrite jobs for one Markdown file. ``list_by_session(..., user_id=user_id)`` below is the access boundary for this metadata-only recovery endpoint. It must return an empty list rather than a thread-ownership 404 when a legacy thread's display login differs from its persisted creator identity. """ user_id = await get_current_user(request) job_store = getattr(request.app.state, "deep_research_job_store", None) if job_store is None: raise HTTPException(status_code=503, detail="改写任务服务暂不可用。") virtual_path = path.strip() scope = _artifact_rewrite_scope(thread_id, virtual_path) rows = await job_store.list_by_session(scope, user_id=user_id, limit=limit) return {"items": [_artifact_job_response(row) for row in rows if _artifact_job_matches(row, thread_id=thread_id, virtual_path=virtual_path)]} @router.get("/document-rewrite/jobs/{job_id}/stream", response_class=StreamingResponse) @require_permission("threads", "read", owner_check=True) async def stream_durable_document_rewrite_job( job_id: str, request: Request, thread_id: str = Query(..., min_length=1), path: str = Query(..., min_length=1), after: int = Query(default=0, ge=0), ) -> StreamingResponse: user_id = await get_current_user(request) job_store = getattr(request.app.state, "deep_research_job_store", None) if job_store is None: raise HTTPException(status_code=503, detail="改写任务服务暂不可用。") job = await job_store.get(job_id, user_id=user_id) if job is None or not _artifact_job_matches(job, thread_id=thread_id, virtual_path=path.strip()): raise HTTPException(status_code=404, detail="改写任务不存在或无权访问。") stream = await _stream_durable_artifact_rewrite_job(request=request, job=job, after=after) return StreamingResponse( stream, media_type="text/event-stream", headers={"X-Accel-Buffering": "no", "Cache-Control": "no-cache"}, ) @router.post("/document-rewrite/jobs/{job_id}/cancel") @require_permission("threads", "write", owner_check=True) async def cancel_durable_document_rewrite_job( job_id: str, request: Request, thread_id: str = Query(..., min_length=1), path: str = Query(..., min_length=1), ) -> dict: user_id = await get_current_user(request) job_store = getattr(request.app.state, "deep_research_job_store", None) if job_store is None: raise HTTPException(status_code=503, detail="改写任务服务暂不可用。") job = await job_store.get(job_id, user_id=user_id) if job is None or not _artifact_job_matches(job, thread_id=thread_id, virtual_path=path.strip()): raise HTTPException(status_code=404, detail="改写任务不存在或无权访问。") cancelled = await job_store.request_cancel(job_id, user_id=user_id) if cancelled is None: raise HTTPException(status_code=409, detail="该改写任务当前无法停止。") if not cancelled.get("lease_owner"): await job_store.finalize_cancel(job_id) return {"jobId": job_id, "status": "cancelled"} @router.post( "/document-rewrite", summary="Rewrite a complete Markdown artifact (SSE stream)", response_class=StreamingResponse, ) @require_permission("threads", "write", owner_check=True) async def rewrite_document( body: DocumentRewriteRequest, request: Request, config: AppConfig = Depends(get_config), ) -> StreamingResponse: """Stream a full Markdown rewrite and atomically replace the artifact. The generated text is never appended to the real file while tokens are arriving. The final compare-and-swap protects a user edit made while the rewrite was in progress. """ requested_path = body.path.strip() user_id = await get_current_user(request) version_store = getattr(request.app.state, "document_rewrite_version_store", None) report_outline = body.report_outline.strip() if body.report_outline else "" instruction = _instruction_with_report_outline(body.instruction, report_outline) suffix = Path(requested_path).suffix.lower() if suffix not in {".md", ".markdown"}: raise HTTPException(status_code=400, detail="AI 全文改写仅支持 Markdown 文件。") if ".skill/" in requested_path: raise HTTPException(status_code=400, detail="技能归档中的文件不支持全文改写。") actual_path = await aresolve_thread_virtual_path(body.thread_id, requested_path) if not actual_path.exists() or not actual_path.is_file(): raise HTTPException(status_code=404, detail=f"Artifact not found: {requested_path}") try: original = await asyncio.to_thread(actual_path.read_text, encoding="utf-8") except UnicodeDecodeError as exc: raise HTTPException(status_code=400, detail="该文件不是 UTF-8 文本,不能进行全文改写。") from exc if not original.strip(): raise HTTPException(status_code=400, detail="空 Markdown 文件不能进行全文改写。") if len(original) > _DOCUMENT_REWRITE_MAX_CHARS: raise HTTPException( status_code=400, detail=f"文件超过 {_DOCUMENT_REWRITE_MAX_CHARS} 字符,请先选择章节改写。", ) base_hash = _content_hash(original) if body.expected_hash and body.expected_hash != base_hash: raise HTTPException(status_code=409, detail="文件已变化,请刷新后重新发起全文改写。") if not await _reserve_document_rewrite(actual_path): raise HTTPException(status_code=409, detail="该文件正在执行全文改写,请等待当前任务完成后再试。") try: style_instruction = _DOCUMENT_STYLE_INSTRUCTIONS.get( (body.style or "").strip(), _DOCUMENT_STYLE_INSTRUCTIONS["professional"], ) resolved_name = body.model_name or (config.models[0].name if config.models else None) model_cfg = config.get_model_config(resolved_name) if resolved_name else None thinking_enabled = bool(model_cfg and getattr(model_cfg, "supports_thinking", False)) except Exception: await _release_document_rewrite(actual_path) raise async def generate(): draft_parts: list[str] = [] try: yield _sse( "snapshot_fixed", { "fileName": actual_path.name, "originalCharCount": _count_document_chars(original), "baseHash": base_hash, }, ) yield _sse( "requirements_started", {}, ) model = create_chat_model( name=body.model_name, thinking_enabled=thinking_enabled, app_config=config, ) requirement_parts: list[str] = [] requirements_fallback = False try: async for delta in _astream_rewrite_requirements( model, file_name=actual_path.name, instruction=instruction, style_instruction=style_instruction, original=original, ): requirement_parts.append(delta) yield _sse("requirements_delta", {"delta": delta}) except asyncio.CancelledError: raise except Exception: # noqa: BLE001 logger.exception( "writing/document-rewrite requirements planning failed: thread=%s path=%s", body.thread_id, requested_path, ) requirements_fallback = True requirements_plan = "".join(requirement_parts).strip()[:3_000] if not requirements_plan: requirements_fallback = True requirements_plan = _build_fallback_rewrite_requirements_plan( instruction=instruction, style_instruction=style_instruction, ) yield _sse( "requirements_completed", {"content": requirements_plan, "fallback": requirements_fallback}, ) messages = [ SystemMessage(content=_DOCUMENT_REWRITE_SYSTEM_PROMPT), HumanMessage( content=( f"文件名:{actual_path.name}\n" f"改写要求:{instruction}\n" f"风格要求:{style_instruction}\n\n" "【已确认的改写执行计划开始】\n" f"{requirements_plan}\n" "【已确认的改写执行计划结束】\n\n" "【原始 Markdown 开始】\n" f"{original}\n" "【原始 Markdown 结束】" ) ), ] yield _sse( "requirements_resolved", { "instruction": instruction, "style": body.style or "professional", "modelName": resolved_name or "system_default", "plan": requirements_plan, }, ) inline_think_filter = InlineThinkTagFilter() async for chunk in model.astream(messages, config={"run_name": "document_rewrite"}): thinking_delta, text_delta = _extract_chunk_parts(chunk) if thinking_delta: yield _sse("thinking", {"chunk": thinking_delta}) text_delta, inline_thinking = inline_think_filter.push_parts(text_delta) if inline_thinking: yield _sse("thinking", {"chunk": inline_thinking}) if text_delta: draft_parts.append(text_delta) yield _sse( "rewrite_delta", { "delta": text_delta, "charCount": _count_document_chars("".join(draft_parts)), }, ) tail, inline_thinking_tail = inline_think_filter.finish_parts() if inline_thinking_tail: yield _sse("thinking", {"chunk": inline_thinking_tail}) if tail: draft_parts.append(tail) yield _sse( "rewrite_delta", { "delta": tail, "charCount": _count_document_chars("".join(draft_parts)), }, ) rewritten = "".join(draft_parts).strip() yield _sse("validation_started", {}) validation_error = _markdown_validation_error(rewritten) if validation_error: yield _sse("error", {"message": validation_error, "stage": "validation"}) return heading_count_before = sum(1 for line in original.splitlines() if line.lstrip().startswith("#")) heading_count_after = sum(1 for line in rewritten.splitlines() if line.lstrip().startswith("#")) warnings = ["标题数量发生变化,请在前后对比中确认结构。"] if heading_count_before != heading_count_after else [] yield _sse( "validation_completed", { "passed": True, "warnings": warnings, }, ) yield _sse( "comparison_ready", build_document_rewrite_summary( source_display_name=actual_path.name, instruction=instruction, model_name=resolved_name or "system_default", original=original, rewritten=rewritten, validation_warnings=warnings, ), ) yield _sse("commit_started", {}) current = await asyncio.to_thread(actual_path.read_text, encoding="utf-8") if _content_hash(current) != base_hash: yield _sse( "conflict", { "message": "文件在改写期间已被修改,已保留候选稿但没有覆盖原文件。", }, ) return committed_hash = _content_hash(rewritten) await asyncio.to_thread(_atomic_write_text, actual_path, rewritten) version_id: str | None = None version_snapshot_failed = False if version_store is not None: try: version_id = f"drv_{uuid.uuid4().hex}" await version_store.create( id=version_id, user_id=user_id, source_type="artifact", source_id=body.thread_id, source_path=requested_path, original_content=original, original_hash=base_hash, committed_hash=committed_hash, instruction=instruction, model_name=resolved_name, ) except Exception: # noqa: BLE001 # A version snapshot enables undo but is not allowed to # discard a successfully committed document on failure. logger.exception( "could not persist artifact rewrite version for %s; keeping rewritten file", actual_path, ) version_id = None version_snapshot_failed = True else: try: await version_store.prune_source( user_id=user_id, source_type="artifact", source_id=body.thread_id, source_path=requested_path, ) except Exception: # noqa: BLE001 logger.exception("could not prune artifact rewrite versions for %s", actual_path) yield _sse( "committed", { "content": rewritten, "committedHash": committed_hash, "versionId": version_id, "versionSnapshotFailed": version_snapshot_failed, }, ) yield _sse( "done", { "committedHash": committed_hash, "versionId": version_id, "versionSnapshotFailed": version_snapshot_failed, }, ) except asyncio.CancelledError: raise except Exception as exc: # noqa: BLE001 logger.exception( "writing/document-rewrite failed: thread=%s path=%s", body.thread_id, requested_path, ) yield _sse("error", {"message": str(exc), "stage": "generation"}) finally: await _release_document_rewrite(actual_path) return StreamingResponse( generate(), media_type="text/event-stream", headers={"X-Accel-Buffering": "no", "Cache-Control": "no-cache"}, ) @router.post("/document-rewrite/versions/{version_id}/undo") @require_permission("threads", "write", owner_check=True) async def undo_document_rewrite( version_id: str, body: DocumentRewriteUndoRequest, request: Request, ) -> dict: """Restore the pre-rewrite snapshot if nobody edited the artifact since.""" user_id = await get_current_user(request) version_store = getattr(request.app.state, "document_rewrite_version_store", None) if version_store is None: raise HTTPException(status_code=503, detail="改写版本服务暂不可用,无法执行撤销。") version = await version_store.get(version_id, user_id=user_id, source_type="artifact") if version is None: raise HTTPException(status_code=404, detail="未找到可撤销的改写版本。") if version.get("reverted_at") is not None: raise HTTPException(status_code=409, detail="该改写版本已撤销。") if version.get("source_id") != body.thread_id or version.get("source_path") != body.path.strip(): raise HTTPException(status_code=409, detail="撤销目标与改写版本不一致。") actual_path = await aresolve_thread_virtual_path(body.thread_id, body.path.strip()) if not actual_path.exists() or not actual_path.is_file(): raise HTTPException(status_code=404, detail="原文件已不存在,无法撤销。") if not await _reserve_document_rewrite(actual_path): raise HTTPException(status_code=409, detail="该文件已有全文改写或版本恢复任务正在执行,请稍后重试。") try: current = await asyncio.to_thread(actual_path.read_text, encoding="utf-8") if _content_hash(current) != version["committed_hash"]: raise HTTPException(status_code=409, detail="文件在改写后已被修改,不能覆盖式撤销。") original_content = str(version.get("original_content") or "") restored_hash = _content_hash(original_content) # Retain the state being undone before touching the real file, so the # user can safely return to this rewritten version once. redo_version_id = f"drv_{uuid.uuid4().hex}" await asyncio.to_thread(_atomic_write_text, actual_path, original_content) try: await version_store.create( id=redo_version_id, user_id=user_id, source_type="artifact", source_id=body.thread_id, source_path=body.path.strip(), original_content=current, original_hash=version["committed_hash"], committed_hash=restored_hash, instruction="恢复 AI 全文改写前版本", model_name=version.get("model_name"), ) except Exception: current_after_restore = await asyncio.to_thread(actual_path.read_text, encoding="utf-8") if _content_hash(current_after_restore) == restored_hash: await asyncio.to_thread(_atomic_write_text, actual_path, current) raise try: await version_store.prune_source( user_id=user_id, source_type="artifact", source_id=body.thread_id, source_path=body.path.strip(), ) except Exception: # noqa: BLE001 logger.exception("could not prune artifact rewrite versions for %s", actual_path) if not await version_store.mark_reverted(version_id, user_id=user_id): # Do not report a completed undo when the old snapshot could not # be consumed; restore the committed document if no user edit won # the very short persistence window. current_after_restore = await asyncio.to_thread(actual_path.read_text, encoding="utf-8") if _content_hash(current_after_restore) == restored_hash: await asyncio.to_thread(_atomic_write_text, actual_path, current) raise HTTPException(status_code=409, detail="撤销版本状态已变化,未覆盖文件。") return { "content": original_content, "restoredHash": restored_hash, "versionId": version_id, "redoVersionId": redo_version_id, } finally: await _release_document_rewrite(actual_path) @router.post( "/rewrite", summary="AI Text Rewrite (SSE stream)", response_class=StreamingResponse, ) async def rewrite_selection( body: RewriteRequest, config: AppConfig = Depends(get_config), ) -> StreamingResponse: action, messages = build_rewrite_messages(body) # Auto-enable thinking when the selected model supports it resolved_name = body.model_name or (config.models[0].name if config.models else None) model_cfg = config.get_model_config(resolved_name) if resolved_name else None thinking_enabled = bool(model_cfg and getattr(model_cfg, "supports_thinking", False)) async def generate(): try: model = create_chat_model(name=body.model_name, thinking_enabled=thinking_enabled, app_config=config) inline_think_filter = InlineThinkTagFilter() async for chunk in model.astream(messages, config={"run_name": "writing_rewrite"}): thinking_delta, text_delta = _extract_chunk_parts(chunk) if thinking_delta: yield _sse("thinking", {"chunk": thinking_delta}) text_delta, inline_thinking = inline_think_filter.push_parts(text_delta) if inline_thinking: yield _sse("thinking", {"chunk": inline_thinking}) if text_delta: yield _sse("text", {"chunk": text_delta}) tail, inline_thinking_tail = inline_think_filter.finish_parts() if inline_thinking_tail: yield _sse("thinking", {"chunk": inline_thinking_tail}) if tail: yield _sse("text", {"chunk": tail}) yield _sse("done", {}) except Exception as e: logger.exception("writing/rewrite failed: thread=%s action=%s", body.thread_id, action) yield _sse("error", {"message": str(e)}) return StreamingResponse( generate(), media_type="text/event-stream", headers={"X-Accel-Buffering": "no", "Cache-Control": "no-cache"}, )