Files
dodo/backend/calendar_router.py
T
bboysoul 97c1bd4e87
ci / gitleaks (push) Successful in 7s
ci / docker (push) Successful in 3m29s
feat: auto-refresh calendar subscriptions
2026-09-21 15:36:09 +08:00

195 lines
6.6 KiB
Python

from datetime import datetime, timedelta
from uuid import UUID
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
from fastapi import APIRouter, Depends, HTTPException, Query, Response
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from . import calendar as calendar_service
from .auth import current_user
from .calendar_refresh import refresh_subscription_cache
from .db import get_db
from .models import CalendarSubscription, User
router = APIRouter(prefix="/api/v1", tags=["calendar"])
MAX_WINDOW = timedelta(days=366)
class SubscriptionCreate(BaseModel):
name: str = Field(min_length=1, max_length=120)
url: str = Field(min_length=1, max_length=2000)
color: str = Field(default="#f15a29", pattern=r"^#[0-9A-Fa-f]{6}$")
enabled: bool = True
@field_validator("name")
@classmethod
def clean_name(cls, value: str) -> str:
value = value.strip()
if not value:
raise ValueError("name cannot be blank")
return value
class SubscriptionUpdate(BaseModel):
name: str | None = Field(default=None, min_length=1, max_length=120)
url: str | None = Field(default=None, min_length=1, max_length=2000)
color: str | None = Field(default=None, pattern=r"^#[0-9A-Fa-f]{6}$")
enabled: bool | None = None
@field_validator("name")
@classmethod
def clean_name(cls, value: str | None) -> str | None:
if value is None:
return None
value = value.strip()
if not value:
raise ValueError("name cannot be blank")
return value
@model_validator(mode="after")
def reject_nulls(self):
for field in self.model_fields_set:
if getattr(self, field) is None:
raise ValueError(f"{field} cannot be null")
return self
class SubscriptionOut(BaseModel):
model_config = ConfigDict(from_attributes=True)
id: UUID
name: str
url: str
color: str
enabled: bool
refreshed_at: datetime | None
last_error: str | None
stale: bool
def _out(row: CalendarSubscription) -> dict:
return {
"id": row.id,
"name": row.name,
"url": row.url,
"color": row.color,
"enabled": row.enabled,
"refreshed_at": row.refreshed_at,
"last_error": row.last_error,
"stale": bool(row.last_error and row.ics_cache),
}
async def _owned(db: AsyncSession, user_id: UUID, subscription_id: UUID) -> CalendarSubscription:
row = await db.scalar(select(CalendarSubscription).where(
CalendarSubscription.id == subscription_id,
CalendarSubscription.user_id == user_id,
))
if row is None:
raise HTTPException(404, "calendar subscription not found")
return row
@router.get("/calendar-subscriptions", response_model=list[SubscriptionOut])
async def list_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, CalendarSubscription.id))).all()
return [_out(row) for row in rows]
@router.post("/calendar-subscriptions", response_model=SubscriptionOut, status_code=201)
async def create_subscription(
payload: SubscriptionCreate,
user: User = Depends(current_user),
db: AsyncSession = Depends(get_db),
):
calendar_service.validate_calendar_url(payload.url)
row = CalendarSubscription(user_id=user.id, **payload.model_dump())
db.add(row)
await db.flush()
await refresh_subscription_cache(db, row)
await db.refresh(row)
return _out(row)
@router.patch("/calendar-subscriptions/{subscription_id}", response_model=SubscriptionOut)
async def update_subscription(
subscription_id: UUID,
payload: SubscriptionUpdate,
user: User = Depends(current_user),
db: AsyncSession = Depends(get_db),
):
row = await _owned(db, user.id, subscription_id)
changes = payload.model_dump(exclude_unset=True)
if "url" in changes:
calendar_service.validate_calendar_url(changes["url"])
if changes["url"] != row.url:
row.ics_cache = row.etag = row.last_modified = row.refreshed_at = row.last_error = None
for key, value in changes.items():
setattr(row, key, value)
await db.commit()
await db.refresh(row)
return _out(row)
@router.delete("/calendar-subscriptions/{subscription_id}", status_code=204)
async def delete_subscription(
subscription_id: UUID,
user: User = Depends(current_user),
db: AsyncSession = Depends(get_db),
):
row = await _owned(db, user.id, subscription_id)
await db.delete(row)
await db.commit()
return Response(status_code=204)
@router.post("/calendar-subscriptions/{subscription_id}/refresh", response_model=SubscriptionOut)
async def refresh_subscription(
subscription_id: UUID,
user: User = Depends(current_user),
db: AsyncSession = Depends(get_db),
):
row = await _owned(db, user.id, subscription_id)
await refresh_subscription_cache(db, row)
await db.refresh(row)
return _out(row)
@router.get("/calendar-events")
async def calendar_events(
start: datetime = Query(),
end: datetime = Query(),
user: User = Depends(current_user),
db: AsyncSession = Depends(get_db),
):
if start.tzinfo is None or end.tzinfo is None or end <= start or end - start > MAX_WINDOW:
raise HTTPException(422, "start/end must be timezone-aware and span at most 366 days")
try:
ZoneInfo(user.timezone)
except ZoneInfoNotFoundError as exc:
raise HTTPException(422, "user timezone is invalid") from exc
rows = (await db.scalars(select(CalendarSubscription).where(
CalendarSubscription.user_id == user.id,
CalendarSubscription.enabled.is_(True),
).order_by(CalendarSubscription.created_at, CalendarSubscription.id))).all()
events = []
sources = []
for row in rows:
if not row.ics_cache:
await refresh_subscription_cache(db, row)
try:
parsed = calendar_service.parse_ics_events(
row.ics_cache or "", row.name, row.color, start, end, user.timezone,
source_id=str(row.id),
)
events.extend(parsed)
except ValueError as exc:
row.last_error = str(exc)
await db.commit()
sources.append({"id": row.id, "name": row.name, "stale": bool(row.last_error)})
events.sort(key=lambda item: (item["starts_at"], item["title"], item["id"]))
return {"events": events, "sources": sources}