import calendar import csv import io import re from datetime import UTC, date, datetime, time, timedelta from pathlib import Path from uuid import UUID from zoneinfo import ZoneInfo from fastapi import APIRouter, Depends, File, HTTPException, Query, Response, UploadFile from fastapi.responses import FileResponse from pydantic import BaseModel, Field, model_validator from sqlalchemy import case, delete, func, select, update from sqlalchemy.ext.asyncio import AsyncSession from .auth import current_user from .config import get_settings from .db import get_db from .lunar_support import ( lunar_label_with_year, lunar_to_solar_safe, next_lunar_occurrence, solar_to_lunar_parts, solar_to_lunar_text, ) from .models import ( Attachment, AuditLog, Countdown, Folder, Habit, HabitLog, HabitPause, RecurrenceException, RecurrenceTemplate, Task, TaskList, User, new_id, utcnow, ) router = APIRouter(prefix="/api/v1") def audit(db: AsyncSession, user_id: UUID, action: str, entity_type: str, entity_id=None, **details): db.add(AuditLog(user_id=user_id, action=action, entity_type=entity_type, entity_id=entity_id, details=details)) class RecurrenceCreate(BaseModel): task_id: UUID rrule: str = Field(min_length=5, max_length=1000) class RecurrenceChange(BaseModel): title: str | None = Field(None, min_length=1, max_length=500) due_at: datetime | None = None rrule: str | None = None class OccurrenceComplete(BaseModel): occurrence_at: datetime _RRULE_PART = re.compile(r"^[A-Z]+=[A-Z0-9,+-]+$") _WEEKDAYS = {"MO": 0, "TU": 1, "WE": 2, "TH": 3, "FR": 4, "SA": 5, "SU": 6} def parse_rrule(value: str) -> dict[str, str]: parts = {} for part in value.upper().split(";"): if not _RRULE_PART.fullmatch(part): raise HTTPException(422, "无效的 RRULE") key, val = part.split("=", 1) parts[key] = val allowed = {"FREQ", "INTERVAL", "BYDAY", "BYMONTHDAY", "BYMONTH", "COUNT", "UNTIL"} if set(parts) - allowed: raise HTTPException(422, "重复规则包含不支持的字段") if parts.get("FREQ") not in {"DAILY", "WEEKLY", "MONTHLY", "YEARLY"}: raise HTTPException(422, "仅支持 DAILY、WEEKLY、MONTHLY、YEARLY") if "BYDAY" in parts: weekdays = parts["BYDAY"].split(",") if not weekdays or any(day not in _WEEKDAYS for day in weekdays): raise HTTPException(422, "无效的重复星期") try: month_days = [int(day) for day in parts.get("BYMONTHDAY", "").split(",") if day] months = [int(month) for month in parts.get("BYMONTH", "").split(",") if month] if any(day < 1 or day > 31 for day in month_days) or any(month < 1 or month > 12 for month in months): raise ValueError if "UNTIL" in parts: datetime.fromisoformat(parts["UNTIL"]) if "INTERVAL" in parts and int(parts["INTERVAL"]) < 1: raise ValueError if "COUNT" in parts and int(parts["COUNT"]) < 1: raise ValueError except ValueError as exc: raise HTTPException(422, "无效的 RRULE 数字") from exc return parts def occurrences(rule: str, starts: datetime, start: datetime, end: datetime, cutoff=None): parts = parse_rrule(rule) interval = int(parts.get("INTERVAL", 1)) count = int(parts.get("COUNT", 100000)) if "UNTIL" in parts: until = datetime.fromisoformat(parts["UNTIL"]) until = until.replace(tzinfo=UTC) if until.tzinfo is None else until.astimezone(UTC) else: until = end if starts.tzinfo is None: starts = starts.replace(tzinfo=UTC) if start.tzinfo is None: start = start.replace(tzinfo=UTC) if end.tzinfo is None: end = end.replace(tzinfo=UTC) if cutoff is not None: cutoff = cutoff.replace(tzinfo=UTC) if cutoff.tzinfo is None else cutoff.astimezone(UTC) until = min(until, cutoff) if until.tzinfo is None: until = until.replace(tzinfo=UTC) result = [] cursor = starts emitted = 0 while cursor <= until and emitted < count: include = False if parts["FREQ"] == "DAILY": include = (cursor.date() - starts.date()).days % interval == 0 elif parts["FREQ"] == "WEEKLY": days = {_WEEKDAYS[x] for x in parts.get("BYDAY", list(_WEEKDAYS)[starts.weekday()]).split(",")} include = cursor.weekday() in days and ((cursor.date() - starts.date()).days // 7) % interval == 0 elif parts["FREQ"] == "MONTHLY": month_delta = (cursor.year - starts.year) * 12 + cursor.month - starts.month month_days = {int(x) for x in parts.get("BYMONTHDAY", str(starts.day)).split(",")} include = month_delta % interval == 0 and cursor.day in month_days else: years = cursor.year - starts.year months = {int(x) for x in parts.get("BYMONTH", str(starts.month)).split(",")} month_days = {int(x) for x in parts.get("BYMONTHDAY", str(starts.day)).split(",")} include = years % interval == 0 and cursor.month in months and cursor.day in month_days if include and cursor >= starts: emitted += 1 if start <= cursor <= end: result.append(cursor) cursor += timedelta(days=1) return result def is_occurrence(rule: str, starts: datetime, at: datetime) -> bool: """Return True when `at` is the exact timestamp of an occurrence this rule generates.""" parts = parse_rrule(rule) if starts.tzinfo is None: starts = starts.replace(tzinfo=UTC) else: starts = starts.astimezone(UTC) if at.tzinfo is None: at = at.replace(tzinfo=UTC) else: at = at.astimezone(UTC) if at < starts: return False if "UNTIL" in parts: until = datetime.fromisoformat(parts["UNTIL"]) until = until.replace(tzinfo=UTC) if until.tzinfo is None else until.astimezone(UTC) if at > until: return False # Exact-match validation via the same canonical generator used for recurrence # mutations, so COUNT/UNTIL, time-of-day, BYDAY/BYMONTHDAY and sparse yearly # rules all share one behavior. candidates = occurrences(rule, starts, starts, at, None) return at in candidates def is_occurrence_utc(rule: str, starts: datetime, at: datetime) -> bool: return is_occurrence(rule, starts, at) async def owned_task(db, user_id, task_id): task = await db.scalar(select(Task).where(Task.id == task_id, Task.user_id == user_id, Task.deleted_at.is_(None))) if not task: raise HTTPException(404, "任务不存在") return task async def owned_recurrence(db, user_id, recurrence_id): row = await db.scalar(select(RecurrenceTemplate).where(RecurrenceTemplate.id == recurrence_id, RecurrenceTemplate.user_id == user_id)) if not row: raise HTTPException(404, "重复规则不存在") return row def ensure_real_occurrence(template: RecurrenceTemplate, occurrence_at: datetime) -> None: """Reject occurrence_at values that are not valid occurrences of this recurrence rule.""" if not is_occurrence(template.rrule, template.starts_at, occurrence_at): raise HTTPException(422, "occurrence_at 不是该重复规则的有效发生时刻") if template.ends_at: at = occurrence_at.replace(tzinfo=UTC) if occurrence_at.tzinfo is None else occurrence_at.astimezone(UTC) ends = template.ends_at.replace(tzinfo=UTC) if template.ends_at.tzinfo is None else template.ends_at.astimezone(UTC) if at > ends: raise HTTPException(422, "occurrence_at 晚于该重复规则的有效截止时间") @router.get("/tasks/{task_id}/recurrence") async def get_task_recurrence(task_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): await owned_task(db, user.id, task_id) row = await db.scalar(select(RecurrenceTemplate).where( RecurrenceTemplate.task_id == task_id, RecurrenceTemplate.user_id == user.id )) if row is None: return None return {"id": row.id, "task_id": row.task_id, "rrule": row.rrule, "starts_at": row.starts_at, "ends_at": row.ends_at} @router.post("/recurrences", status_code=201) async def create_recurrence(payload: RecurrenceCreate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): task = await owned_task(db, user.id, payload.task_id) if not task.due_at: raise HTTPException(422, "重复任务需要截止时间") parse_rrule(payload.rrule) if await db.scalar(select(RecurrenceTemplate.id).where(RecurrenceTemplate.task_id == task.id)): raise HTTPException(409, "任务已有重复规则") row = RecurrenceTemplate(user_id=user.id, task_id=task.id, rrule=payload.rrule.upper(), starts_at=task.due_at) db.add(row) await db.commit(); await db.refresh(row) return {"id": row.id, "task_id": row.task_id, "rrule": row.rrule, "starts_at": row.starts_at} async def upsert_exception(db, template_id, at): row = await db.scalar(select(RecurrenceException).where(RecurrenceException.template_id == template_id, RecurrenceException.occurrence_at == at)) if not row: row = RecurrenceException(template_id=template_id, occurrence_at=at) db.add(row) await db.flush() return row @router.patch("/recurrences/{recurrence_id}") async def edit_recurrence(recurrence_id: UUID, payload: RecurrenceChange, scope: str = Query("all", pattern="^(this|this-and-future|all)$"), occurrence_at: datetime | None = None, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): template = await owned_recurrence(db, user.id, recurrence_id) task = await owned_task(db, user.id, template.task_id) if scope == "this": if not occurrence_at: raise HTTPException(422, "需要 occurrence_at") ensure_real_occurrence(template, occurrence_at) row = await upsert_exception(db, template.id, occurrence_at) if payload.title is not None: row.title = payload.title if payload.due_at is not None: row.due_at = payload.due_at elif scope == "this-and-future": if not occurrence_at: raise HTTPException(422, "需要 occurrence_at") ensure_real_occurrence(template, occurrence_at) template.ends_at = occurrence_at - timedelta(microseconds=1) if payload.rrule: parse_rrule(payload.rrule) db.add(RecurrenceTemplate(user_id=user.id, task_id=task.id, rrule=payload.rrule, starts_at=payload.due_at or occurrence_at)) elif payload.title: task.title = payload.title else: if payload.rrule: parse_rrule(payload.rrule); template.rrule = payload.rrule.upper() if payload.title is not None: task.title = payload.title if payload.due_at is not None: task.due_at = payload.due_at; template.starts_at = payload.due_at await db.commit() return {"id": template.id, "scope": scope} @router.post("/recurrences/{recurrence_id}/complete") async def complete_occurrence(recurrence_id: UUID, payload: OccurrenceComplete, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): template = await owned_recurrence(db, user.id, recurrence_id) ensure_real_occurrence(template, payload.occurrence_at) row = await upsert_exception(db, template.id, payload.occurrence_at); row.completed = True audit(db, user.id, "complete", "task", template.task_id, occurrence_at=payload.occurrence_at.isoformat()) await db.commit(); return {"completed": True} @router.delete("/recurrences/{recurrence_id}", status_code=204) async def delete_recurrence(recurrence_id: UUID, scope: str = Query("all", pattern="^(this|this-and-future|all)$"), occurrence_at: datetime | None = None, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): template = await owned_recurrence(db, user.id, recurrence_id) if scope == "this": if not occurrence_at: raise HTTPException(422, "需要 occurrence_at") ensure_real_occurrence(template, occurrence_at) row = await upsert_exception(db, template.id, occurrence_at); row.deleted = True elif scope == "this-and-future": if not occurrence_at: raise HTTPException(422, "需要 occurrence_at") ensure_real_occurrence(template, occurrence_at) template.ends_at = occurrence_at - timedelta(microseconds=1) else: await db.delete(template) await db.commit(); return Response(status_code=204) def _clamped_date(year: int, month: int, day: int) -> date: return date(year, month, min(day, calendar.monthrange(year, month)[1])) def countdown_occurrence(event_date: date, repeat_rule: str, today: date) -> date: """Return the first date-only occurrence on or after today.""" if repeat_rule == "none" or event_date >= today: return event_date if repeat_rule == "weekly": days = (event_date.weekday() - today.weekday()) % 7 return today + timedelta(days=days) if repeat_rule == "monthly": candidate = _clamped_date(today.year, today.month, event_date.day) if candidate < today: year = today.year + (today.month == 12) month = 1 if today.month == 12 else today.month + 1 candidate = _clamped_date(year, month, event_date.day) return candidate candidate = _clamped_date(today.year, event_date.month, event_date.day) if candidate < today: candidate = _clamped_date(today.year + 1, event_date.month, event_date.day) return candidate def countdown_status(display_date: date, today: date) -> tuple[int, str]: days = (display_date - today).days if days > 0: return days, f"还有 {days} 天" if days == 0: return 0, "就是今天" return days, f"已经 {-days} 天" class CountdownInput(BaseModel): title: str = Field(min_length=1, max_length=200) event_date: date calendar_mode: str = Field("solar", pattern="^(solar|lunar)$") lunar_month: int | None = None lunar_day: int | None = None ignore_year: bool = False kind: str = Field("countdown", pattern="^(countdown|anniversary|birthday)$") repeat_rule: str = Field("none", pattern="^(none|weekly|monthly|yearly)$") icon: str = Field("📅", min_length=1, max_length=32) pinned: bool = False @model_validator(mode="after") def calendar_fields_valid(self): if self.calendar_mode == "solar": if self.lunar_month is not None or self.lunar_day is not None: raise ValueError("solar countdown cannot include lunar fields") return self if self.lunar_month is None or self.lunar_day is None: raise ValueError("lunar_month and lunar_day are required") if self.lunar_month == 0 or not -12 <= self.lunar_month <= 12: raise ValueError("lunar_month must be 1..12 or -1..-12 for leap months") if not 1 <= self.lunar_day <= 30: raise ValueError("lunar_day must be 1..30") converted = lunar_to_solar_safe(self.event_date.year, self.lunar_month, self.lunar_day) if converted is None: raise ValueError("lunar date does not exist in the selected year") self.event_date = converted return self class CountdownUpdate(BaseModel): title: str | None = Field(None, min_length=1, max_length=200) event_date: date | None = None calendar_mode: str | None = Field(None, pattern="^(solar|lunar)$") lunar_month: int | None = None lunar_day: int | None = None ignore_year: bool | None = None kind: str | None = Field(None, pattern="^(countdown|anniversary|birthday)$") repeat_rule: str | None = Field(None, pattern="^(none|weekly|monthly|yearly)$") icon: str | None = Field(None, min_length=1, max_length=32) @model_validator(mode="after") def reject_explicit_nulls(self): for field in self.model_fields_set: if getattr(self, field) is None: raise ValueError(f"{field} cannot be null") if self.lunar_month is not None and (self.lunar_month == 0 or not -12 <= self.lunar_month <= 12): raise ValueError("lunar_month must be 1..12 or -1..-12 for leap months") if self.lunar_day is not None and not 1 <= self.lunar_day <= 30: raise ValueError("lunar_day must be 1..30") return self def countdown_dict(row: Countdown, today: date | None = None): today = today or datetime.now(ZoneInfo("Asia/Shanghai")).date() if row.calendar_mode == "lunar" and (row.ignore_year or row.repeat_rule != "none"): # 农历按年重复:忽略年份或指定重复时,都按“每年农历”语义计算下一次 display_date = next_lunar_occurrence( row.lunar_month, row.lunar_day, True, "yearly", today ) or row.event_date effective_repeat = "yearly" else: repeat_rule = "yearly" if row.ignore_year else row.repeat_rule display_date = countdown_occurrence(row.event_date, repeat_rule, today) effective_repeat = repeat_rule days, day_text = countdown_status(display_date, today) lunar_year = None lunar_text = None if row.calendar_mode == "lunar": lunar_year, _, _ = solar_to_lunar_parts(row.event_date) lunar_text = ( solar_to_lunar_text(display_date) if row.ignore_year or row.repeat_rule != "none" else lunar_label_with_year(row.event_date) ) return { "id": row.id, "title": row.title, "event_date": row.event_date, "display_date": display_date, "calendar_mode": row.calendar_mode, "lunar_year": lunar_year, "lunar_month": row.lunar_month, "lunar_day": row.lunar_day, "ignore_year": row.ignore_year, "repeat_rule": effective_repeat, "lunar_text": lunar_text, "kind": row.kind, "icon": row.icon, "pinned": row.pinned, "archived_at": row.archived_at, "days": days, "day_text": day_text, "created_at": row.created_at, "updated_at": row.updated_at, } async def owned_countdown(db, user_id, countdown_id): row = await db.scalar(select(Countdown).where(Countdown.id == countdown_id, Countdown.user_id == user_id)) if not row: raise HTTPException(404, "倒数日不存在") return row @router.post("/countdowns", status_code=201) async def create_countdown(payload: CountdownInput, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): if payload.pinned: await db.execute(update(Countdown).where(Countdown.user_id == user.id).values(pinned=False)) row = Countdown(user_id=user.id, **payload.model_dump()) db.add(row) await db.flush() audit(db, user.id, "create", "countdown", row.id) await db.commit() await db.refresh(row) return countdown_dict(row) @router.get("/countdowns") async def list_countdowns(archived: bool = False, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): condition = Countdown.archived_at.is_not(None) if archived else Countdown.archived_at.is_(None) rows = (await db.scalars( select(Countdown) .where(Countdown.user_id == user.id, condition) .order_by(Countdown.pinned.desc(), Countdown.event_date, Countdown.created_at) )).all() result = [countdown_dict(row) for row in rows] return sorted(result, key=lambda item: (not item["pinned"], item["display_date"], item["created_at"])) @router.patch("/countdowns/{countdown_id}") async def edit_countdown(countdown_id: UUID, payload: CountdownUpdate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): row = await owned_countdown(db, user.id, countdown_id) values = payload.model_dump(exclude_unset=True) combined = { "event_date": values.get("event_date", row.event_date), "calendar_mode": values.get("calendar_mode", row.calendar_mode), "lunar_month": values.get("lunar_month", row.lunar_month), "lunar_day": values.get("lunar_day", row.lunar_day), "ignore_year": values.get("ignore_year", row.ignore_year), } if combined["calendar_mode"] == "solar": if combined["lunar_month"] is not None or combined["lunar_day"] is not None: if "calendar_mode" not in values: raise HTTPException(422, "公历倒数日不能设置农历日期") combined["lunar_month"] = combined["lunar_day"] = None else: if combined["lunar_month"] is None or combined["lunar_day"] is None: raise HTTPException(422, "农历倒数日需要月份和日期") converted = lunar_to_solar_safe( combined["event_date"].year, combined["lunar_month"], combined["lunar_day"] ) if converted is None: raise HTTPException(422, "所选年份不存在该农历日期") combined["event_date"] = converted values.update(combined) for key, value in values.items(): setattr(row, key, value) row.updated_at = utcnow() audit(db, user.id, "update", "countdown", row.id) await db.commit() await db.refresh(row) return countdown_dict(row) @router.post("/countdowns/{countdown_id}/pin") async def pin_countdown(countdown_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): row = await owned_countdown(db, user.id, countdown_id) if row.archived_at is not None: raise HTTPException(409, "已归档倒数日不能置顶") await db.execute(update(Countdown).where(Countdown.user_id == user.id, Countdown.id != row.id).values(pinned=False)) row.pinned = True row.updated_at = utcnow() await db.commit() await db.refresh(row) return countdown_dict(row) @router.delete("/countdowns/{countdown_id}", status_code=204) async def archive_countdown(countdown_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): row = await owned_countdown(db, user.id, countdown_id) if row.archived_at is None: row.archived_at = utcnow() row.pinned = False audit(db, user.id, "archive", "countdown", row.id) await db.commit() return Response(status_code=204) @router.post("/countdowns/{countdown_id}/restore") async def restore_countdown(countdown_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): row = await owned_countdown(db, user.id, countdown_id) if row.archived_at is None: raise HTTPException(409, "倒数日未归档") row.archived_at = None row.updated_at = utcnow() audit(db, user.id, "restore", "countdown", row.id) await db.commit() await db.refresh(row) return countdown_dict(row) @router.delete("/countdowns/{countdown_id}/purge", status_code=204) async def purge_countdown(countdown_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): row = await owned_countdown(db, user.id, countdown_id) if row.archived_at is None: raise HTTPException(409, "请先归档再永久删除") await db.delete(row) await db.commit() return Response(status_code=204) class HabitCreate(BaseModel): name: str = Field(min_length=1, max_length=200) kind: str = Field("boolean", pattern="^(boolean|numeric)$") target: float = Field(1, gt=0) max_value: float | None = Field(None, gt=0) schedule_type: str = Field("daily", pattern="^(daily|weekly|monthly|interval)$") weekdays: list[int] | None = None month_days: list[int] | None = None interval_days: int | None = Field(None, ge=1) start_date: date = Field(default_factory=date.today) @model_validator(mode="after") def schedule_valid(self): if self.schedule_type == "interval" and not self.interval_days: raise ValueError("interval_days required") if self.kind == "boolean": self.target = 1; self.max_value = 1 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) class HabitLogEdit(BaseModel): value: float = Field(ge=0) class PauseInput(BaseModel): start_date: date end_date: date @model_validator(mode="after") def ordered(self): if self.end_date < self.start_date: raise ValueError("invalid range") return self 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, "position": h.position} async def owned_habit(db, user_id, habit_id): row = await db.scalar(select(Habit).where(Habit.id == habit_id, Habit.user_id == user_id)) if not row: raise HTTPException(404, "习惯不存在") return row @router.post("/habits", status_code=201) async def create_habit(payload: HabitCreate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): 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.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}") async def edit_habit(habit_id: UUID, payload: HabitCreate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): row = await owned_habit(db, user.id, habit_id) for key, value in payload.model_dump(exclude={"weekdays", "month_days"}).items(): setattr(row, key, value) row.weekdays = ",".join(map(str, payload.weekdays)) if payload.weekdays else None; row.month_days = ",".join(map(str, payload.month_days)) if payload.month_days else None await db.commit(); return habit_dict(row) @router.delete("/habits/{habit_id}", status_code=204) async def archive_habit(habit_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): row = await owned_habit(db, user.id, habit_id); row.archived_at = utcnow(); await db.commit(); return Response(status_code=204) @router.post("/habits/{habit_id}/logs") async def add_habit_log(habit_id: UUID, payload: HabitLogInput, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): habit = await owned_habit(db, user.id, habit_id) row = await db.scalar(select(HabitLog).where(HabitLog.habit_id == habit.id, HabitLog.day == payload.day)) value = min((row.value if row else 0) + payload.value, habit.max_value or float("inf")) if habit.kind == "boolean": value = 1 if row: row.value = value; row.updated_at = utcnow() else: row = HabitLog(habit_id=habit.id, day=payload.day, value=value); db.add(row) await db.commit(); return {"day": row.day, "value": row.value} @router.put("/habits/{habit_id}/logs/{day}") async def edit_habit_log(habit_id: UUID, day: date, payload: HabitLogEdit, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): habit = await owned_habit(db, user.id, habit_id) value = min(payload.value, habit.max_value or float("inf")); value = float(bool(value)) if habit.kind == "boolean" else value row = await db.scalar(select(HabitLog).where(HabitLog.habit_id == habit.id, HabitLog.day == day)) if row: row.value = value; row.updated_at = utcnow() else: row = HabitLog(habit_id=habit.id, day=day, value=value); db.add(row) await db.commit(); return {"day": row.day, "value": row.value} @router.get("/habits/{habit_id}/logs") async def habit_logs(habit_id: UUID, from_date: date | None = Query(default=None, alias="from"), to_date: date | None = Query(default=None, alias="to"), user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): habit = await owned_habit(db, user.id, habit_id) condition = HabitLog.habit_id == habit.id if from_date is not None: condition = condition & (HabitLog.day >= from_date) if to_date is not None: condition = condition & (HabitLog.day <= to_date) return [{"day": x.day, "value": x.value} for x in (await db.scalars(select(HabitLog).where(condition).order_by(HabitLog.day.desc()))).all()] @router.post("/habits/{habit_id}/pauses", status_code=201) async def pause_habit(habit_id: UUID, payload: PauseInput, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): habit = await owned_habit(db, user.id, habit_id); row = HabitPause(habit_id=habit.id, **payload.model_dump()); db.add(row); await db.commit(); await db.refresh(row); return {"id": row.id, **payload.model_dump()} def scheduled(h, day): if day < h.start_date: return False if h.schedule_type == "daily": return True if h.schedule_type == "weekly": return day.weekday() in {int(x) for x in (h.weekdays or "").split(",") if x} if h.schedule_type == "monthly": return day.day in {int(x) for x in (h.month_days or "").split(",") if x} return (day - h.start_date).days % h.interval_days == 0 @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.position, Habit.created_at))).all()) if not habits: return {"days": days, "habits": []} habit_ids = [habit.id for habit in habits] week_logs = (await db.scalars(select(HabitLog).where(HabitLog.habit_id.in_(habit_ids), HabitLog.day.between(days[0], days[-1])))).all() stats_rows = (await db.execute( select( HabitLog.habit_id, func.sum(HabitLog.value), func.count(HabitLog.id), func.sum(case((HabitLog.value >= Habit.target, 1), else_=0)), ) .join(Habit, Habit.id == HabitLog.habit_id) .where(HabitLog.habit_id.in_(habit_ids)) .group_by(HabitLog.habit_id) )).all() pause_rows = (await db.scalars(select(HabitPause).where(HabitPause.habit_id.in_(habit_ids), HabitPause.end_date >= days[0], HabitPause.start_date <= days[-1]))).all() logs_by_habit = {} stats_by_habit = {} pauses_by_habit = {} for log in week_logs: logs_by_habit.setdefault(log.habit_id, {})[log.day] = log.value for habit_id, total, logged_days, completed_days in stats_rows: stats_by_habit[habit_id] = {"total": total or 0, "completed_days": completed_days or 0, "logged_days": logged_days or 0} for pause in pause_rows: pauses_by_habit.setdefault(pause.habit_id, []).append(pause) output = [] for habit in habits: logs = logs_by_habit.get(habit.id, {}) pauses = pauses_by_habit.get(habit.id, []) data = habit_dict(habit) data["cells"] = [{"day": day, "scheduled": scheduled(habit, day), "paused": any(pause.start_date <= day <= pause.end_date for pause in pauses), "value": logs.get(day, 0)} for day in days] data["stats"] = stats_by_habit.get(habit.id, {"total": 0, "completed_days": 0, "logged_days": 0}) output.append(data) return {"days": days, "habits": output} @router.get("/habits/{habit_id}/stats") async def habit_stats(habit_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): habit = await owned_habit(db, user.id, habit_id); logs = list((await db.scalars(select(HabitLog).where(HabitLog.habit_id == habit.id).order_by(HabitLog.day))).all()) return {"total": sum(x.value for x in logs), "completed_days": sum(x.value >= habit.target for x in logs), "logged_days": len(logs)} _ALLOWED_MIME = {"text/plain", "text/csv", "application/pdf", "image/jpeg", "image/png", "image/gif", "application/json", "application/zip"} @router.post("/tasks/{task_id}/attachments", status_code=201) async def upload_attachment(task_id: UUID, file: UploadFile = File(...), user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): await owned_task(db, user.id, task_id) name = Path(file.filename or "").name if not name or name != file.filename or file.content_type not in _ALLOWED_MIME: raise HTTPException(400, "文件名或类型不允许") limit = get_settings().attachment_max_mb * 1024 * 1024; content = await file.read(limit + 1) if len(content) > limit: raise HTTPException(413, "文件过大") root = Path(get_settings().attachment_dir).resolve(); root.mkdir(parents=True, exist_ok=True); storage = str(new_id()) (root / storage).write_bytes(content) row = Attachment(user_id=user.id, task_id=task_id, filename=name, storage_name=storage, mime_type=file.content_type, size=len(content)); db.add(row); await db.commit(); await db.refresh(row) return {"id": row.id, "task_id": row.task_id, "filename": row.filename, "mime_type": row.mime_type, "size": row.size} async def owned_attachment(db, user_id, attachment_id): row = await db.scalar(select(Attachment).where(Attachment.id == attachment_id, Attachment.user_id == user_id)) if not row: raise HTTPException(404, "附件不存在") return row @router.get("/attachments/{attachment_id}") async def download_attachment(attachment_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): row = await owned_attachment(db, user.id, attachment_id); path = Path(get_settings().attachment_dir).resolve() / row.storage_name if not path.is_file(): raise HTTPException(404, "附件文件不存在") return FileResponse(path, media_type=row.mime_type, filename=row.filename) @router.delete("/attachments/{attachment_id}", status_code=204) async def delete_attachment(attachment_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): row = await owned_attachment(db, user.id, attachment_id); (Path(get_settings().attachment_dir).resolve() / row.storage_name).unlink(missing_ok=True); await db.delete(row); await db.commit(); return Response(status_code=204) def read_ticktick(content: bytes): try: text = content.decode("utf-8-sig") except UnicodeDecodeError as exc: raise HTTPException(422, "CSV 必须为 UTF-8") from exc reader = csv.DictReader(io.StringIO(text)); required = {"Title", "ID"} if not reader.fieldnames or not required <= set(reader.fieldnames): raise HTTPException(422, "CSV 缺少 Title 或 ID") rows = []; errors = [] for index, row in enumerate(reader, 2): if not row.get("Title", "").strip() or not row.get("ID", "").strip(): errors.append({"row": index, "error": "Title/ID required"}) else: rows.append(row) return rows, errors @router.post("/import/ticktick/preview") async def preview_ticktick(file: UploadFile = File(...), user: User = Depends(current_user)): rows, errors = read_ticktick(await file.read()); return {"valid": len(rows), "invalid": len(errors), "errors": errors, "sample": rows[:10]} @router.post("/import/ticktick") async def import_ticktick(file: UploadFile = File(...), user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): rows, errors = read_ticktick(await file.read()) if errors: raise HTTPException(422, errors) inbox = await db.scalar(select(TaskList).where(TaskList.user_id == user.id, TaskList.is_inbox.is_(True))) external_ids = [row["ID"].strip() for row in rows if row.get("ID", "").strip()] existing_ids = set( ( await db.scalars( select(Task.external_id).where(Task.user_id == user.id, Task.external_id.in_(external_ids)) ) ).all() ) if external_ids else set() imported = skipped = 0 for raw in rows: external_id = raw["ID"].strip() if not external_id or external_id in existing_ids: skipped += 1 continue due = None if raw.get("Due Date"): try: due = datetime.combine(date.fromisoformat(raw["Due Date"][:10]), time.min, tzinfo=UTC) except ValueError: raise HTTPException(422, f"无效日期: {raw['Due Date']}") task = Task(user_id=user.id, list_id=inbox.id, title=raw["Title"].strip(), completed=raw.get("Status", "0").lower() in {"1", "completed", "true"}, due_at=due, external_id=external_id); db.add(task); existing_ids.add(external_id); imported += 1 audit(db, user.id, "import", "task", count=imported); await db.commit(); return {"imported": imported, "skipped": skipped} @router.get("/export") async def export_json(user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): def serialize(row, fields): return {f: (str(v) if isinstance((v := getattr(row, f)), UUID) else v.isoformat() if isinstance(v, (date, datetime)) else v) for f in fields} folders = list((await db.scalars(select(Folder).where(Folder.user_id == user.id))).all()); lists = list((await db.scalars(select(TaskList).where(TaskList.user_id == user.id))).all()); tasks = list((await db.scalars(select(Task).where(Task.user_id == user.id))).all()); habits = list((await db.scalars(select(Habit).where(Habit.user_id == user.id))).all()); countdowns = list((await db.scalars(select(Countdown).where(Countdown.user_id == user.id))).all()) recurrences = list((await db.scalars(select(RecurrenceTemplate).where(RecurrenceTemplate.user_id == user.id))).all()) return { "version": 1, "exported_at": utcnow(), "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], "recurrences": [serialize(x, ["id", "task_id", "rrule", "starts_at", "ends_at"]) for x in recurrences], "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], } @router.post("/restore") async def restore_json(payload: dict, mode: str = Query("merge", pattern="^(merge|replace)$"), user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): if payload.get("version") != 1: raise HTTPException(422, "不支持的备份版本") if mode == "replace": await db.execute(delete(Countdown).where(Countdown.user_id == user.id)) await db.execute(delete(Task).where(Task.user_id == user.id)) await db.execute(delete(Habit).where(Habit.user_id == user.id)) await db.execute(delete(TaskList).where(TaskList.user_id == user.id)) await db.execute(delete(Folder).where(Folder.user_id == user.id)) id_map = {} task_id_map = {} for raw in payload.get("folders", []): old = raw["id"] row = Folder(user_id=user.id, name=raw["name"], position=raw.get("position", 0)) db.add(row) await db.flush() id_map[old] = row.id inbox = None for raw in payload.get("lists", []): row = TaskList( user_id=user.id, folder_id=id_map.get(raw.get("folder_id")), name=raw["name"], is_inbox=raw.get("is_inbox", False), position=raw.get("position", 0), ) db.add(row) await db.flush() id_map[raw["id"]] = row.id if row.is_inbox: inbox = row if not inbox: inbox = TaskList(user_id=user.id, name="收集箱", is_inbox=True) db.add(inbox) await db.flush() restored = 0 pending_tasks = [] for raw in payload.get("tasks", []): ext = raw.get("external_id") existing = await db.scalar(select(Task).where(Task.user_id == user.id, Task.external_id == ext)) if ext else None if existing and mode == "merge": task_id_map[raw["id"]] = existing.id continue row = Task( user_id=user.id, list_id=id_map.get(raw.get("list_id"), inbox.id), title=raw["title"], description=raw.get("description", ""), priority=raw.get("priority", 0), completed=raw.get("completed", False), due_at=datetime.fromisoformat(raw["due_at"]) if raw.get("due_at") else None, external_id=ext, ) db.add(row) await db.flush() task_id_map[raw["id"]] = row.id pending_tasks.append((row.id, raw.get("parent_id"))) restored += 1 for task_id, old_parent_id in pending_tasks: if old_parent_id and old_parent_id in task_id_map: await db.execute(update(Task).where(Task.id == task_id, Task.user_id == user.id).values(parent_id=task_id_map[old_parent_id])) for raw in payload.get("recurrences", []): task_id = task_id_map.get(raw.get("task_id")) if not task_id: continue db.add(RecurrenceTemplate( user_id=user.id, task_id=task_id, rrule=raw["rrule"], starts_at=datetime.fromisoformat(raw["starts_at"]), ends_at=datetime.fromisoformat(raw["ends_at"]) if raw.get("ends_at") else None, )) for raw in payload.get("habits", []): row = Habit( user_id=user.id, name=raw["name"], kind=raw.get("kind", "boolean"), target=raw.get("target", 1), max_value=raw.get("max_value"), schedule_type=raw.get("schedule_type", "daily"), weekdays=raw.get("weekdays"), month_days=raw.get("month_days"), 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()) has_pinned_countdown = bool(await db.scalar(select(Countdown.id).where(Countdown.user_id == user.id, Countdown.pinned.is_(True)))) for raw in payload.get("countdowns", []): try: source_id = UUID(raw["id"]) except (KeyError, TypeError, ValueError) as exc: raise HTTPException(422, "无效的倒数日备份 ID") from exc if mode == "merge" and source_id in existing_countdown_ids: continue try: item = CountdownInput( title=raw["title"], event_date=date.fromisoformat(raw["event_date"]), calendar_mode=raw.get("calendar_mode", "solar"), lunar_month=raw.get("lunar_month"), lunar_day=raw.get("lunar_day"), ignore_year=bool(raw.get("ignore_year", False)), kind=raw.get("kind", "countdown"), repeat_rule=raw.get("repeat_rule", "none"), icon=raw.get("icon", "📅"), pinned=bool(raw.get("pinned", False)) and not has_pinned_countdown, ) # Backups store the canonical solar anchor; validation above converts # lunar input again, so preserve the exact exported anchor on restore. item.event_date = date.fromisoformat(raw["event_date"]) except (KeyError, TypeError, ValueError) as exc: raise HTTPException(422, "无效的倒数日备份数据") from exc row = Countdown( id=source_id, user_id=user.id, **item.model_dump(), archived_at=datetime.fromisoformat(raw["archived_at"]) if raw.get("archived_at") else None, created_at=datetime.fromisoformat(raw["created_at"]) if raw.get("created_at") else utcnow(), updated_at=datetime.fromisoformat(raw["updated_at"]) if raw.get("updated_at") else utcnow(), ) if row.archived_at is not None: row.pinned = False has_pinned_countdown = has_pinned_countdown or row.pinned db.add(row) existing_countdown_ids.add(source_id) audit(db, user.id, "restore", "backup", count=restored, mode=mode) await db.commit() return {"restored": restored, "mode": mode} @router.get("/audit-logs") async def audit_logs(limit: int = Query(100, ge=1, le=500), user: User = Depends(current_user), db: AsyncSession = Depends(get_db)): rows = (await db.scalars(select(AuditLog).where(AuditLog.user_id == user.id).order_by(AuditLog.created_at.desc()).limit(limit))).all() return [{"id": x.id, "action": x.action, "entity_type": x.entity_type, "entity_id": x.entity_id, "details": x.details, "created_at": x.created_at} for x in rows]