175a4ad0d1
Apply sliding-window 2/s limits and BannedIP checks on write routes. Also force SVG attachment disposition and claim burn-after-read in DB before streaming, deleting the file via BackgroundTask after the response. Co-authored-by: Cursor <cursoragent@cursor.com>
199 lines
5.0 KiB
Python
199 lines
5.0 KiB
Python
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
|