This commit is contained in:
+52
-10
@@ -45,6 +45,7 @@ from .schemas import (
|
||||
TaskDetailOut,
|
||||
TaskOut,
|
||||
TaskPage,
|
||||
TaskReorder,
|
||||
TaskUpdate,
|
||||
UserOut,
|
||||
)
|
||||
@@ -434,16 +435,16 @@ async def restore_list(
|
||||
|
||||
|
||||
|
||||
def _encode_cursor(created_at: datetime, task_id: UUID) -> str:
|
||||
raw = json.dumps([created_at.isoformat(), str(task_id)]).encode()
|
||||
def _encode_cursor(position: int, created_at: datetime, task_id: UUID) -> str:
|
||||
raw = json.dumps([position, created_at.isoformat(), str(task_id)]).encode()
|
||||
return base64.urlsafe_b64encode(raw).decode().rstrip("=")
|
||||
|
||||
|
||||
def _decode_cursor(cursor: str) -> tuple[datetime, UUID]:
|
||||
def _decode_cursor(cursor: str) -> tuple[int, datetime, UUID]:
|
||||
try:
|
||||
raw = base64.urlsafe_b64decode(cursor + "=" * (-len(cursor) % 4))
|
||||
timestamp, task_id = json.loads(raw)
|
||||
return datetime.fromisoformat(timestamp), UUID(task_id)
|
||||
position, timestamp, task_id = json.loads(raw)
|
||||
return int(position), datetime.fromisoformat(timestamp), UUID(task_id)
|
||||
except (ValueError, TypeError, json.JSONDecodeError) as exc:
|
||||
raise HTTPException(status_code=422, detail="无效的游标") from exc
|
||||
|
||||
@@ -468,7 +469,14 @@ async def create_task(
|
||||
if parent is None:
|
||||
raise HTTPException(status_code=400, detail="父任务必须是同一清单的顶层任务")
|
||||
data = payload.model_dump()
|
||||
task = Task(user_id=user.id, **data)
|
||||
parent_filter = Task.parent_id == payload.parent_id if payload.parent_id else Task.parent_id.is_(None)
|
||||
max_position = await db.scalar(select(func.max(Task.position)).where(
|
||||
Task.user_id == user.id,
|
||||
Task.list_id == payload.list_id,
|
||||
parent_filter,
|
||||
Task.deleted_at.is_(None),
|
||||
))
|
||||
task = Task(user_id=user.id, position=(max_position if max_position is not None else -1) + 1, **data)
|
||||
db.add(task)
|
||||
await db.flush()
|
||||
audit(db, user.id, "create", "task", task.id)
|
||||
@@ -516,20 +524,24 @@ async def list_tasks(
|
||||
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)
|
||||
ordering = (Task.position, Task.created_at, Task.id)
|
||||
if page is not None:
|
||||
size = page_size or limit
|
||||
items = list((await db.scalars(query.order_by(*ordering).offset((page - 1) * size).limit(size))).all())
|
||||
return TaskPage(items=await _task_details(db, items), total=total, page=page, page_size=size)
|
||||
if cursor:
|
||||
created_at, task_id = _decode_cursor(cursor)
|
||||
position, created_at, task_id = _decode_cursor(cursor)
|
||||
query = query.where(
|
||||
or_(Task.created_at > created_at, (Task.created_at == created_at) & (Task.id > task_id))
|
||||
or_(
|
||||
Task.position > position,
|
||||
(Task.position == position) & (Task.created_at > created_at),
|
||||
(Task.position == position) & (Task.created_at == created_at) & (Task.id > task_id),
|
||||
)
|
||||
)
|
||||
rows = list((await db.scalars(query.order_by(*ordering).limit(limit + 1))).all())
|
||||
has_more = len(rows) > limit
|
||||
items = rows[:limit]
|
||||
next_cursor = _encode_cursor(items[-1].created_at, items[-1].id) if has_more else None
|
||||
next_cursor = _encode_cursor(items[-1].position, items[-1].created_at, items[-1].id) if has_more else None
|
||||
return TaskPage(items=await _task_details(db, items), next_cursor=next_cursor, total=total, page=1, page_size=limit)
|
||||
|
||||
|
||||
@@ -561,6 +573,36 @@ async def _task_detail(db: AsyncSession, task: Task) -> TaskDetailOut:
|
||||
return (await _task_details(db, [task]))[0]
|
||||
|
||||
|
||||
@app.put("/api/v1/tasks/reorder", status_code=204)
|
||||
async def reorder_tasks(
|
||||
payload: TaskReorder,
|
||||
user: User = Depends(current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
rows = list((await db.scalars(select(Task).where(
|
||||
Task.id.in_(payload.task_ids), Task.user_id == user.id, Task.deleted_at.is_(None)
|
||||
))).all())
|
||||
if len(rows) != len(payload.task_ids):
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
parent_scopes = {row.parent_id for row in rows}
|
||||
if len(parent_scopes) != 1:
|
||||
raise HTTPException(status_code=400, detail="只能调整同一层级任务的顺序")
|
||||
parent_id = next(iter(parent_scopes))
|
||||
scope_query = select(Task).where(
|
||||
Task.user_id == user.id,
|
||||
Task.deleted_at.is_(None),
|
||||
Task.parent_id.is_(None) if parent_id is None else Task.parent_id == parent_id,
|
||||
).order_by(Task.position, Task.created_at, Task.id)
|
||||
scope_rows = list((await db.scalars(scope_query)).all())
|
||||
requested = set(payload.task_ids)
|
||||
ordered_rows = iter([next(row for row in rows if row.id == task_id) for task_id in payload.task_ids])
|
||||
merged = [next(ordered_rows) if row.id in requested else row for row in scope_rows]
|
||||
for position, row in enumerate(merged):
|
||||
row.position = position
|
||||
await db.commit()
|
||||
return Response(status_code=204)
|
||||
|
||||
|
||||
@app.get("/api/v1/tasks/{task_id}", response_model=TaskDetailOut)
|
||||
async def get_task(
|
||||
task_id: UUID,
|
||||
|
||||
@@ -162,6 +162,7 @@ class Habit(Base):
|
||||
interval_days: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
start_date: Mapped[date] = mapped_column(Date)
|
||||
archived_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
position: Mapped[int] = mapped_column(Integer, default=0)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow)
|
||||
|
||||
|
||||
|
||||
+36
-5
@@ -525,6 +525,16 @@ class HabitCreate(BaseModel):
|
||||
return self
|
||||
|
||||
|
||||
class HabitReorder(BaseModel):
|
||||
habit_ids: list[UUID] = Field(min_length=1)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def unique_ids(self):
|
||||
if len(self.habit_ids) != len(set(self.habit_ids)):
|
||||
raise ValueError("habit_ids must be unique")
|
||||
return self
|
||||
|
||||
|
||||
class HabitLogInput(BaseModel):
|
||||
day: date
|
||||
value: float = Field(gt=0)
|
||||
@@ -541,7 +551,7 @@ class PauseInput(BaseModel):
|
||||
|
||||
|
||||
def habit_dict(h):
|
||||
return {"id": h.id, "name": h.name, "kind": h.kind, "target": h.target, "max_value": h.max_value, "schedule_type": h.schedule_type, "weekdays": [int(x) for x in h.weekdays.split(",")] if h.weekdays else None, "month_days": [int(x) for x in h.month_days.split(",")] if h.month_days else None, "interval_days": h.interval_days, "start_date": h.start_date, "archived_at": h.archived_at}
|
||||
return {"id": h.id, "name": h.name, "kind": h.kind, "target": h.target, "max_value": h.max_value, "schedule_type": h.schedule_type, "weekdays": [int(x) for x in h.weekdays.split(",")] if h.weekdays else None, "month_days": [int(x) for x in h.month_days.split(",")] if h.month_days else None, "interval_days": h.interval_days, "start_date": h.start_date, "archived_at": h.archived_at, "position": h.position}
|
||||
|
||||
|
||||
async def owned_habit(db, user_id, habit_id):
|
||||
@@ -552,14 +562,34 @@ async def owned_habit(db, user_id, habit_id):
|
||||
|
||||
@router.post("/habits", status_code=201)
|
||||
async def create_habit(payload: HabitCreate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
row = Habit(user_id=user.id, **payload.model_dump(exclude={"weekdays", "month_days"}), weekdays=",".join(map(str, payload.weekdays)) if payload.weekdays else None, month_days=",".join(map(str, payload.month_days)) if payload.month_days else None)
|
||||
max_position = await db.scalar(select(func.max(Habit.position)).where(Habit.user_id == user.id))
|
||||
row = Habit(user_id=user.id, position=(max_position if max_position is not None else -1) + 1, **payload.model_dump(exclude={"weekdays", "month_days"}), weekdays=",".join(map(str, payload.weekdays)) if payload.weekdays else None, month_days=",".join(map(str, payload.month_days)) if payload.month_days else None)
|
||||
db.add(row); await db.commit(); await db.refresh(row); return habit_dict(row)
|
||||
|
||||
|
||||
@router.get("/habits")
|
||||
async def list_habits(archived: bool = False, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
condition = Habit.archived_at.is_not(None) if archived else Habit.archived_at.is_(None)
|
||||
return [habit_dict(h) for h in (await db.scalars(select(Habit).where(Habit.user_id == user.id, condition).order_by(Habit.created_at))).all()]
|
||||
return [habit_dict(h) for h in (await db.scalars(select(Habit).where(Habit.user_id == user.id, condition).order_by(Habit.position, Habit.created_at))).all()]
|
||||
|
||||
|
||||
@router.put("/habits/reorder", status_code=204)
|
||||
async def reorder_habits(payload: HabitReorder, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
rows = list((await db.scalars(select(Habit).where(
|
||||
Habit.id.in_(payload.habit_ids), Habit.user_id == user.id, Habit.archived_at.is_(None)
|
||||
))).all())
|
||||
if len(rows) != len(payload.habit_ids):
|
||||
raise HTTPException(404, "习惯不存在")
|
||||
scope_rows = list((await db.scalars(select(Habit).where(
|
||||
Habit.user_id == user.id, Habit.archived_at.is_(None)
|
||||
).order_by(Habit.position, Habit.created_at))).all())
|
||||
requested = set(payload.habit_ids)
|
||||
ordered_rows = iter([next(row for row in rows if row.id == habit_id) for habit_id in payload.habit_ids])
|
||||
merged = [next(ordered_rows) if row.id in requested else row for row in scope_rows]
|
||||
for position, row in enumerate(merged):
|
||||
row.position = position
|
||||
await db.commit()
|
||||
return Response(status_code=204)
|
||||
|
||||
|
||||
@router.patch("/habits/{habit_id}")
|
||||
@@ -623,7 +653,7 @@ def scheduled(h, day):
|
||||
@router.get("/habits/grid")
|
||||
async def habits_grid(week: date, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
start = week - timedelta(days=week.weekday()); days = [start + timedelta(days=i) for i in range(7)]
|
||||
habits = list((await db.scalars(select(Habit).where(Habit.user_id == user.id, Habit.archived_at.is_(None)).order_by(Habit.created_at))).all())
|
||||
habits = list((await db.scalars(select(Habit).where(Habit.user_id == user.id, Habit.archived_at.is_(None)).order_by(Habit.position, Habit.created_at))).all())
|
||||
if not habits:
|
||||
return {"days": days, "habits": []}
|
||||
habit_ids = [habit.id for habit in habits]
|
||||
@@ -755,7 +785,7 @@ async def export_json(user: User = Depends(current_user), db: AsyncSession = Dep
|
||||
"folders": [serialize(x, ["id", "name", "position", "deleted_at"]) for x in folders],
|
||||
"lists": [serialize(x, ["id", "folder_id", "name", "is_inbox", "position", "deleted_at"]) for x in lists],
|
||||
"tasks": [serialize(x, ["id", "list_id", "parent_id", "title", "description", "priority", "completed", "due_at", "external_id", "deleted_at"]) for x in tasks],
|
||||
"habits": [serialize(x, ["id", "name", "kind", "target", "max_value", "schedule_type", "weekdays", "month_days", "interval_days", "start_date", "archived_at"]) for x in habits],
|
||||
"habits": [serialize(x, ["id", "name", "kind", "target", "max_value", "schedule_type", "weekdays", "month_days", "interval_days", "start_date", "archived_at", "position"]) for x in habits],
|
||||
"countdowns": [serialize(x, ["id", "title", "event_date", "calendar_mode", "lunar_month", "lunar_day", "ignore_year", "kind", "repeat_rule", "icon", "pinned", "archived_at", "created_at", "updated_at"]) for x in countdowns],
|
||||
}
|
||||
|
||||
@@ -835,6 +865,7 @@ async def restore_json(payload: dict, mode: str = Query("merge", pattern="^(merg
|
||||
interval_days=raw.get("interval_days"),
|
||||
start_date=date.fromisoformat(raw["start_date"]),
|
||||
archived_at=datetime.fromisoformat(raw["archived_at"]) if raw.get("archived_at") else None,
|
||||
position=raw.get("position", 0),
|
||||
)
|
||||
db.add(row)
|
||||
existing_countdown_ids = set((await db.scalars(select(Countdown.id).where(Countdown.user_id == user.id))).all())
|
||||
|
||||
@@ -96,6 +96,16 @@ class TaskUpdate(BaseModel):
|
||||
return self
|
||||
|
||||
|
||||
class TaskReorder(BaseModel):
|
||||
task_ids: list[UUID] = Field(min_length=1)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def unique_ids(self):
|
||||
if len(self.task_ids) != len(set(self.task_ids)):
|
||||
raise ValueError("task_ids must be unique")
|
||||
return self
|
||||
|
||||
|
||||
class TaskOut(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
id: UUID
|
||||
|
||||
Reference in New Issue
Block a user