1024 lines
45 KiB
Python
1024 lines
45 KiB
Python
"""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"},
|
||
)
|