实现模板解析与模板管理基础链路
This commit is contained in:
+14
-3
@@ -1,7 +1,8 @@
|
||||
import uvicorn
|
||||
from contextlib import asynccontextmanager
|
||||
from fastapi import FastAPI
|
||||
from fastapi import FastAPI, HTTPException, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse
|
||||
from database import init_db, engine
|
||||
from config import settings
|
||||
from routers import templates, models, generate, export
|
||||
@@ -32,14 +33,24 @@ app.include_router(generate.router, prefix="/api/v1/generate", tags=["生成管
|
||||
app.include_router(export.router, prefix="/api/v1/export", tags=["导出管理"])
|
||||
|
||||
|
||||
@app.exception_handler(HTTPException)
|
||||
async def http_exception_handler(_: Request, exc: HTTPException):
|
||||
return JSONResponse(status_code=exc.status_code, content={"code": -1, "message": exc.detail})
|
||||
|
||||
|
||||
@app.exception_handler(Exception)
|
||||
async def unhandled_exception_handler(_: Request, exc: Exception):
|
||||
return JSONResponse(status_code=500, content={"code": -1, "message": str(exc) or "服务器内部错误"})
|
||||
|
||||
|
||||
@app.get("/")
|
||||
async def root():
|
||||
return {"message": "AI 文档模板生成系统 API", "version": settings.APP_VERSION}
|
||||
return {"code": 0, "data": {"name": settings.APP_NAME, "version": settings.APP_VERSION}, "message": "ok"}
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "ok"}
|
||||
return {"code": 0, "data": {"status": "ok"}, "message": "ok"}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -0,0 +1,224 @@
|
||||
import asyncio
|
||||
import os
|
||||
import tempfile
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from io import BytesIO
|
||||
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from config import settings
|
||||
from database import get_db
|
||||
from models.paragraph import Paragraph
|
||||
from models.template import Template
|
||||
from schemas.schemas import Response, TemplateSave
|
||||
from services.minio_client import minio_client
|
||||
from services.template_parser import parse_template
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _build_object_path(filename: str) -> tuple[str, str]:
|
||||
ext = os.path.splitext(filename)[1].lower()
|
||||
date_prefix = datetime.now().strftime("%Y%m%d")
|
||||
object_name = f"{date_prefix}/{uuid.uuid4().hex}{ext}"
|
||||
return ext, object_name
|
||||
|
||||
|
||||
def _serialize_paragraph(paragraph: Paragraph) -> dict:
|
||||
return {
|
||||
"id": paragraph.id,
|
||||
"template_id": paragraph.template_id,
|
||||
"sort_index": paragraph.sort_index,
|
||||
"title": paragraph.title,
|
||||
"content": paragraph.content,
|
||||
"style_json": paragraph.style_json,
|
||||
"is_table": paragraph.is_table,
|
||||
"table_json": paragraph.table_json,
|
||||
"edit_mode": paragraph.edit_mode,
|
||||
"model_id": paragraph.model_id,
|
||||
"need_prompt": paragraph.need_prompt,
|
||||
"prompt_text": paragraph.prompt_text,
|
||||
"need_file": paragraph.need_file,
|
||||
"file_note": paragraph.file_note,
|
||||
"output_format": paragraph.output_format,
|
||||
}
|
||||
|
||||
|
||||
def _serialize_template(template: Template) -> dict:
|
||||
return {
|
||||
"id": template.id,
|
||||
"name": template.name,
|
||||
"description": template.description,
|
||||
"file_path": template.file_path,
|
||||
"paragraph_count": template.paragraph_count,
|
||||
"status": template.status,
|
||||
"created_at": template.created_at,
|
||||
"updated_at": template.updated_at,
|
||||
}
|
||||
|
||||
|
||||
@router.get("")
|
||||
async def list_templates(
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(20, ge=1, le=100),
|
||||
keyword: str = Query("", alias="q"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
filters = []
|
||||
if keyword:
|
||||
filters.append(Template.name.like(f"%{keyword}%"))
|
||||
|
||||
total_stmt = select(func.count(Template.id))
|
||||
list_stmt = select(Template).order_by(Template.id.desc())
|
||||
if filters:
|
||||
total_stmt = total_stmt.where(*filters)
|
||||
list_stmt = list_stmt.where(*filters)
|
||||
|
||||
total = (await db.execute(total_stmt)).scalar_one()
|
||||
result = await db.execute(list_stmt.offset((page - 1) * page_size).limit(page_size))
|
||||
items = [_serialize_template(item) for item in result.scalars().all()]
|
||||
return Response(
|
||||
data={"items": items, "total": total, "page": page, "page_size": page_size}
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{template_id}")
|
||||
async def get_template(template_id: int, db: AsyncSession = Depends(get_db)):
|
||||
template = await db.get(Template, template_id)
|
||||
if template is None:
|
||||
raise HTTPException(status_code=404, detail="模板不存在")
|
||||
|
||||
result = await db.execute(
|
||||
select(Paragraph)
|
||||
.where(Paragraph.template_id == template_id)
|
||||
.order_by(Paragraph.sort_index.asc(), Paragraph.id.asc())
|
||||
)
|
||||
paragraphs = [_serialize_paragraph(item) for item in result.scalars().all()]
|
||||
payload = _serialize_template(template)
|
||||
payload["paragraphs"] = paragraphs
|
||||
return Response(data=payload)
|
||||
|
||||
|
||||
@router.post("/upload")
|
||||
async def upload_template(file: UploadFile = File(...), db: AsyncSession = Depends(get_db)):
|
||||
if not file.filename:
|
||||
raise HTTPException(status_code=400, detail="文件名不能为空")
|
||||
|
||||
ext, object_name = _build_object_path(file.filename)
|
||||
if ext != ".docx":
|
||||
raise HTTPException(status_code=400, detail="模板仅支持 .docx 格式")
|
||||
|
||||
content = await file.read()
|
||||
if not content:
|
||||
raise HTTPException(status_code=400, detail="上传文件不能为空")
|
||||
if len(content) > settings.MAX_UPLOAD_SIZE:
|
||||
raise HTTPException(status_code=400, detail="文件大小超过限制")
|
||||
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=ext) as temp_file:
|
||||
temp_file.write(content)
|
||||
temp_path = temp_file.name
|
||||
|
||||
try:
|
||||
parsed_items = await asyncio.to_thread(parse_template, temp_path)
|
||||
finally:
|
||||
if os.path.exists(temp_path):
|
||||
os.remove(temp_path)
|
||||
|
||||
await asyncio.to_thread(
|
||||
minio_client.put_object,
|
||||
settings.MINIO_BUCKET_TEMPLATES,
|
||||
object_name,
|
||||
BytesIO(content),
|
||||
len(content),
|
||||
file.content_type or "application/vnd.openxmlformats-officedocument.wordprocessingml.document",
|
||||
)
|
||||
|
||||
template = Template(
|
||||
name=os.path.splitext(file.filename)[0],
|
||||
description="",
|
||||
file_path=f"{settings.MINIO_BUCKET_TEMPLATES}/{object_name}",
|
||||
paragraph_count=len(parsed_items),
|
||||
status="draft",
|
||||
)
|
||||
db.add(template)
|
||||
await db.flush()
|
||||
|
||||
paragraph_rows: list[Paragraph] = []
|
||||
for item in parsed_items:
|
||||
paragraph = Paragraph(
|
||||
template_id=template.id,
|
||||
sort_index=item.sort_index,
|
||||
title=item.title,
|
||||
content=item.content,
|
||||
style_json=item.style_json,
|
||||
is_table=item.is_table,
|
||||
table_json=item.table_json,
|
||||
)
|
||||
db.add(paragraph)
|
||||
paragraph_rows.append(paragraph)
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(template)
|
||||
for paragraph in paragraph_rows:
|
||||
await db.refresh(paragraph)
|
||||
|
||||
payload = _serialize_template(template)
|
||||
payload["paragraphs"] = [_serialize_paragraph(item) for item in paragraph_rows]
|
||||
return Response(data=payload)
|
||||
|
||||
|
||||
@router.put("/{template_id}/paragraphs")
|
||||
async def save_template_paragraphs(
|
||||
template_id: int,
|
||||
body: TemplateSave,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
template = await db.get(Template, template_id)
|
||||
if template is None:
|
||||
raise HTTPException(status_code=404, detail="模板不存在")
|
||||
|
||||
result = await db.execute(select(Paragraph).where(Paragraph.template_id == template_id))
|
||||
paragraph_map = {item.id: item for item in result.scalars().all()}
|
||||
|
||||
for config in body.paragraphs:
|
||||
paragraph = paragraph_map.get(config.id)
|
||||
if paragraph is None:
|
||||
continue
|
||||
paragraph.sort_index = config.sort_index
|
||||
paragraph.title = config.title
|
||||
paragraph.edit_mode = config.edit_mode
|
||||
paragraph.model_id = config.model_id
|
||||
paragraph.need_prompt = config.need_prompt
|
||||
paragraph.prompt_text = config.prompt_text
|
||||
paragraph.need_file = config.need_file
|
||||
paragraph.file_note = config.file_note
|
||||
paragraph.output_format = config.output_format
|
||||
|
||||
await db.commit()
|
||||
return Response(data={"template_id": template_id, "saved": len(body.paragraphs)})
|
||||
|
||||
|
||||
@router.delete("/{template_id}")
|
||||
async def delete_template(template_id: int, db: AsyncSession = Depends(get_db)):
|
||||
template = await db.get(Template, template_id)
|
||||
if template is None:
|
||||
raise HTTPException(status_code=404, detail="模板不存在")
|
||||
|
||||
result = await db.execute(select(Paragraph).where(Paragraph.template_id == template_id))
|
||||
for paragraph in result.scalars().all():
|
||||
await db.delete(paragraph)
|
||||
|
||||
file_path = template.file_path or ""
|
||||
if "/" in file_path:
|
||||
bucket, object_name = file_path.split("/", 1)
|
||||
try:
|
||||
await asyncio.to_thread(minio_client.remove_object, bucket, object_name)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
await db.delete(template)
|
||||
await db.commit()
|
||||
return Response(data={"id": template_id})
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
from minio import Minio
|
||||
from config import settings
|
||||
|
||||
# MinIO 客户端
|
||||
minio_client = Minio(
|
||||
settings.MINIO_ENDPOINT,
|
||||
access_key=settings.MINIO_ACCESS_KEY,
|
||||
secret_key=settings.MINIO_SECRET_KEY,
|
||||
secure=settings.MINIO_USE_SSL,
|
||||
)
|
||||
|
||||
|
||||
async def init_buckets():
|
||||
"""初始化 MinIO 存储桶(应用启动时调用)"""
|
||||
buckets = [
|
||||
settings.MINIO_BUCKET_TEMPLATES, # 原始模板文件
|
||||
settings.MINIO_BUCKET_UPLOADS, # 用户上传的参考文件
|
||||
settings.MINIO_BUCKET_OUTPUTS, # 生成的文档
|
||||
]
|
||||
for bucket in buckets:
|
||||
if not minio_client.bucket_exists(bucket):
|
||||
minio_client.make_bucket(bucket)
|
||||
print(f"[MinIO] 创建存储桶: {bucket}")
|
||||
|
||||
|
||||
def get_file_url(bucket: str, object_name: str) -> str:
|
||||
"""获取文件的公开访问 URL"""
|
||||
if settings.MINIO_USE_SSL:
|
||||
protocol = "https"
|
||||
else:
|
||||
protocol = "http"
|
||||
return f"{protocol}://{settings.MINIO_ENDPOINT}/{bucket}/{object_name}"
|
||||
|
||||
|
||||
def get_presigned_url(bucket: str, object_name: str, expires: int = 3600) -> str:
|
||||
"""获取预签名下载 URL(带过期时间)"""
|
||||
return minio_client.presigned_get_object(bucket, object_name, expires=expires)
|
||||
@@ -0,0 +1,216 @@
|
||||
import json
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
|
||||
from docx import Document
|
||||
from docx.document import Document as DocumentObject
|
||||
from docx.oxml.ns import qn
|
||||
from docx.oxml.table import CT_Tbl
|
||||
from docx.oxml.text.paragraph import CT_P
|
||||
from docx.table import Table
|
||||
from docx.text.paragraph import Paragraph
|
||||
from docx.enum.text import WD_ALIGN_PARAGRAPH
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParsedParagraph:
|
||||
sort_index: int
|
||||
title: str
|
||||
content: str
|
||||
style_json: str
|
||||
is_table: bool
|
||||
table_json: str
|
||||
|
||||
|
||||
def _iter_block_items(document: DocumentObject) -> Iterator[Paragraph | Table]:
|
||||
body = document.element.body
|
||||
for child in body.iterchildren():
|
||||
if isinstance(child, CT_P):
|
||||
yield Paragraph(child, document)
|
||||
elif isinstance(child, CT_Tbl):
|
||||
yield Table(child, document)
|
||||
|
||||
|
||||
def _safe_pt(value: object) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return round(float(value.pt), 2)
|
||||
except AttributeError:
|
||||
return None
|
||||
|
||||
|
||||
def _safe_indent(value: object) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return round(float(value.pt), 2)
|
||||
except AttributeError:
|
||||
return None
|
||||
|
||||
|
||||
def _alignment_name(value: WD_ALIGN_PARAGRAPH | None) -> str:
|
||||
if value is None:
|
||||
return "LEFT"
|
||||
return getattr(value, "name", "LEFT")
|
||||
|
||||
|
||||
def _heading_level(style_name: str) -> int | None:
|
||||
if not style_name:
|
||||
return None
|
||||
normalized = style_name.lower().replace(" ", "")
|
||||
if normalized.startswith("heading"):
|
||||
level = normalized.replace("heading", "")
|
||||
if level.isdigit():
|
||||
return int(level)
|
||||
return None
|
||||
|
||||
|
||||
def _get_run_font_info(paragraph: Paragraph) -> dict:
|
||||
for run in paragraph.runs:
|
||||
if not run.text.strip():
|
||||
continue
|
||||
r_fonts = getattr(run._element.rPr, "rFonts", None) if run._element.rPr is not None else None
|
||||
east_asia = r_fonts.get(qn("w:eastAsia")) if r_fonts is not None else None
|
||||
color = None
|
||||
if run.font.color is not None and run.font.color.rgb is not None:
|
||||
color = str(run.font.color.rgb)
|
||||
return {
|
||||
"name": run.font.name,
|
||||
"eastAsia": east_asia,
|
||||
"size": _safe_pt(run.font.size),
|
||||
"bold": bool(run.bold) if run.bold is not None else False,
|
||||
"italic": bool(run.italic) if run.italic is not None else False,
|
||||
"color": color or "000000",
|
||||
}
|
||||
return {
|
||||
"name": None,
|
||||
"eastAsia": None,
|
||||
"size": None,
|
||||
"bold": False,
|
||||
"italic": False,
|
||||
"color": "000000",
|
||||
}
|
||||
|
||||
|
||||
def _capture_paragraph_style(paragraph: Paragraph, level: int) -> dict:
|
||||
fmt = paragraph.paragraph_format
|
||||
return {
|
||||
"font": _get_run_font_info(paragraph),
|
||||
"paragraph": {
|
||||
"alignment": _alignment_name(paragraph.alignment),
|
||||
"spaceBefore": _safe_pt(fmt.space_before),
|
||||
"spaceAfter": _safe_pt(fmt.space_after),
|
||||
"lineSpacing": fmt.line_spacing,
|
||||
"firstLineIndent": _safe_indent(fmt.first_line_indent),
|
||||
},
|
||||
"headingLevel": level,
|
||||
}
|
||||
|
||||
|
||||
def _get_cell_style(cell) -> dict:
|
||||
paragraph = cell.paragraphs[0] if cell.paragraphs else None
|
||||
font_info = _get_run_font_info(paragraph) if paragraph is not None else {
|
||||
"name": None,
|
||||
"eastAsia": None,
|
||||
"size": None,
|
||||
"bold": False,
|
||||
"italic": False,
|
||||
"color": "000000",
|
||||
}
|
||||
return {
|
||||
"font": font_info,
|
||||
"shading": None,
|
||||
"alignment": _alignment_name(paragraph.alignment) if paragraph is not None else "LEFT",
|
||||
"borders": {"top": None, "bottom": None, "left": None, "right": None},
|
||||
}
|
||||
|
||||
|
||||
def _extract_table_data(table: Table) -> dict:
|
||||
rows = len(table.rows)
|
||||
cols = max((len(row.cells) for row in table.rows), default=0)
|
||||
grid_span: dict[str, int] = {}
|
||||
cell_styles: list[dict] = []
|
||||
matrix: list[list[str]] = []
|
||||
|
||||
for row_index, row in enumerate(table.rows):
|
||||
row_values: list[str] = []
|
||||
for col_index, cell in enumerate(row.cells):
|
||||
text = "\n".join(paragraph.text.strip() for paragraph in cell.paragraphs if paragraph.text.strip())
|
||||
row_values.append(text)
|
||||
tc_pr = cell._tc.tcPr
|
||||
grid_span_value = None
|
||||
if tc_pr is not None and tc_pr.gridSpan is not None:
|
||||
grid_span_value = tc_pr.gridSpan.val
|
||||
if grid_span_value:
|
||||
grid_span[f"{row_index}-{col_index}"] = int(grid_span_value)
|
||||
cell_styles.append(_get_cell_style(cell))
|
||||
matrix.append(row_values)
|
||||
|
||||
return {
|
||||
"rows": rows,
|
||||
"cols": cols,
|
||||
"gridSpan": grid_span,
|
||||
"cellStyles": cell_styles,
|
||||
"tableWidth": None,
|
||||
"data": matrix,
|
||||
}
|
||||
|
||||
|
||||
def parse_template(file_path: str) -> list[ParsedParagraph]:
|
||||
document = Document(file_path)
|
||||
parsed: list[ParsedParagraph] = []
|
||||
current_item: ParsedParagraph | None = None
|
||||
loose_table_count = 0
|
||||
|
||||
for block in _iter_block_items(document):
|
||||
if isinstance(block, Paragraph):
|
||||
text = block.text.strip()
|
||||
if not text:
|
||||
continue
|
||||
|
||||
level = _heading_level(block.style.name if block.style is not None else "")
|
||||
if level is not None:
|
||||
current_item = ParsedParagraph(
|
||||
sort_index=len(parsed) + 1,
|
||||
title=text,
|
||||
content="",
|
||||
style_json=json.dumps(_capture_paragraph_style(block, level), ensure_ascii=False),
|
||||
is_table=False,
|
||||
table_json="{}",
|
||||
)
|
||||
parsed.append(current_item)
|
||||
continue
|
||||
|
||||
if current_item is None:
|
||||
current_item = ParsedParagraph(
|
||||
sort_index=len(parsed) + 1,
|
||||
title="未命名段落",
|
||||
content=text,
|
||||
style_json=json.dumps(_capture_paragraph_style(block, 0), ensure_ascii=False),
|
||||
is_table=False,
|
||||
table_json="{}",
|
||||
)
|
||||
parsed.append(current_item)
|
||||
else:
|
||||
current_item.content = "\n".join(filter(None, [current_item.content, text]))
|
||||
else:
|
||||
table_data = _extract_table_data(block)
|
||||
table_text = f"[表格] {table_data['rows']} 行 {table_data['cols']} 列"
|
||||
if current_item is None:
|
||||
loose_table_count += 1
|
||||
current_item = ParsedParagraph(
|
||||
sort_index=len(parsed) + 1,
|
||||
title=f"表格_{loose_table_count}",
|
||||
content=table_text,
|
||||
style_json="{}",
|
||||
is_table=True,
|
||||
table_json=json.dumps(table_data, ensure_ascii=False),
|
||||
)
|
||||
parsed.append(current_item)
|
||||
else:
|
||||
current_item.is_table = True
|
||||
current_item.table_json = json.dumps(table_data, ensure_ascii=False)
|
||||
current_item.content = "\n".join(filter(None, [current_item.content, table_text]))
|
||||
|
||||
return parsed
|
||||
Reference in New Issue
Block a user