This commit is contained in:
+1
-300
@@ -1,21 +1,12 @@
|
||||
import asyncio
|
||||
import csv
|
||||
import http.client
|
||||
import io
|
||||
import ipaddress
|
||||
import re
|
||||
import socket
|
||||
import ssl
|
||||
import urllib.parse
|
||||
from datetime import UTC, date, datetime, time, timedelta
|
||||
from pathlib import Path
|
||||
from uuid import UUID
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from dateutil.rrule import rrulestr
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, Response, UploadFile
|
||||
from fastapi.responses import FileResponse
|
||||
from icalendar import Calendar
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from sqlalchemy import case, delete, func, select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@@ -26,17 +17,14 @@ from .db import get_db
|
||||
from .models import (
|
||||
Attachment,
|
||||
AuditLog,
|
||||
CalendarSubscription,
|
||||
Folder,
|
||||
Habit,
|
||||
HabitLog,
|
||||
HabitPause,
|
||||
RecurrenceException,
|
||||
RecurrenceTemplate,
|
||||
Tag,
|
||||
Task,
|
||||
TaskList,
|
||||
TaskTag,
|
||||
User,
|
||||
new_id,
|
||||
utcnow,
|
||||
@@ -262,265 +250,6 @@ async def delete_recurrence(recurrence_id: UUID, scope: str = Query("all", patte
|
||||
await db.commit(); return Response(status_code=204)
|
||||
|
||||
|
||||
class CalendarSubscriptionCreate(BaseModel):
|
||||
name: str = Field(min_length=1, max_length=120)
|
||||
url: str = Field(min_length=1, max_length=2000)
|
||||
color: str = Field(default="#f15a29", min_length=1, max_length=32)
|
||||
|
||||
|
||||
def calendar_subscription_dict(row: CalendarSubscription):
|
||||
return {"id": row.id, "name": row.name, "url": row.url, "color": row.color, "enabled": row.enabled}
|
||||
|
||||
|
||||
LOCAL_TZ = ZoneInfo("Asia/Shanghai")
|
||||
|
||||
|
||||
def _is_forbidden_ip(value: str) -> bool:
|
||||
ip = ipaddress.ip_address(value)
|
||||
return ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_multicast or ip.is_reserved
|
||||
|
||||
|
||||
def _localized_ics_datetime(value, default_tz=LOCAL_TZ):
|
||||
if isinstance(value, datetime):
|
||||
return value if value.tzinfo else value.replace(tzinfo=default_tz), False
|
||||
if isinstance(value, date):
|
||||
return datetime.combine(value, time.min, tzinfo=default_tz), True
|
||||
raise ValueError("unsupported ICS datetime")
|
||||
|
||||
|
||||
def _as_utc(value, default_tz=LOCAL_TZ):
|
||||
localized, all_day = _localized_ics_datetime(value, default_tz)
|
||||
return localized.astimezone(UTC), all_day
|
||||
|
||||
|
||||
def _event_duration(event, starts_at: datetime, all_day: bool):
|
||||
if event.get("dtend"):
|
||||
ends_at, _ = _as_utc(event.decoded("dtend"))
|
||||
return max(ends_at - starts_at, timedelta())
|
||||
if event.get("duration"):
|
||||
return event.decoded("duration")
|
||||
return timedelta(days=1) if all_day else timedelta(hours=1)
|
||||
|
||||
|
||||
def _excluded_starts(event):
|
||||
excluded = set()
|
||||
exdates = event.get("exdate")
|
||||
if not exdates:
|
||||
return excluded
|
||||
if not isinstance(exdates, list):
|
||||
exdates = [exdates]
|
||||
for exdate in exdates:
|
||||
for item in getattr(exdate, "dts", []):
|
||||
excluded.add(_as_utc(item.dt)[0])
|
||||
return excluded
|
||||
|
||||
|
||||
def _append_calendar_event(events, source_name, color, event, starts_at, ends_at, all_day):
|
||||
uid = str(event.get("uid") or "")
|
||||
title = str(event.get("summary") or "未命名事件").strip() or "未命名事件"
|
||||
events.append({
|
||||
"id": uid or f"{source_name}-{starts_at.isoformat()}-{title}",
|
||||
"title": title,
|
||||
"starts_at": starts_at,
|
||||
"ends_at": ends_at,
|
||||
"all_day": all_day,
|
||||
"source_name": source_name,
|
||||
"color": color,
|
||||
})
|
||||
|
||||
|
||||
def _overlaps(starts_at: datetime, ends_at: datetime | None, window_start: datetime | None, window_end: datetime | None) -> bool:
|
||||
if not window_start or not window_end:
|
||||
return True
|
||||
return starts_at < window_end and (ends_at or starts_at) > window_start
|
||||
|
||||
|
||||
def parse_ics_events(content: str, source_name: str, color: str, window_start: datetime | None = None, window_end: datetime | None = None):
|
||||
calendar = Calendar.from_ical(content)
|
||||
components = list(calendar.walk("VEVENT"))
|
||||
overrides = {}
|
||||
for event in components:
|
||||
recurrence_id = event.get("recurrence-id")
|
||||
if recurrence_id:
|
||||
overrides[(str(event.get("uid") or ""), _as_utc(event.decoded("recurrence-id"))[0])] = event
|
||||
|
||||
events = []
|
||||
for event in components:
|
||||
if not event.get("dtstart") or event.get("recurrence-id") or str(event.get("status") or "").upper() == "CANCELLED":
|
||||
continue
|
||||
localized_start, all_day = _localized_ics_datetime(event.decoded("dtstart"))
|
||||
starts_at = localized_start.astimezone(UTC)
|
||||
duration = _event_duration(event, starts_at, all_day)
|
||||
excluded = _excluded_starts(event)
|
||||
uid = str(event.get("uid") or "")
|
||||
if event.get("rrule") and window_start and window_end:
|
||||
rule_text = event.get("rrule").to_ical().decode()
|
||||
rule = rrulestr(rule_text, dtstart=localized_start)
|
||||
for occurrence in rule.between(window_start - duration, window_end, inc=True):
|
||||
occurrence = occurrence.astimezone(UTC) if occurrence.tzinfo else occurrence.replace(tzinfo=LOCAL_TZ).astimezone(UTC)
|
||||
if occurrence in excluded or (uid, occurrence) in overrides:
|
||||
continue
|
||||
ends_at = occurrence + duration
|
||||
if _overlaps(occurrence, ends_at, window_start, window_end):
|
||||
_append_calendar_event(events, source_name, color, event, occurrence, ends_at, all_day)
|
||||
continue
|
||||
ends_at = starts_at + duration
|
||||
if _overlaps(starts_at, ends_at, window_start, window_end):
|
||||
_append_calendar_event(events, source_name, color, event, starts_at, ends_at, all_day)
|
||||
|
||||
for event in overrides.values():
|
||||
if not event.get("dtstart") or str(event.get("status") or "").upper() == "CANCELLED":
|
||||
continue
|
||||
starts_at, all_day = _as_utc(event.decoded("dtstart"))
|
||||
duration = _event_duration(event, starts_at, all_day)
|
||||
ends_at = starts_at + duration
|
||||
if _overlaps(starts_at, ends_at, window_start, window_end):
|
||||
_append_calendar_event(events, source_name, color, event, starts_at, ends_at, all_day)
|
||||
|
||||
return events
|
||||
|
||||
|
||||
def _validated_calendar_target(url: str):
|
||||
parsed = urllib.parse.urlparse(url)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
|
||||
raise HTTPException(422, "日历订阅只支持 http/https 链接")
|
||||
if parsed.username or parsed.password:
|
||||
raise HTTPException(422, "日历订阅链接不能包含账号密码")
|
||||
port = parsed.port or (443 if parsed.scheme == "https" else 80)
|
||||
try:
|
||||
addresses = socket.getaddrinfo(parsed.hostname, port, type=socket.SOCK_STREAM)
|
||||
except socket.gaierror as exc:
|
||||
raise HTTPException(422, "日历订阅域名无法解析") from exc
|
||||
for item in addresses:
|
||||
ip = item[4][0]
|
||||
if not _is_forbidden_ip(ip):
|
||||
return parsed, ip, port
|
||||
raise HTTPException(422, "日历订阅不能指向内网地址")
|
||||
|
||||
|
||||
def _validate_public_calendar_url(url: str):
|
||||
_validated_calendar_target(url)
|
||||
|
||||
|
||||
def _request_pinned_calendar_url(url: str):
|
||||
parsed, ip, port = _validated_calendar_target(url)
|
||||
target_host = f"[{ip}]" if ":" in ip else ip
|
||||
path = urllib.parse.urlunparse(("", "", parsed.path or "/", parsed.params, parsed.query, ""))
|
||||
headers = {"Host": parsed.hostname or "", "User-Agent": "dodo-calendar-fetch/1.0"}
|
||||
if parsed.port:
|
||||
headers["Host"] = f"{headers['Host']}:{parsed.port}"
|
||||
if parsed.scheme == "https":
|
||||
connection = http.client.HTTPSConnection(target_host, port=port, timeout=15, context=ssl.create_default_context())
|
||||
else:
|
||||
connection = http.client.HTTPConnection(target_host, port=port, timeout=15)
|
||||
try:
|
||||
if parsed.scheme == "https":
|
||||
raw = socket.create_connection((ip, port), timeout=15)
|
||||
sock = ssl.create_default_context().wrap_socket(raw, server_hostname=parsed.hostname)
|
||||
connection.sock = sock
|
||||
connection.request("GET", path, headers=headers)
|
||||
response = connection.getresponse()
|
||||
body = response.read(2_000_001)
|
||||
return response.status, response.getheaders(), response.getheader("Content-Type") or "", body
|
||||
except OSError as exc:
|
||||
raise HTTPException(502, f"订阅拉取失败: {exc}") from exc
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
|
||||
_CAL_CACHE: dict[str, tuple[float, list[dict]]] = {}
|
||||
_CAL_CACHE_TTL = 300.0
|
||||
|
||||
|
||||
def fetch_calendar_events(url: str, source_name: str, color: str, window_start: datetime | None = None, window_end: datetime | None = None):
|
||||
cache_key = f"{source_name}|{url}|{window_start.isoformat() if window_start else ''}|{window_end.isoformat() if window_end else ''}"
|
||||
now = datetime.now(UTC).timestamp()
|
||||
hit = _CAL_CACHE.get(cache_key)
|
||||
if hit and now - hit[0] < _CAL_CACHE_TTL:
|
||||
return [dict(item) for item in hit[1]]
|
||||
current_url = url
|
||||
for _ in range(4):
|
||||
status, headers, content_type, body = _request_pinned_calendar_url(current_url)
|
||||
if status in {301, 302, 303, 307, 308}:
|
||||
location = dict(headers).get("Location")
|
||||
if not location:
|
||||
raise HTTPException(502, "订阅重定向缺少目标地址")
|
||||
current_url = urllib.parse.urljoin(current_url, location)
|
||||
continue
|
||||
if status >= 400:
|
||||
raise HTTPException(502, f"订阅拉取失败: HTTP {status}")
|
||||
if content_type.split(";", 1)[0].lower() not in {"text/calendar", "text/plain", "application/octet-stream"}:
|
||||
raise HTTPException(422, "订阅链接没有返回 ICS 日历内容")
|
||||
if len(body) > 2_000_000:
|
||||
raise HTTPException(413, "日历订阅内容超过 2MB")
|
||||
parsed = parse_ics_events(body.decode("utf-8", errors="ignore"), source_name, color, window_start, window_end)
|
||||
if len(_CAL_CACHE) >= 64:
|
||||
_CAL_CACHE.clear()
|
||||
_CAL_CACHE[cache_key] = (now, parsed)
|
||||
return [dict(item) for item in parsed]
|
||||
raise HTTPException(502, "日历订阅重定向次数过多")
|
||||
|
||||
|
||||
@router.get("/calendar-subscriptions")
|
||||
async def list_calendar_subscriptions(user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
rows = (await db.scalars(select(CalendarSubscription).where(CalendarSubscription.user_id == user.id).order_by(CalendarSubscription.created_at))).all()
|
||||
return [calendar_subscription_dict(row) for row in rows]
|
||||
|
||||
|
||||
@router.post("/calendar-subscriptions", status_code=201)
|
||||
async def create_calendar_subscription(payload: CalendarSubscriptionCreate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
_validate_public_calendar_url(payload.url)
|
||||
row = CalendarSubscription(user_id=user.id, **payload.model_dump())
|
||||
db.add(row)
|
||||
await db.commit()
|
||||
await db.refresh(row)
|
||||
return calendar_subscription_dict(row)
|
||||
|
||||
|
||||
@router.get("/calendar-subscriptions/today")
|
||||
async def today_calendar_events(
|
||||
day: date,
|
||||
user: User = Depends(current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
rows = list((await db.scalars(select(CalendarSubscription).where(CalendarSubscription.user_id == user.id, CalendarSubscription.enabled.is_(True)).order_by(CalendarSubscription.created_at))).all())
|
||||
if not rows:
|
||||
return []
|
||||
tomorrow = day + timedelta(days=1)
|
||||
day_start = datetime.combine(day, time.min, tzinfo=LOCAL_TZ).astimezone(UTC)
|
||||
day_end = datetime.combine(tomorrow, time.min, tzinfo=LOCAL_TZ).astimezone(UTC)
|
||||
|
||||
async def pull(row: CalendarSubscription):
|
||||
return await asyncio.to_thread(fetch_calendar_events, row.url, row.name, row.color, day_start, day_end)
|
||||
|
||||
results = await asyncio.gather(*(pull(row) for row in rows), return_exceptions=True)
|
||||
events = []
|
||||
for row, result in zip(rows, results, strict=True):
|
||||
if isinstance(result, Exception):
|
||||
continue
|
||||
for item in result:
|
||||
events.append({
|
||||
"id": item["id"],
|
||||
"title": item["title"],
|
||||
"starts_at": item["starts_at"].isoformat(),
|
||||
"ends_at": item["ends_at"].isoformat() if item.get("ends_at") else None,
|
||||
"all_day": item["all_day"],
|
||||
"source_name": row.name,
|
||||
"color": row.color,
|
||||
})
|
||||
events.sort(key=lambda item: (item["all_day"] is False, item["starts_at"], item["title"]))
|
||||
return events
|
||||
|
||||
|
||||
@router.delete("/calendar-subscriptions/{subscription_id}", status_code=204)
|
||||
async def delete_calendar_subscription(subscription_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
row = await db.scalar(select(CalendarSubscription).where(CalendarSubscription.id == subscription_id, CalendarSubscription.user_id == user.id))
|
||||
if not row:
|
||||
raise HTTPException(404, "日历订阅不存在")
|
||||
await db.delete(row)
|
||||
await db.commit()
|
||||
return Response(status_code=204)
|
||||
|
||||
|
||||
class HabitCreate(BaseModel):
|
||||
name: str = Field(min_length=1, max_length=200)
|
||||
@@ -763,25 +492,13 @@ async def import_ticktick(file: UploadFile = File(...), user: User = Depends(cur
|
||||
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()); tags = list((await db.scalars(select(Tag).where(Tag.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())
|
||||
task_tags = [
|
||||
{"task_id": str(x.task_id), "tag_id": str(x.tag_id)}
|
||||
for x in (
|
||||
await db.scalars(
|
||||
select(TaskTag)
|
||||
.join(Task, Task.id == TaskTag.task_id)
|
||||
.where(Task.user_id == user.id)
|
||||
)
|
||||
).all()
|
||||
]
|
||||
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())
|
||||
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],
|
||||
"tags": [serialize(x, ["id", "name", "color"]) for x in tags],
|
||||
"tasks": [serialize(x, ["id", "list_id", "parent_id", "title", "description", "priority", "completed", "due_at", "external_id", "deleted_at"]) for x in tasks],
|
||||
"task_tags": task_tags,
|
||||
"habits": [serialize(x, ["id", "name", "kind", "target", "max_value", "schedule_type", "weekdays", "month_days", "interval_days", "start_date", "archived_at"]) for x in habits],
|
||||
}
|
||||
|
||||
@@ -791,15 +508,12 @@ async def restore_json(payload: dict, mode: str = Query("merge", pattern="^(merg
|
||||
if payload.get("version") != 1:
|
||||
raise HTTPException(422, "不支持的备份版本")
|
||||
if mode == "replace":
|
||||
await db.execute(delete(TaskTag).where(TaskTag.task_id.in_(select(Task.id).where(Task.user_id == user.id))))
|
||||
await db.execute(delete(Task).where(Task.user_id == user.id))
|
||||
await db.execute(delete(Tag).where(Tag.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 = {}
|
||||
tag_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))
|
||||
@@ -824,11 +538,6 @@ async def restore_json(payload: dict, mode: str = Query("merge", pattern="^(merg
|
||||
inbox = TaskList(user_id=user.id, name="收集箱", is_inbox=True)
|
||||
db.add(inbox)
|
||||
await db.flush()
|
||||
for raw in payload.get("tags", []):
|
||||
row = Tag(user_id=user.id, name=raw["name"], color=raw.get("color", "#f15a29"))
|
||||
db.add(row)
|
||||
await db.flush()
|
||||
tag_id_map[raw["id"]] = row.id
|
||||
restored = 0
|
||||
pending_tasks = []
|
||||
for raw in payload.get("tasks", []):
|
||||
@@ -855,14 +564,6 @@ async def restore_json(payload: dict, mode: str = Query("merge", pattern="^(merg
|
||||
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]))
|
||||
task_tag_rows = []
|
||||
for raw in payload.get("task_tags", []):
|
||||
task_id = task_id_map.get(raw.get("task_id"))
|
||||
tag_id = tag_id_map.get(raw.get("tag_id"))
|
||||
if task_id and tag_id:
|
||||
task_tag_rows.append(TaskTag(task_id=task_id, tag_id=tag_id))
|
||||
if task_tag_rows:
|
||||
db.add_all(task_tag_rows)
|
||||
for raw in payload.get("habits", []):
|
||||
row = Habit(
|
||||
user_id=user.id,
|
||||
|
||||
Reference in New Issue
Block a user