from __future__ import annotations import secrets from datetime import datetime, timedelta from pathlib import Path from sqlalchemy import func, select from sqlalchemy.orm import Session from app import config from app.models import Item from app.schemas import ItemCreate, TTLOption TTL_DELTAS: dict[TTLOption, timedelta | None] = { "1h": timedelta(hours=1), "24h": timedelta(days=1), "7d": timedelta(days=7), "never": None, } def _slug() -> str: return secrets.token_urlsafe(8) def _expires_at(ttl: TTLOption, now: datetime | None = None) -> datetime | None: delta = TTL_DELTAS[ttl] if delta is None: return None return (now or datetime.utcnow()) + delta def _title_from_body(body: str) -> str: first = body.split("\n", 1)[0].strip() return first[:255] def _is_unavailable(item: Item, now: datetime | None = None) -> bool: if item.burned: return True now = now or datetime.utcnow() if item.expires_at is not None and item.expires_at <= now: return True return False def item_to_dict(item: Item, *, include_body: bool = True) -> dict: data = { "slug": item.slug, "kind": item.kind, "title": item.title, "file_name": item.file_name, "mime": item.mime, "size_bytes": item.size_bytes, "is_public": item.is_public, "burn_after_read": item.burn_after_read, "expires_at": item.expires_at, "created_at": item.created_at, "view_count": item.view_count, } if include_body: data["body"] = item.body return data def create_text_item(db: Session, payload: ItemCreate, created_ip: str) -> Item: now = datetime.utcnow() title = payload.title if payload.title is not None else _title_from_body(payload.body) item = Item( slug=_slug(), kind="text", title=title[:255], body=payload.body, is_public=payload.is_public, burn_after_read=payload.burn_after_read, expires_at=_expires_at(payload.ttl, now), created_ip=created_ip, created_at=now, ) db.add(item) db.commit() db.refresh(item) return item def create_file_item( db: Session, *, file_name: str, file_path: str, mime: str, size_bytes: int, is_public: bool, burn_after_read: bool, ttl: TTLOption, created_ip: str, title: str | None = None, ) -> Item: now = datetime.utcnow() item = Item( slug=_slug(), kind="file", title=(title or file_name)[:255], body=None, file_name=file_name, file_path=file_path, mime=mime, size_bytes=size_bytes, is_public=is_public, burn_after_read=burn_after_read, expires_at=_expires_at(ttl, now), created_ip=created_ip, created_at=now, ) db.add(item) db.commit() db.refresh(item) return item def resolve_file_path(item: Item) -> Path | None: if not item.file_path: return None path = Path(item.file_path) if not path.is_absolute(): path = config.settings.data_dir / path return path def get_file_item(db: Session, slug: str) -> Item | None: item = db.scalar(select(Item).where(Item.slug == slug)) if item is None or item.kind != "file" or _is_unavailable(item): return None return item def claim_burn(db: Session, item: Item) -> None: """Mark burn consumed in DB before streaming bytes.""" item.burned = True db.commit() def delete_item_file(item: Item) -> None: """Remove file bytes from disk (safe after response has finished streaming).""" path = resolve_file_path(item) if path is not None: try: path.unlink(missing_ok=True) except OSError: pass def burn_file_item(db: Session, item: Item) -> None: """Mark a file item burned and remove its bytes from disk.""" claim_burn(db, item) delete_item_file(item) def get_item(db: Session, slug: str) -> Item | None: item = db.scalar(select(Item).where(Item.slug == slug)) if item is None or _is_unavailable(item): return None item.view_count += 1 # File items burn on download, not metadata fetch. if item.burn_after_read and item.kind != "file": item.burned = True db.commit() db.refresh(item) return item def list_wall( db: Session, *, page: int = 1, page_size: int = 20, ) -> tuple[list[Item], int]: page = max(page, 1) page_size = min(max(page_size, 1), 100) now = datetime.utcnow() filters = [ Item.is_public.is_(True), Item.burned.is_(False), (Item.expires_at.is_(None) | (Item.expires_at > now)), ] total = db.scalar(select(func.count()).select_from(Item).where(*filters)) or 0 items = list( db.scalars( select(Item) .where(*filters) .order_by(Item.created_at.desc()) .offset((page - 1) * page_size) .limit(page_size) ).all() ) return items, total