diff --git a/backend/alembic/versions/5030e3939a26_create_paper_template_schema.py b/backend/alembic/versions/5030e3939a26_create_paper_template_schema.py new file mode 100644 index 0000000..eb09ea5 --- /dev/null +++ b/backend/alembic/versions/5030e3939a26_create_paper_template_schema.py @@ -0,0 +1,117 @@ +"""create paper template schema + +Creates the three tables behind the template configuration feature: + +``section_field`` the reusable heading library (name / level / font size / + colour). Deliberately flat — there is no ``parent_id``, so a + level-2 heading can be reused under any number of parents + and inside any number of templates. +``paper_template`` a named outline: name + abstract. +``template_field`` the join table, and the only place a display position is + recorded (``sort``). It has **no** unique constraint on + ``(template_id, field_id)``: placing the same field twice in + one template is a legitimate layout, and the UI warns about + it rather than the schema forbidding it. + +Note on the foreign keys below: TiDB parses ``FOREIGN KEY`` for compatibility +but does not enforce it. They are declared to document the relationships, and +the integrity they would provide is enforced in the application layer instead. + +Revision ID: 5030e3939a26 +Revises: +Create Date: 2026-09-18 15:46:10.533951 + +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision: str = "5030e3939a26" +down_revision: str | None = None +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +#: ``CURRENT_TIMESTAMP`` rather than ``now()``: it is the spelling MySQL and +#: TiDB both accept as a DATETIME column default without the expression-default +#: parentheses that only newer MySQL versions allow. +_NOW = sa.text("CURRENT_TIMESTAMP") + + +def upgrade() -> None: + op.create_table( + "paper_template", + sa.Column("id", sa.Integer(), autoincrement=True, nullable=False), + sa.Column("name", sa.String(length=255), nullable=False), + sa.Column("abstract", sa.Text(), nullable=True), + sa.Column("created_at", sa.DateTime(), server_default=_NOW, nullable=False), + sa.Column("updated_at", sa.DateTime(), server_default=_NOW, nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + + op.create_table( + "section_field", + sa.Column("id", sa.Integer(), autoincrement=True, nullable=False), + # Name carries the user's own numbering ("1. Introduction"); the API + # stores and returns it verbatim. + sa.Column("name", sa.String(length=255), nullable=False), + # Heading depth as a rendering hint: 1 = "1.", 2 = "1.1". + sa.Column("level", sa.Integer(), server_default=sa.text("1"), nullable=False), + # DECIMAL(4,1) because conventional Chinese sizes are fractional + # (五号 = 10.5pt). + sa.Column( + "font_size", + sa.Numeric(precision=4, scale=1), + server_default=sa.text("12.0"), + nullable=False, + ), + # Canonical RGB as #RRGGBB. + sa.Column( + "font_color", + sa.String(length=7), + server_default=sa.text("'#000000'"), + nullable=False, + ), + sa.Column("created_at", sa.DateTime(), server_default=_NOW, nullable=False), + sa.Column("updated_at", sa.DateTime(), server_default=_NOW, nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + + op.create_table( + "template_field", + sa.Column("id", sa.Integer(), autoincrement=True, nullable=False), + sa.Column("template_id", sa.Integer(), nullable=False), + sa.Column("field_id", sa.Integer(), nullable=False), + # Ascending integer; the sole determinant of render order. + sa.Column("sort", sa.Integer(), nullable=False), + sa.ForeignKeyConstraint( + ["field_id"], ["section_field.id"], ondelete="RESTRICT" + ), + sa.ForeignKeyConstraint( + ["template_id"], ["paper_template.id"], ondelete="CASCADE" + ), + sa.PrimaryKeyConstraint("id"), + ) + # Every read is "the fields of template X, in order", so one composite + # index serves both the filter and the sort. + op.create_index( + "ix_template_field_template_sort", + "template_field", + ["template_id", "sort"], + unique=False, + ) + op.create_index( + "ix_template_field_field_id", "template_field", ["field_id"], unique=False + ) + + +def downgrade() -> None: + # Reverse dependency order: the join table goes before the tables it points + # at. + op.drop_index("ix_template_field_field_id", table_name="template_field") + op.drop_index("ix_template_field_template_sort", table_name="template_field") + op.drop_table("template_field") + op.drop_table("section_field") + op.drop_table("paper_template") diff --git a/backend/app/api/router.py b/backend/app/api/router.py index f7028ea..79b673c 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -2,7 +2,9 @@ from fastapi import APIRouter -from app.api.routes import health +from app.api.routes import health, section_fields, templates api_router = APIRouter() api_router.include_router(health.router) +api_router.include_router(section_fields.router) +api_router.include_router(templates.router) diff --git a/backend/app/api/routes/section_fields.py b/backend/app/api/routes/section_fields.py new file mode 100644 index 0000000..3eef7f9 --- /dev/null +++ b/backend/app/api/routes/section_fields.py @@ -0,0 +1,155 @@ +"""Section-field library endpoints (字段管理). + +The library is global and reusable: a field exists once and is then placed into +any number of templates. That is why deleting a field is guarded — the join +table is the only thing keeping a template's outline intact, and silently +dropping a live heading from every template would be data loss, not cleanup. +""" + +from fastapi import APIRouter, Depends, HTTPException, Query, Response, status +from sqlalchemy.orm import Session + +from app.crud import section_field as crud +from app.db.session import get_db +from app.models import SectionField +from app.schemas.common import BatchDeleteRequest, BatchDeleteResult, PageResult +from app.schemas.section_field import ( + SectionFieldCreate, + SectionFieldRead, + SectionFieldUpdate, +) + +router = APIRouter(prefix="/section-fields", tags=["section-fields"]) + + +def _get_or_404(db: Session, field_id: int) -> SectionField: + field = crud.get(db, field_id) + if field is None: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"字段 {field_id} 不存在", + ) + return field + + +def _assert_unused(db: Session, fields: list[SectionField]) -> None: + """Refuse the delete if any field is still placed in a template. + + All offenders are reported at once rather than one per attempt, so a batch + delete does not turn into trial and error. + """ + counts = crud.usage_counts(db, [field.id for field in fields]) + if not counts: + return + + by_id = {field.id: field.name for field in fields} + blockers = "、".join( + f"“{by_id.get(field_id, field_id)}”({count} 个模板)" + for field_id, count in sorted(counts.items()) + ) + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"以下字段正被模板使用,请先在模板中移除:{blockers}", + ) + + +@router.get( + "", + response_model=PageResult[SectionFieldRead], + summary="List the field library", +) +def list_fields( + db: Session = Depends(get_db), + keyword: str | None = Query(default=None, description="按字段名称模糊搜索"), + level: int | None = Query(default=None, ge=1, le=9, description="按字段等级过滤"), + page: int = Query(default=1, ge=1), + page_size: int = Query(default=20, ge=1, le=500), +) -> PageResult[SectionFieldRead]: + """Browse the field library, grouped by level then creation order.""" + rows, total = crud.list_fields( + db, + keyword=keyword, + level=level, + page=page, + page_size=page_size, + ) + return PageResult.build( + items=[SectionFieldRead.model_validate(row) for row in rows], + total=total, + page=page, + page_size=page_size, + ) + + +@router.post( + "", + response_model=SectionFieldRead, + status_code=status.HTTP_201_CREATED, + summary="Create a section field", +) +def create_field( + payload: SectionFieldCreate, + db: Session = Depends(get_db), +) -> SectionFieldRead: + """Add a heading to the library, with its typography.""" + return SectionFieldRead.model_validate(crud.create(db, payload)) + + +@router.post( + "/batch-delete", + response_model=BatchDeleteResult, + summary="Delete several section fields", +) +def batch_delete_fields( + payload: BatchDeleteRequest, + db: Session = Depends(get_db), +) -> BatchDeleteResult: + """Delete the given fields, refusing wholesale if any is still in use.""" + fields = crud.get_many(db, payload.ids) + _assert_unused(db, fields) + + for field in fields: + crud.delete(db, field) + return BatchDeleteResult(deleted=len(fields)) + + +@router.get( + "/{field_id}", + response_model=SectionFieldRead, + summary="Fetch one section field", +) +def get_field(field_id: int, db: Session = Depends(get_db)) -> SectionFieldRead: + """Return a single field.""" + return SectionFieldRead.model_validate(_get_or_404(db, field_id)) + + +@router.patch( + "/{field_id}", + response_model=SectionFieldRead, + summary="Update a section field", +) +def update_field( + field_id: int, + payload: SectionFieldUpdate, + db: Session = Depends(get_db), +) -> SectionFieldRead: + """Rename or restyle a field. + + The change is visible in every template that places the field, because + templates store a reference rather than a copy. + """ + field = _get_or_404(db, field_id) + return SectionFieldRead.model_validate(crud.update(db, field, payload)) + + +@router.delete( + "/{field_id}", + status_code=status.HTTP_204_NO_CONTENT, + summary="Delete a section field", +) +def delete_field(field_id: int, db: Session = Depends(get_db)) -> Response: + """Delete an unused field.""" + field = _get_or_404(db, field_id) + _assert_unused(db, [field]) + crud.delete(db, field) + return Response(status_code=status.HTTP_204_NO_CONTENT) diff --git a/backend/app/api/routes/templates.py b/backend/app/api/routes/templates.py new file mode 100644 index 0000000..be63fad --- /dev/null +++ b/backend/app/api/routes/templates.py @@ -0,0 +1,167 @@ +"""Paper-template endpoints (模板管理). + +A template is created from two pieces of free text plus a free selection of +library fields. Nothing constrains the selection: the same field may be picked +twice, the picks may arrive in any order, and the only thing that decides how +the outline reads is ``sort``. +""" + +from fastapi import APIRouter, Depends, HTTPException, Query, Response, status +from sqlalchemy.orm import Session + +from app.crud import paper_template as crud +from app.db.session import get_db +from app.models import PaperTemplate +from app.schemas.common import BatchDeleteRequest, BatchDeleteResult, PageResult +from app.schemas.paper_template import ( + PaperTemplateCreate, + PaperTemplateListItem, + PaperTemplateRead, + PaperTemplateUpdate, +) + +router = APIRouter(prefix="/templates", tags=["templates"]) + + +def _get_or_404(db: Session, template_id: int) -> PaperTemplate: + template = crud.get(db, template_id) + if template is None: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"模板 {template_id} 不存在", + ) + return template + + +def _assert_name_free(db: Session, name: str, *, exclude_id: int | None = None) -> None: + if crud.name_taken(db, name, exclude_id=exclude_id): + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"模板名称“{name}”已存在", + ) + + +def _assert_fields_exist(db: Session, field_ids: list[int]) -> None: + """Reject a selection that references fields the library does not have. + + Checked in Python rather than by a foreign key because TiDB parses but does + not enforce ``FOREIGN KEY``, so an unchecked write would happily leave a + template pointing at nothing. + """ + missing = crud.missing_field_ids(db, field_ids) + if missing: + joined = "、".join(str(field_id) for field_id in missing) + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"字段不存在:{joined}", + ) + + +@router.get( + "", + response_model=PageResult[PaperTemplateListItem], + summary="List paper templates", +) +def list_templates( + db: Session = Depends(get_db), + keyword: str | None = Query(default=None, description="按模板名称或摘要模糊搜索"), + page: int = Query(default=1, ge=1), + page_size: int = Query(default=20, ge=1, le=200), +) -> PageResult[PaperTemplateListItem]: + """Browse templates, most recently edited first.""" + items, total = crud.list_templates( + db, keyword=keyword, page=page, page_size=page_size + ) + return PageResult.build(items=items, total=total, page=page, page_size=page_size) + + +@router.post( + "", + response_model=PaperTemplateRead, + status_code=status.HTTP_201_CREATED, + summary="Create a paper template", +) +def create_template( + payload: PaperTemplateCreate, + db: Session = Depends(get_db), +) -> PaperTemplateRead: + """Create a template from a name, an abstract, and a field selection.""" + _assert_name_free(db, payload.name) + _assert_fields_exist(db, [item.field_id for item in payload.fields]) + + template = crud.create( + db, + name=payload.name, + abstract=payload.abstract, + fields=payload.fields, + ) + return PaperTemplateRead.from_model(template) + + +@router.post( + "/batch-delete", + response_model=BatchDeleteResult, + summary="Delete several paper templates", +) +def batch_delete_templates( + payload: BatchDeleteRequest, + db: Session = Depends(get_db), +) -> BatchDeleteResult: + """Delete the given templates and all of their placement rows.""" + return BatchDeleteResult(deleted=crud.delete_many(db, payload.ids)) + + +@router.get( + "/{template_id}", + response_model=PaperTemplateRead, + summary="Fetch one paper template with its outline", +) +def get_template(template_id: int, db: Session = Depends(get_db)) -> PaperTemplateRead: + """Return a template; ``fields`` arrives already ordered by ``sort``.""" + return PaperTemplateRead.from_model(_get_or_404(db, template_id)) + + +@router.patch( + "/{template_id}", + response_model=PaperTemplateRead, + summary="Update a paper template", +) +def update_template( + template_id: int, + payload: PaperTemplateUpdate, + db: Session = Depends(get_db), +) -> PaperTemplateRead: + """Update the name, the abstract, the field selection, or any combination. + + ``fields`` is a full replacement when present. Omitting it leaves the + outline untouched; sending ``[]`` clears it. + """ + template = _get_or_404(db, template_id) + + if payload.name is not None: + _assert_name_free(db, payload.name, exclude_id=template_id) + if payload.fields is not None: + _assert_fields_exist(db, [item.field_id for item in payload.fields]) + + updated = crud.update( + db, + template, + name=payload.name, + abstract=payload.abstract, + # Both "field omitted" and "field set to null" arrive as None; only + # model_fields_set records which one the client actually sent. + abstract_provided="abstract" in payload.model_fields_set, + fields=payload.fields, + ) + return PaperTemplateRead.from_model(updated) + + +@router.delete( + "/{template_id}", + status_code=status.HTTP_204_NO_CONTENT, + summary="Delete a paper template", +) +def delete_template(template_id: int, db: Session = Depends(get_db)) -> Response: + """Delete a template. Library fields it referenced are left alone.""" + crud.delete(db, _get_or_404(db, template_id)) + return Response(status_code=status.HTTP_204_NO_CONTENT) diff --git a/backend/app/crud/__init__.py b/backend/app/crud/__init__.py index 2aabb63..38f82e4 100644 --- a/backend/app/crud/__init__.py +++ b/backend/app/crud/__init__.py @@ -1,7 +1,10 @@ """Data-access helpers. Functions here take a :class:`sqlalchemy.orm.Session` and operate on ORM -models. There are none yet — this package is a placeholder for that layer. +models. They own the commit: a route calls one function and gets back either a +persisted object or ``None``. """ -__all__: list[str] = [] +from app.crud import paper_template, section_field + +__all__ = ["paper_template", "section_field"] diff --git a/backend/app/crud/filters.py b/backend/app/crud/filters.py new file mode 100644 index 0000000..13e065f --- /dev/null +++ b/backend/app/crud/filters.py @@ -0,0 +1,19 @@ +"""Query-building helpers shared by the CRUD modules.""" + +LIKE_ESCAPE = "\\" + + +def like_pattern(keyword: str) -> str: + """Turn user input into a safe ``%keyword%`` LIKE pattern. + + A user searching for ``50%`` or ``a_b`` means those characters literally, + but ``%`` and ``_`` are LIKE wildcards. Escaping them — backslash first, + or the escape characters added for ``%`` would be doubled — keeps the + search behaving the way the search box implies. + """ + escaped = ( + keyword.replace(LIKE_ESCAPE, LIKE_ESCAPE * 2) + .replace("%", f"{LIKE_ESCAPE}%") + .replace("_", f"{LIKE_ESCAPE}_") + ) + return f"%{escaped}%" diff --git a/backend/app/crud/paper_template.py b/backend/app/crud/paper_template.py new file mode 100644 index 0000000..2ed8f5d --- /dev/null +++ b/backend/app/crud/paper_template.py @@ -0,0 +1,180 @@ +"""Data access for paper templates and their ordered field selections.""" + +from collections.abc import Sequence + +from sqlalchemy import func, or_, select +from sqlalchemy.orm import Session + +from app.crud.filters import LIKE_ESCAPE, like_pattern +from app.models import PaperTemplate, SectionField, TemplateField +from app.schemas.paper_template import PaperTemplateListItem, TemplateFieldInput + + +def _conditions(keyword: str | None) -> list: + """Search both the name and the abstract from one box.""" + if not keyword: + return [] + pattern = like_pattern(keyword) + return [ + or_( + PaperTemplate.name.like(pattern, escape=LIKE_ESCAPE), + PaperTemplate.abstract.like(pattern, escape=LIKE_ESCAPE), + ) + ] + + +def _field_count_column(): + """A correlated ``COUNT`` of the template's placement rows.""" + return ( + select(func.count(TemplateField.id)) + .where(TemplateField.template_id == PaperTemplate.id) + .correlate(PaperTemplate) + .scalar_subquery() + ) + + +def list_templates( + db: Session, + *, + keyword: str | None = None, + page: int = 1, + page_size: int = 20, +) -> tuple[list[PaperTemplateListItem], int]: + """Return one page of templates with their field counts, plus the total. + + The outline itself is left out: a table row needs the count, not the rows. + Recently edited templates come first, since that is what a user returns to. + """ + conditions = _conditions(keyword) + + total = db.scalar( + select(func.count(PaperTemplate.id)).where(*conditions) + ) or 0 + + stmt = ( + select(PaperTemplate, _field_count_column().label("field_count")) + .where(*conditions) + .order_by(PaperTemplate.updated_at.desc(), PaperTemplate.id.desc()) + .offset((page - 1) * page_size) + .limit(page_size) + ) + + items = [ + PaperTemplateListItem( + id=template.id, + name=template.name, + abstract=template.abstract, + field_count=field_count, + created_at=template.created_at, + updated_at=template.updated_at, + ) + for template, field_count in db.execute(stmt).all() + ] + return items, total + + +def get(db: Session, template_id: int) -> PaperTemplate | None: + """Return one template with its outline already loaded, or ``None``. + + ``items`` (ordered by ``sort``) and each item's ``field`` are both + configured for eager loading on the relationships, so this is a fixed + number of queries rather than one per field. + """ + return db.get(PaperTemplate, template_id) + + +def name_taken(db: Session, name: str, *, exclude_id: int | None = None) -> bool: + """Whether another template already uses ``name``.""" + stmt = select(func.count(PaperTemplate.id)).where(PaperTemplate.name == name) + if exclude_id is not None: + stmt = stmt.where(PaperTemplate.id != exclude_id) + return bool(db.scalar(stmt)) + + +def missing_field_ids(db: Session, field_ids: Sequence[int]) -> list[int]: + """Which of ``field_ids`` do not exist in the field library. + + Returning the offenders rather than a boolean lets the API name them, which + is the difference between a usable error and "invalid request". + """ + wanted = set(field_ids) + if not wanted: + return [] + found = set( + db.scalars(select(SectionField.id).where(SectionField.id.in_(wanted))).all() + ) + return sorted(wanted - found) + + +def _build_items(fields: Sequence[TemplateFieldInput]) -> list[TemplateField]: + """Materialise the client's selection into placement rows.""" + return [TemplateField(field_id=item.field_id, sort=item.sort) for item in fields] + + +def create( + db: Session, + *, + name: str, + abstract: str | None, + fields: Sequence[TemplateFieldInput], +) -> PaperTemplate: + """Insert a template together with its ordered selection.""" + template = PaperTemplate(name=name, abstract=abstract) + template.items = _build_items(fields) + db.add(template) + db.commit() + db.refresh(template) + return template + + +def update( + db: Session, + template: PaperTemplate, + *, + name: str | None = None, + abstract: str | None = None, + abstract_provided: bool = False, + fields: Sequence[TemplateFieldInput] | None = None, +) -> PaperTemplate: + """Apply a partial update. ``fields=None`` leaves the selection untouched. + + ``abstract_provided`` distinguishes "clear the abstract" from "leave it" — + both arrive as ``None`` in the payload, and only the caller knows which the + client meant. + """ + if name is not None: + template.name = name + if abstract_provided: + template.abstract = abstract + if fields is not None: + # delete-orphan removes the dropped rows on flush. TiDB does not + # enforce ON DELETE CASCADE, so this ORM-level cascade is the only + # thing cleaning up the join table. + template.items.clear() + db.flush() + template.items.extend(_build_items(fields)) + + db.commit() + db.refresh(template) + return template + + +def delete(db: Session, template: PaperTemplate) -> None: + """Delete a template and its placement rows.""" + db.delete(template) + db.commit() + + +def delete_many(db: Session, template_ids: Sequence[int]) -> int: + """Delete several templates, returning how many actually existed.""" + if not template_ids: + return 0 + templates = list( + db.scalars( + select(PaperTemplate).where(PaperTemplate.id.in_(list(template_ids))) + ).all() + ) + for template in templates: + db.delete(template) + db.commit() + return len(templates) diff --git a/backend/app/crud/section_field.py b/backend/app/crud/section_field.py new file mode 100644 index 0000000..35c2571 --- /dev/null +++ b/backend/app/crud/section_field.py @@ -0,0 +1,118 @@ +"""Data access for the reusable section-field library (字段管理).""" + +from collections.abc import Sequence + +from sqlalchemy import Select, func, select +from sqlalchemy.orm import Session + +from app.crud.filters import LIKE_ESCAPE, like_pattern +from app.models import SectionField, TemplateField +from app.schemas.section_field import SectionFieldCreate, SectionFieldUpdate + + +def _conditions(keyword: str | None, level: int | None) -> list: + """Translate the list filters into SQLAlchemy predicates.""" + conditions = [] + if keyword: + conditions.append( + SectionField.name.like(like_pattern(keyword), escape=LIKE_ESCAPE) + ) + if level is not None: + conditions.append(SectionField.level == level) + return conditions + + +def _ordered(stmt: Select) -> Select: + """Apply the library's canonical browse order. + + Grouped by ``level`` first so the picker reads as an outline, then by + insertion order so a field stays where the user put it. Deliberately *not* + ordered by ``name``: names carry hand-written numbering ("1.", "10.", + "2."), and string ordering would scramble it. + """ + return stmt.order_by(SectionField.level.asc(), SectionField.id.asc()) + + +def list_fields( + db: Session, + *, + keyword: str | None = None, + level: int | None = None, + page: int = 1, + page_size: int = 20, +) -> tuple[list[SectionField], int]: + """Return one page of the field library, plus the unpaged total.""" + conditions = _conditions(keyword, level) + + total = db.scalar( + select(func.count(SectionField.id)).where(*conditions) + ) or 0 + + stmt = _ordered( + select(SectionField) + .where(*conditions) + .offset((page - 1) * page_size) + .limit(page_size) + ) + return list(db.scalars(stmt).all()), total + + +def get(db: Session, field_id: int) -> SectionField | None: + """Return one field, or ``None``.""" + return db.get(SectionField, field_id) + + +def get_many(db: Session, field_ids: Sequence[int]) -> list[SectionField]: + """Return every field whose id is in ``field_ids`` (missing ids ignored).""" + if not field_ids: + return [] + stmt = select(SectionField).where(SectionField.id.in_(list(field_ids))) + return list(db.scalars(stmt).all()) + + +def usage_counts(db: Session, field_ids: Sequence[int]) -> dict[int, int]: + """Count how many *distinct templates* place each field. + + Drives the refusal message when a field in use is deleted. ``DISTINCT`` + matters because one template may legitimately place the same field twice. + """ + if not field_ids: + return {} + stmt = ( + select( + TemplateField.field_id, + func.count(func.distinct(TemplateField.template_id)), + ) + .where(TemplateField.field_id.in_(list(field_ids))) + .group_by(TemplateField.field_id) + ) + return {field_id: count for field_id, count in db.execute(stmt).all()} + + +def create(db: Session, data: SectionFieldCreate) -> SectionField: + """Insert a field.""" + field = SectionField(**data.model_dump()) + db.add(field) + db.commit() + db.refresh(field) + return field + + +def update(db: Session, field: SectionField, data: SectionFieldUpdate) -> SectionField: + """Apply a partial update to a field. + + ``exclude_unset`` is what makes PATCH semantics work: a key the client did + not send leaves the column alone, while an explicit ``null`` — which the + schemas reject for every nullable-typed column here — would not. + """ + for key, value in data.model_dump(exclude_unset=True).items(): + setattr(field, key, value) + db.commit() + db.refresh(field) + return field + + +def delete(db: Session, field: SectionField) -> None: + """Delete a field. Callers must check :func:`usage_counts` first.""" + db.delete(field) + db.commit() diff --git a/backend/app/models/__init__.py b/backend/app/models/__init__.py index 7f2b9d9..2dfc885 100644 --- a/backend/app/models/__init__.py +++ b/backend/app/models/__init__.py @@ -1,15 +1,18 @@ """SQLAlchemy ORM models. -This package is intentionally empty. - -No model classes exist yet, and no tables are created by this project at -import time. The schema will be introduced through Alembic migrations: - - alembic revision --autogenerate -m "describe the change" - alembic upgrade head - -When the first model is written, add it here (or in a submodule imported from -here) so that ``Base.metadata`` — and therefore autogenerate — can see it. +Importing this package registers every model on ``Base.metadata``, which is +what ``alembic revision --autogenerate`` inspects. A new model therefore has to +be added to the imports below, not only to its own module. """ -__all__: list[str] = [] +from app.models.mixins import TimestampMixin +from app.models.paper_template import PaperTemplate +from app.models.section_field import SectionField +from app.models.template_field import TemplateField + +__all__ = [ + "PaperTemplate", + "SectionField", + "TemplateField", + "TimestampMixin", +] diff --git a/backend/app/models/mixins.py b/backend/app/models/mixins.py new file mode 100644 index 0000000..84b99b8 --- /dev/null +++ b/backend/app/models/mixins.py @@ -0,0 +1,26 @@ +"""Column mixins shared by the ORM models.""" + +from datetime import datetime + +from sqlalchemy import DateTime, func +from sqlalchemy.orm import Mapped, mapped_column + + +class TimestampMixin: + """Adds ``created_at`` / ``updated_at`` to a model. + + Timestamps are generated by the database on insert and refreshed by the + ORM on update, so a row written by any client carries a server clock value. + """ + + created_at: Mapped[datetime] = mapped_column( + DateTime, + nullable=False, + server_default=func.now(), + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime, + nullable=False, + server_default=func.now(), + onupdate=func.now(), + ) diff --git a/backend/app/models/paper_template.py b/backend/app/models/paper_template.py new file mode 100644 index 0000000..35c191e --- /dev/null +++ b/backend/app/models/paper_template.py @@ -0,0 +1,46 @@ +"""Paper templates (模板表). + +A template is a named, ordered selection of library fields — the outline a +paper is written against. It stores no typography of its own: font size and +colour are read from the referenced :class:`~app.models.section_field.SectionField`, +so correcting a field's styling updates every template that uses it. +""" + +from typing import TYPE_CHECKING + +from sqlalchemy import String, Text +from sqlalchemy.orm import Mapped, mapped_column, relationship + +from app.db.base import Base +from app.models.mixins import TimestampMixin + +if TYPE_CHECKING: # pragma: no cover - typing only + from app.models.template_field import TemplateField + + +class PaperTemplate(TimestampMixin, Base): + """A named outline: a template name, a summary, and its ordered fields.""" + + __tablename__ = "paper_template" + + id: Mapped[int] = mapped_column(primary_key=True, autoincrement=True) + + name: Mapped[str] = mapped_column(String(255), nullable=False) + + #: Free-text abstract (摘要) describing when to use this template. + abstract: Mapped[str | None] = mapped_column(Text, nullable=True) + + #: The template's fields, always handed to callers in display order. + #: + #: ``delete-orphan`` is doing real work here: TiDB parses but does not + #: enforce ``ON DELETE CASCADE``, so removing a template's rows in the join + #: table is the ORM's job, not the database's. + items: Mapped[list["TemplateField"]] = relationship( + back_populates="template", + cascade="all, delete-orphan", + order_by="TemplateField.sort, TemplateField.id", + lazy="selectin", + ) + + def __repr__(self) -> str: # pragma: no cover - debugging aid + return f"" diff --git a/backend/app/models/section_field.py b/backend/app/models/section_field.py new file mode 100644 index 0000000..22f63bb --- /dev/null +++ b/backend/app/models/section_field.py @@ -0,0 +1,77 @@ +"""The reusable section-field library (字段表). + +A *section field* is one heading a paper can contain — "1. Introduction", +"2.1 Dataset", "0 Abstract". Fields live in this one global library and are +never owned by a template. + +Why there is no ``parent_id`` +----------------------------- +A tree would make a field usable under exactly one parent, so a level-2 heading +such as "Background" could not sit under both "1. Introduction" and +"2. Related Work". Hierarchy is therefore expressed only by +:attr:`SectionField.level` (1, 2, 3 ...), which is a *rendering hint* — it +drives indentation and numbering semantics in the UI — while a field stays +free to be attached to any number of templates and any number of parents +within them. + +The number that the reader sees is part of :attr:`name` and is written by the +user ("1. Introduction", "0 Abstract"). Nothing derives or rewrites it. + +Ordering inside a template is *not* stored here: it lives in +``template_field.sort``. This table has no ``sort`` column on purpose, so that +one library field can occupy a different position in every template. +""" + +from decimal import Decimal + +from sqlalchemy import Integer, Numeric, String, text +from sqlalchemy.orm import Mapped, mapped_column + +from app.db.base import Base +from app.models.mixins import TimestampMixin + + +class SectionField(TimestampMixin, Base): + """A single reusable heading, with the typography it should render in.""" + + __tablename__ = "section_field" + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + + #: Display name, numbering included. Stored and rendered verbatim. + name: Mapped[str] = mapped_column(String(255), nullable=False) + + #: Heading depth: 1 for "1.", 2 for "1.1", 3 for "1.1.1". Rendering hint + #: only — it does not link fields to each other. + level: Mapped[int] = mapped_column( + Integer, + nullable=False, + default=1, + server_default=text("1"), + ) + + #: Font size in points. ``Numeric`` rather than ``Integer`` because the + #: conventional Chinese sizes are fractional (五号 = 10.5pt, + #: 小四 = 12pt). + font_size: Mapped[Decimal] = mapped_column( + Numeric(4, 1), + nullable=False, + default=Decimal("12.0"), + server_default=text("12.0"), + ) + + #: Font colour as ``#RRGGBB`` — the canonical RGB encoding. The API + #: normalises ``rgb(20, 30, 40)``, ``#abc`` and bare ``aabbcc`` to it on + #: write, so the column always holds one comparable format. + font_color: Mapped[str] = mapped_column( + String(7), + nullable=False, + default="#000000", + server_default=text("'#000000'"), + ) + + def __repr__(self) -> str: # pragma: no cover - debugging aid + return ( + f"" + ) diff --git a/backend/app/models/template_field.py b/backend/app/models/template_field.py new file mode 100644 index 0000000..c2c76b1 --- /dev/null +++ b/backend/app/models/template_field.py @@ -0,0 +1,73 @@ +"""The template <-> field join table (关联表), which owns the display order. + +This table is the heart of the design. It carries exactly one piece of +information that belongs to neither side alone: ``sort`` — where this field +sits in *this* template. + +Consequences worth stating explicitly, because they are the requirements: + +* A field can appear in many templates, at a different position in each. +* A field can legitimately appear **more than once** in the same template + (e.g. a level-2 "Background" under both "1. Introduction" and + "2. Related Work"), so there is deliberately **no** unique constraint on + ``(template_id, field_id)``. The UI warns about repeats; it does not forbid + them. +* The user picks fields in any order they like. Nothing about the selection + order is stored — readers order strictly by ``sort``, then by ``id`` as a + stable tie-breaker. +""" + +from typing import TYPE_CHECKING + +from sqlalchemy import ForeignKey, Index, Integer +from sqlalchemy.orm import Mapped, mapped_column, relationship + +from app.db.base import Base + +if TYPE_CHECKING: # pragma: no cover - typing only + from app.models.paper_template import PaperTemplate + from app.models.section_field import SectionField + + +class TemplateField(Base): + """One placement of one library field inside one template.""" + + __tablename__ = "template_field" + + __table_args__ = ( + # Every read is "the fields of template X, in order", so the index + # covers both the filter and the sort. + Index("ix_template_field_template_sort", "template_id", "sort"), + Index("ix_template_field_field_id", "field_id"), + ) + + id: Mapped[int] = mapped_column(primary_key=True, autoincrement=True) + + template_id: Mapped[int] = mapped_column( + ForeignKey("paper_template.id", ondelete="CASCADE"), + nullable=False, + ) + + #: ``RESTRICT``: a field still placed in a template must not vanish from + #: under it. The API enforces this in Python as well, since TiDB does not + #: enforce foreign keys itself. + field_id: Mapped[int] = mapped_column( + ForeignKey("section_field.id", ondelete="RESTRICT"), + nullable=False, + ) + + #: Display position within the template. Plain ascending integer — lower + #: sorts render first regardless of the field's ``level``. + sort: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + + template: Mapped["PaperTemplate"] = relationship(back_populates="items") + + #: Eager-loaded because every template read needs the field's name and + #: typography; lazy loading would emit one query per row. + field: Mapped["SectionField"] = relationship(lazy="joined") + + def __repr__(self) -> str: # pragma: no cover - debugging aid + return ( + f"" + ) diff --git a/backend/app/schemas/__init__.py b/backend/app/schemas/__init__.py index 1bd246d..d5a1ef6 100644 --- a/backend/app/schemas/__init__.py +++ b/backend/app/schemas/__init__.py @@ -1 +1,39 @@ -"""Pydantic models describing request and response payloads.""" +"""Pydantic schemas for request and response payloads.""" + +from app.schemas.common import ( + BatchDeleteRequest, + BatchDeleteResult, + PageResult, + normalize_hex_color, +) +from app.schemas.paper_template import ( + PaperTemplateCreate, + PaperTemplateListItem, + PaperTemplateRead, + PaperTemplateUpdate, + TemplateFieldInput, + TemplateFieldRead, +) +from app.schemas.section_field import ( + SectionFieldCreate, + SectionFieldRead, + SectionFieldRef, + SectionFieldUpdate, +) + +__all__ = [ + "BatchDeleteRequest", + "BatchDeleteResult", + "PageResult", + "PaperTemplateCreate", + "PaperTemplateListItem", + "PaperTemplateRead", + "PaperTemplateUpdate", + "SectionFieldCreate", + "SectionFieldRead", + "SectionFieldRef", + "SectionFieldUpdate", + "TemplateFieldInput", + "TemplateFieldRead", + "normalize_hex_color", +] diff --git a/backend/app/schemas/common.py b/backend/app/schemas/common.py new file mode 100644 index 0000000..7056166 --- /dev/null +++ b/backend/app/schemas/common.py @@ -0,0 +1,121 @@ +"""Shared request/response shapes and value normalisation helpers.""" + +import re +from decimal import Decimal +from typing import Generic, TypeVar + +from pydantic import BaseModel, Field, field_serializer, field_validator + +T = TypeVar("T") + + +class PageResult(BaseModel, Generic[T]): + """One page of a list endpoint.""" + + items: list[T] + total: int + page: int + page_size: int + pages: int = 0 + + @classmethod + def build( + cls, + *, + items: list[T], + total: int, + page: int, + page_size: int, + ) -> "PageResult[T]": + """Assemble a page and derive ``pages`` from the page size. + + ``ceil`` is done in integers so an empty result reports 0 pages rather + than a misleading 1. + """ + pages = (total + page_size - 1) // page_size if page_size > 0 else 0 + return cls(items=items, total=total, page=page, page_size=page_size, pages=pages) + + +class BatchDeleteRequest(BaseModel): + """Body for the batch-delete endpoints.""" + + ids: list[int] = Field(min_length=1, description="Primary keys to delete") + + +class BatchDeleteResult(BaseModel): + """How many rows a batch delete actually removed.""" + + deleted: int + + +# ``#rgb`` / ``#rrggbb`` / ``rrggbb`` — the ``#`` is optional. +_HEX_RE = re.compile(r"^#?([0-9a-fA-F]{3}|[0-9a-fA-F]{6})$") +# ``rgb(r, g, b)`` / ``rgba(r, g, b, a)`` — alpha, when present, is discarded: +# the stored value is plain RGB. +_RGB_RE = re.compile( + r"^rgba?\(\s*(\d{1,3})\s*,\s*(\d{1,3})\s*,\s*(\d{1,3})\s*(?:,\s*[\d.]+\s*)?\)$" +) + +COLOR_ERROR = "颜色格式无效,请使用 #RRGGBB 或 rgb(r,g,b),例如 #FF0000 / rgb(255, 0, 0)" + + +def normalize_hex_color(value: str) -> str: + """Normalise any accepted colour spelling to canonical ``#RRGGBB``. + + The database column is fixed-width ``CHAR(7)``, so exactly one spelling has + to survive the write. Accepting the common variants costs little and means + neither the colour picker nor a hand-written API call can produce a value + the column cannot hold. + + Raises: + ValueError: if the value is not a recognisable colour. + """ + text = value.strip() + + rgb_match = _RGB_RE.match(text) + if rgb_match: + channels = [int(part) for part in rgb_match.groups()] + if any(channel > 255 for channel in channels): + raise ValueError(COLOR_ERROR) + return "#{:02X}{:02X}{:02X}".format(*channels) + + hex_match = _HEX_RE.match(text) + if hex_match: + digits = hex_match.group(1) + if len(digits) == 3: # #abc -> #AABBCC + digits = "".join(char * 2 for char in digits) + return "#" + digits.upper() + + raise ValueError(COLOR_ERROR) + + +class TypographyMixin: + """Shared handling of the two typography columns, ``font_size`` / + ``font_color``, for every schema that carries them. + + A plain mixin rather than a :class:`BaseModel` subclass: pydantic v2 + collects ``field_validator`` / ``field_serializer`` from non-model bases in + the MRO, so a schema picks these up with ``class Foo(TypographyMixin, + BaseModel)`` and keeps its own base-model configuration. + + Two rules, both about making the wire format unsurprising: + + * **In** — any accepted colour spelling is normalised to ``#RRGGBB``, so a + ``CHAR(7)`` column always receives something it can hold. + * **Out** — ``font_size`` is emitted as a JSON *number*, not the string + pydantic produces for ``Decimal`` by default. A client should not have to + parse a font size before putting it in a CSS rule. + """ + + @field_validator("font_color", mode="before", check_fields=False) + @classmethod + def _normalise_color(cls, value: object) -> object: + # Non-strings pass through so an omitted colour stays omitted. The + # colour normaliser tolerates a None-valued optional field this way. + if isinstance(value, str): + return normalize_hex_color(value) + return value + + @field_serializer("font_size", check_fields=False) + def _serialize_font_size(self, value: Decimal) -> float: + return float(value) diff --git a/backend/app/schemas/paper_template.py b/backend/app/schemas/paper_template.py new file mode 100644 index 0000000..d21e139 --- /dev/null +++ b/backend/app/schemas/paper_template.py @@ -0,0 +1,123 @@ +"""Request/response schemas for paper templates (模板管理).""" + +from datetime import datetime +from decimal import Decimal + +from pydantic import BaseModel, ConfigDict, Field + +from app.models.paper_template import PaperTemplate +from app.models.template_field import TemplateField +from app.schemas.common import TypographyMixin + + +class TemplateFieldInput(BaseModel): + """One placement of a library field inside a template. + + The client sends the whole ordered selection on create and on update; the + server replaces the stored rows with it. ``sort`` is a plain ascending + integer and is the *only* thing that decides render order. + """ + + field_id: int + sort: int = Field(default=0, ge=-1_000_000, le=1_000_000) + + +class TemplateFieldRead(TypographyMixin, BaseModel): + """A template's field, flattened with the library row it points at. + + Persistence keeps the two apart — a placement row plus the library row it + references — but a client rendering an outline wants one flat record, so + the typography is denormalised into the response. That is also what keeps + the outline in sync: there is only ever one copy of a field's name and + styling, in the library. + """ + + id: int + field_id: int + sort: int + name: str + level: int + font_size: Decimal + font_color: str + + @classmethod + def from_model(cls, item: TemplateField) -> "TemplateFieldRead": + """Flatten a placement row and its library row into one record.""" + return cls( + id=item.id, + field_id=item.field_id, + sort=item.sort, + name=item.field.name, + level=item.field.level, + font_size=item.field.font_size, + font_color=item.field.font_color, + ) + + +class PaperTemplateBase(BaseModel): + """Shared body of the create/update payloads.""" + + name: str = Field(min_length=1, max_length=255) + abstract: str | None = None + + +class PaperTemplateCreate(PaperTemplateBase): + """Payload for ``POST /templates``. + + The field selection is free: any subset, any order, repeats allowed. The + server stores it as given and orders it by ``sort`` on read. + """ + + fields: list[TemplateFieldInput] = Field(default_factory=list) + + +class PaperTemplateUpdate(BaseModel): + """Payload for ``PATCH /templates/{id}``. + + ``fields`` is a full replacement when present — omit it to leave the + selection untouched. + """ + + name: str | None = Field(default=None, min_length=1, max_length=255) + abstract: str | None = None + fields: list[TemplateFieldInput] | None = None + + +class PaperTemplateListItem(BaseModel): + """A template as it appears in the list table — no field rows.""" + + model_config = ConfigDict(from_attributes=True) + + id: int + name: str + abstract: str | None + field_count: int + created_at: datetime + updated_at: datetime + + +class PaperTemplateRead(BaseModel): + """A template with its outline, already ordered by ``sort``.""" + + model_config = ConfigDict(from_attributes=True) + + id: int + name: str + abstract: str | None + #: The outline, ordered by ``sort`` — the ORM relationship sets that + #: ordering, so this list is display-ready with no client-side sorting. + fields: list[TemplateFieldRead] + created_at: datetime + updated_at: datetime + + @classmethod + def from_model(cls, template: PaperTemplate) -> "PaperTemplateRead": + """Build the response from a template and its placement rows.""" + return cls( + id=template.id, + name=template.name, + abstract=template.abstract, + fields=[TemplateFieldRead.from_model(item) for item in template.items], + created_at=template.created_at, + updated_at=template.updated_at, + ) diff --git a/backend/app/schemas/section_field.py b/backend/app/schemas/section_field.py new file mode 100644 index 0000000..9bfbe13 --- /dev/null +++ b/backend/app/schemas/section_field.py @@ -0,0 +1,60 @@ +"""Request/response schemas for the section-field library (字段管理).""" + +from datetime import datetime +from decimal import Decimal + +from pydantic import BaseModel, ConfigDict, Field + +from app.schemas.common import TypographyMixin + + +class SectionFieldBase(TypographyMixin, BaseModel): + """Fields a client may set on a section field.""" + + #: Display name, numbering included — e.g. ``"1. Introduction"``. Stored + #: and rendered verbatim; the API never rewrites it. + name: str = Field(min_length=1, max_length=255) + + #: Heading depth. 1 = "1.", 2 = "1.1", 3 = "1.1.1". Rendering hint only. + level: int = Field(default=1, ge=1, le=9) + + #: Font size in points. Fractional because 五号 = 10.5pt. + font_size: Decimal = Field(default=Decimal("12.0"), gt=0, le=99) + + #: ``#RRGGBB``. ``rgb(...)`` and shorthand forms are normalised on write. + font_color: str = Field(default="#000000", max_length=32) + + +class SectionFieldCreate(SectionFieldBase): + """Payload for ``POST /section-fields``.""" + + +class SectionFieldUpdate(TypographyMixin, BaseModel): + """Payload for ``PATCH /section-fields/{id}`` — every part optional. + + The colour normaliser tolerates ``None`` here; the validator returns + non-strings untouched, so an omitted colour stays omitted. + """ + + name: str | None = Field(default=None, min_length=1, max_length=255) + level: int | None = Field(default=None, ge=1, le=9) + font_size: Decimal | None = Field(default=None, gt=0, le=99) + font_color: str | None = Field(default=None, max_length=32) + + +class SectionFieldRead(SectionFieldBase): + """A stored section field.""" + + model_config = ConfigDict(from_attributes=True) + + id: int + created_at: datetime + updated_at: datetime + + +class SectionFieldRef(BaseModel): + """How many templates currently place a field — used to explain a refusal.""" + + field_id: int + field_name: str + template_count: int diff --git a/backend/scripts/seed.py b/backend/scripts/seed.py new file mode 100644 index 0000000..2aae4ca --- /dev/null +++ b/backend/scripts/seed.py @@ -0,0 +1,197 @@ +"""Seed the section-field library and a couple of starter templates. + +Idempotent: fields are matched by name and templates by name, so running it +twice adds nothing. ``--reset`` empties the three tables first, which is the +way to get back to the seeded baseline after experimenting. + +Usage (from ``backend/``):: + + .venv/bin/python scripts/seed.py + .venv/bin/python scripts/seed.py --reset +""" + +from __future__ import annotations + +import argparse +import sys +from decimal import Decimal +from pathlib import Path + +# Running this as a plain script puts scripts/ on sys.path, not backend/, so +# `import app` would fail. Anchor to the backend directory instead. +BACKEND_DIR = Path(__file__).resolve().parents[1] +if str(BACKEND_DIR) not in sys.path: + sys.path.insert(0, str(BACKEND_DIR)) + +from sqlalchemy import delete, select # noqa: E402 +from sqlalchemy.orm import Session # noqa: E402 + +from app.db.session import SessionLocal # noqa: E402 +from app.models import PaperTemplate, SectionField, TemplateField # noqa: E402 + +# --- the field library ------------------------------------------------------- +# +# (name, level, font_size, font_color) +# +# The number in the name is written by hand and stored verbatim — nothing in +# the system derives or rewrites it. Level is only a rendering hint: it decides +# indentation in the outline, never parentage. +# +# Colours: headings are near-black, running text is black, and the two +# unnumbered front/back sections (Abstract, References) are grey so they read +# as apparatus rather than as body sections. +BODY = "#000000" +HEADING = "#1F1F1F" +SUB = "#404040" +GREY = "#595959" + +FIELDS: list[tuple[str, int, Decimal, str]] = [ + ("0 Abstract", 1, Decimal("10.5"), GREY), + ("1 Introduction", 1, Decimal("12.0"), HEADING), + ("1.1 Research Background", 2, Decimal("10.5"), SUB), + ("1.2 Problem Statement", 2, Decimal("10.5"), SUB), + ("1.3 Contributions", 2, Decimal("10.5"), SUB), + ("2 Related Work", 1, Decimal("12.0"), HEADING), + ("2.1 Prior Approaches", 2, Decimal("10.5"), SUB), + ("2.2 Limitations of Existing Work", 2, Decimal("10.5"), SUB), + ("3 Method", 1, Decimal("12.0"), HEADING), + ("3.1 Problem Formulation", 2, Decimal("10.5"), SUB), + ("3.2 Framework Overview", 2, Decimal("10.5"), SUB), + ("3.3 Implementation Details", 2, Decimal("10.5"), SUB), + ("4 Experiments", 1, Decimal("12.0"), HEADING), + ("4.1 Datasets", 2, Decimal("10.5"), SUB), + ("4.2 Experimental Setup", 2, Decimal("10.5"), SUB), + ("4.3 Main Results", 2, Decimal("10.5"), SUB), + ("4.4 Ablation Study", 2, Decimal("10.5"), SUB), + ("5 Discussion", 1, Decimal("12.0"), BODY), + ("6 Conclusion", 1, Decimal("12.0"), BODY), + ("7 References", 1, Decimal("10.5"), GREY), +] + +# --- starter templates ------------------------------------------------------- +# +# Each entry is (template name, abstract, [field names in display order]). +# `sort` is assigned from the list position (1, 2, 3 ...), which is exactly +# what the UI does when the user picks fields and orders them. +TEMPLATES: list[tuple[str, str, list[str]]] = [ + ( + "标准学术论文(通用)", + "完整的通用学术论文骨架,含摘要、引言、相关工作、方法、实验、讨论与结论。" + "适合期刊或会议长文;先按此模板把结构填满,再按目标期刊微调字段。", + [name for name, *_ in FIELDS], + ), + ( + "四段式短文", + "紧凑的四段式结构,只保留摘要、引言、方法、实验与结论。" + "适合短文、技术报告或初稿阶段的结构搭建。", + [ + "0 Abstract", + "1 Introduction", + "3 Method", + "4 Experiments", + "6 Conclusion", + ], + ), + ( + "方法创新型论文", + "以方法贡献为主线的结构,弱化相关工作、强化方法细节与消融实验。" + "适合以新框架、新算法为主要贡献的投稿。", + [ + "0 Abstract", + "1 Introduction", + "1.2 Problem Statement", + "1.3 Contributions", + "3 Method", + "3.1 Problem Formulation", + "3.2 Framework Overview", + "3.3 Implementation Details", + "4 Experiments", + "4.2 Experimental Setup", + "4.3 Main Results", + "4.4 Ablation Study", + "6 Conclusion", + ], + ), +] + + +def reset(db: Session) -> None: + """Empty the three tables in dependency order.""" + db.execute(delete(TemplateField)) + db.execute(delete(PaperTemplate)) + db.execute(delete(SectionField)) + db.commit() + + +def seed_fields(db: Session) -> dict[str, SectionField]: + """Insert any missing library fields and return name -> field.""" + existing = {field.name: field for field in db.scalars(select(SectionField)).all()} + + created = 0 + for name, level, font_size, font_color in FIELDS: + if name in existing: + continue + field = SectionField( + name=name, level=level, font_size=font_size, font_color=font_color + ) + db.add(field) + existing[name] = field + created += 1 + + db.commit() + for field in existing.values(): + db.refresh(field) + + print(f" section_field : {created} created, {len(existing) - created} already present") + return existing + + +def seed_templates(db: Session, fields: dict[str, SectionField]) -> None: + """Insert any missing starter templates.""" + existing = {name for name in db.scalars(select(PaperTemplate.name)).all()} + + created = 0 + for name, abstract, field_names in TEMPLATES: + if name in existing: + continue + + missing = [field_name for field_name in field_names if field_name not in fields] + if missing: + raise SystemExit(f"template {name!r} references unknown fields: {missing}") + + template = PaperTemplate(name=name, abstract=abstract) + # sort is the list position: 1-based, ascending, with gaps allowed + # because the UI lets the user type any integer. + template.items = [ + TemplateField(field_id=fields[field_name].id, sort=index) + for index, field_name in enumerate(field_names, start=1) + ] + db.add(template) + created += 1 + + db.commit() + print(f" paper_template: {created} created, {len(existing)} already present") + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--reset", + action="store_true", + help="delete all fields and templates before seeding", + ) + args = parser.parse_args() + + with SessionLocal() as db: + if args.reset: + reset(db) + print(" reset : all rows deleted") + + fields = seed_fields(db) + seed_templates(db, fields) + + print("seed complete") + + +if __name__ == "__main__": + main()