This commit is contained in:
+4
-92
@@ -26,7 +26,7 @@ from .auth import (
|
||||
verify_password,
|
||||
)
|
||||
from .db import create_schema, get_db
|
||||
from .models import AppState, Folder, Session, Tag, Task, TaskList, TaskTag, User, utcnow
|
||||
from .models import AppState, Folder, Session, Task, TaskList, User, utcnow
|
||||
from .mvp import audit
|
||||
from .mvp import router as mvp_router
|
||||
from .schemas import (
|
||||
@@ -40,8 +40,6 @@ from .schemas import (
|
||||
LoginRequest,
|
||||
NameUpdate,
|
||||
SessionOut,
|
||||
TagCreate,
|
||||
TagOut,
|
||||
TaskCreate,
|
||||
TaskDetailOut,
|
||||
TaskOut,
|
||||
@@ -186,15 +184,11 @@ async def bootstrap_data(user: User = Depends(current_user), db: AsyncSession =
|
||||
.where(TaskList.user_id == user.id, TaskList.deleted_at.is_(None))
|
||||
.order_by(TaskList.is_inbox.desc(), TaskList.position, TaskList.created_at)
|
||||
)).all())
|
||||
tags = list((await db.scalars(
|
||||
select(Tag).where(Tag.user_id == user.id).order_by(Tag.name, Tag.id)
|
||||
)).all())
|
||||
inbox = next((item for item in lists if item.is_inbox), None)
|
||||
return {
|
||||
"user": UserOut.model_validate(user),
|
||||
"folders": [FolderOut.model_validate(item) for item in folders],
|
||||
"lists": [ListOut.model_validate(item) for item in lists],
|
||||
"tags": [TagOut.model_validate(item) for item in tags],
|
||||
"inbox_id": inbox.id if inbox else None,
|
||||
}
|
||||
|
||||
@@ -420,51 +414,6 @@ async def restore_list(
|
||||
return item
|
||||
|
||||
|
||||
@app.post("/api/v1/tags", response_model=TagOut, status_code=201)
|
||||
async def create_tag(
|
||||
payload: TagCreate,
|
||||
user: User = Depends(current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
existing = await db.scalar(
|
||||
select(Tag.id).where(Tag.user_id == user.id, func.lower(Tag.name) == payload.name.lower())
|
||||
)
|
||||
if existing:
|
||||
raise HTTPException(status_code=409, detail="标签名称已存在")
|
||||
tag = Tag(user_id=user.id, **payload.model_dump())
|
||||
db.add(tag)
|
||||
await db.commit()
|
||||
await db.refresh(tag)
|
||||
return tag
|
||||
|
||||
|
||||
@app.get("/api/v1/tags", response_model=list[TagOut])
|
||||
async def list_tags(user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
return list(
|
||||
(await db.scalars(select(Tag).where(Tag.user_id == user.id).order_by(Tag.name, Tag.id))).all()
|
||||
)
|
||||
|
||||
|
||||
async def _validate_tags(
|
||||
db: AsyncSession, user_id: UUID, tag_ids: list[UUID] | None
|
||||
) -> list[UUID] | None:
|
||||
if tag_ids is None:
|
||||
return None
|
||||
unique_ids = list(dict.fromkeys(tag_ids))
|
||||
if not unique_ids:
|
||||
return []
|
||||
found = set(
|
||||
(await db.scalars(select(Tag.id).where(Tag.user_id == user_id, Tag.id.in_(unique_ids)))).all()
|
||||
)
|
||||
if found != set(unique_ids):
|
||||
raise HTTPException(status_code=404, detail="标签不存在")
|
||||
return unique_ids
|
||||
|
||||
|
||||
async def _replace_tags(db: AsyncSession, task_ids: list[UUID], tag_ids: list[UUID]) -> None:
|
||||
await db.execute(delete(TaskTag).where(TaskTag.task_id.in_(task_ids)))
|
||||
db.add_all(TaskTag(task_id=task_id, tag_id=tag_id) for task_id in task_ids for tag_id in tag_ids)
|
||||
|
||||
|
||||
def _encode_cursor(created_at: datetime, task_id: UUID) -> str:
|
||||
raw = json.dumps([created_at.isoformat(), str(task_id)]).encode()
|
||||
@@ -487,7 +436,6 @@ async def create_task(
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
await _owned_list(db, user.id, payload.list_id)
|
||||
tag_ids = await _validate_tags(db, user.id, payload.tag_ids)
|
||||
if payload.parent_id:
|
||||
parent = await db.scalar(
|
||||
select(Task).where(
|
||||
@@ -500,13 +448,11 @@ async def create_task(
|
||||
)
|
||||
if parent is None:
|
||||
raise HTTPException(status_code=400, detail="父任务必须是同一清单的顶层任务")
|
||||
data = payload.model_dump(exclude={"tag_ids"})
|
||||
data = payload.model_dump()
|
||||
task = Task(user_id=user.id, **data)
|
||||
db.add(task)
|
||||
await db.flush()
|
||||
audit(db, user.id, "create", "task", task.id)
|
||||
if tag_ids:
|
||||
db.add_all(TaskTag(task_id=task.id, tag_id=tag_id) for tag_id in tag_ids)
|
||||
await db.commit()
|
||||
await db.refresh(task)
|
||||
return task
|
||||
@@ -540,11 +486,6 @@ async def list_tasks(
|
||||
query = query.where(Task.due_at < due_to)
|
||||
if q:
|
||||
pattern = f"%{q}%"
|
||||
tag_match = exists(
|
||||
select(TaskTag.task_id)
|
||||
.join(Tag, Tag.id == TaskTag.tag_id)
|
||||
.where(TaskTag.task_id == Task.id, Tag.user_id == user.id, Tag.name.ilike(pattern))
|
||||
)
|
||||
list_match = exists(
|
||||
select(TaskList.id).where(
|
||||
TaskList.id == Task.list_id,
|
||||
@@ -553,7 +494,7 @@ async def list_tasks(
|
||||
)
|
||||
)
|
||||
query = query.where(
|
||||
or_(Task.title.ilike(pattern), Task.description.ilike(pattern), tag_match, list_match)
|
||||
or_(Task.title.ilike(pattern), Task.description.ilike(pattern), list_match)
|
||||
)
|
||||
total = await db.scalar(select(func.count()).select_from(query.order_by(None).subquery())) or 0
|
||||
ordering = (Task.created_at, Task.id)
|
||||
@@ -577,14 +518,6 @@ async def _task_details(db: AsyncSession, tasks: list[Task]) -> list[TaskDetailO
|
||||
if not tasks:
|
||||
return []
|
||||
task_ids = [task.id for task in tasks]
|
||||
tag_rows = (
|
||||
await db.execute(
|
||||
select(TaskTag.task_id, Tag)
|
||||
.join(Tag, Tag.id == TaskTag.tag_id)
|
||||
.where(TaskTag.task_id.in_(task_ids))
|
||||
.order_by(Tag.name, Tag.id)
|
||||
)
|
||||
).all()
|
||||
subtasks = list(
|
||||
(
|
||||
await db.scalars(
|
||||
@@ -594,22 +527,17 @@ async def _task_details(db: AsyncSession, tasks: list[Task]) -> list[TaskDetailO
|
||||
)
|
||||
).all()
|
||||
)
|
||||
tags_by_task: dict[UUID, list[Tag]] = defaultdict(list)
|
||||
subtasks_by_task: dict[UUID, list[Task]] = defaultdict(list)
|
||||
for task_id, tag in tag_rows:
|
||||
tags_by_task[task_id].append(tag)
|
||||
for subtask in subtasks:
|
||||
subtasks_by_task[subtask.parent_id].append(subtask)
|
||||
return [
|
||||
TaskDetailOut(
|
||||
**TaskOut.model_validate(task).model_dump(),
|
||||
tags=tags_by_task[task.id],
|
||||
subtasks=subtasks_by_task[task.id],
|
||||
)
|
||||
for task in tasks
|
||||
]
|
||||
|
||||
|
||||
async def _task_detail(db: AsyncSession, task: Task) -> TaskDetailOut:
|
||||
return (await _task_details(db, [task]))[0]
|
||||
|
||||
@@ -637,7 +565,6 @@ async def update_task(
|
||||
):
|
||||
data = payload.model_dump(exclude_unset=True)
|
||||
expected_version = data.pop("version")
|
||||
tag_ids = await _validate_tags(db, user.id, data.pop("tag_ids", None))
|
||||
if "list_id" in data:
|
||||
await _owned_list(db, user.id, data["list_id"])
|
||||
parent_id = await db.scalar(
|
||||
@@ -670,8 +597,6 @@ async def update_task(
|
||||
if exists_id:
|
||||
raise HTTPException(status_code=409, detail="任务已被更新,请刷新后重试")
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
if tag_ids is not None:
|
||||
await _replace_tags(db, [task_id], tag_ids)
|
||||
if "list_id" in data and task.parent_id is None:
|
||||
await db.execute(
|
||||
update(Task)
|
||||
@@ -773,16 +698,6 @@ async def permanently_delete_task(
|
||||
)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail="回收站中不存在该任务")
|
||||
ids = list(
|
||||
(
|
||||
await db.scalars(
|
||||
select(Task.id).where(
|
||||
or_(Task.id == task.id, Task.parent_id == task.id), Task.user_id == user.id
|
||||
)
|
||||
)
|
||||
).all()
|
||||
)
|
||||
await db.execute(delete(TaskTag).where(TaskTag.task_id.in_(ids)))
|
||||
await db.execute(delete(Task).where(Task.parent_id == task.id, Task.user_id == user.id))
|
||||
await db.delete(task)
|
||||
await db.commit()
|
||||
@@ -812,8 +727,7 @@ async def batch_update_tasks(
|
||||
standalone_children = [task for task in tasks if task.parent_id is not None]
|
||||
if standalone_children:
|
||||
raise HTTPException(status_code=400, detail="子任务不能脱离父任务单独移动")
|
||||
tag_ids = await _validate_tags(db, user.id, payload.tag_ids)
|
||||
changes = payload.model_dump(exclude_unset=True, exclude={"task_ids", "tag_ids", "soft_delete"})
|
||||
changes = payload.model_dump(exclude_unset=True, exclude={"task_ids", "soft_delete"})
|
||||
if payload.soft_delete:
|
||||
changes["deleted_at"] = utcnow()
|
||||
if changes:
|
||||
@@ -856,8 +770,6 @@ async def batch_update_tasks(
|
||||
updated_at=utcnow(),
|
||||
)
|
||||
)
|
||||
if tag_ids is not None:
|
||||
await _replace_tags(db, task_ids, tag_ids)
|
||||
await db.commit()
|
||||
return BatchResult(updated=len(task_ids))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user