diff --git a/src/inkwell/app.py b/src/inkwell/app.py index 667b5b2..c3ff79f 100644 --- a/src/inkwell/app.py +++ b/src/inkwell/app.py @@ -25,7 +25,7 @@ from .notes import bp as notes_bp from .proxy import is_https from .responses import json_error, not_found from .retention import run_sweeper -from .settings import MAX_BODY_MB, get_public_config, get_setting, load_or_create_secret_key, refresh_live +from .settings import MAX_BODY_MB, apply_session_ttl, get_public_config, get_setting, load_or_create_secret_key, refresh_live from .settings_api import bp as settings_bp from .shares_api import bp as shares_bp from .sync import bp as sync_bp, protocol_advertisement @@ -114,8 +114,7 @@ def create_app() -> Quart: async with session_scope() as db: app.secret_key = await load_or_create_secret_key(db) try: - days = int(await get_setting(db, "session_ttl_days")) - app.config["PERMANENT_SESSION_LIFETIME"] = timedelta(days=days) + apply_session_ttl(app, await get_setting(db, "session_ttl_days")) except (ValueError, TypeError, KeyError): pass # The security settings the throttle and the proxy trust read on hot diff --git a/src/inkwell/audit.py b/src/inkwell/audit.py index fabd816..aa3f594 100644 --- a/src/inkwell/audit.py +++ b/src/inkwell/audit.py @@ -18,11 +18,11 @@ from __future__ import annotations import logging import uuid -from datetime import datetime, timedelta, timezone +from datetime import datetime, timezone from sqlalchemy import delete, select -from .common import iso +from .common import expired_before, iso from .db import session_scope from .models.audit_event import AuditEvent from .proxy import client_address @@ -99,9 +99,9 @@ async def sweep_once(*, now: datetime | None = None) -> int: """Delete events older than `audit_retention_days`; 0 keeps them forever.""" async with session_scope() as db: days = int(await get_setting(db, "audit_retention_days")) - if days <= 0: + cutoff = expired_before(now or datetime.now(timezone.utc), days) + if cutoff is None: return 0 - cutoff = (now or datetime.now(timezone.utc)) - timedelta(days=days) deleted = (await db.execute(delete(AuditEvent).where(AuditEvent.at < cutoff))).rowcount await db.commit() return deleted diff --git a/src/inkwell/auth.py b/src/inkwell/auth.py index 85d2215..16220de 100644 --- a/src/inkwell/auth.py +++ b/src/inkwell/auth.py @@ -18,7 +18,7 @@ from .password_resets import INVALID as INVALID_RESET, claim as claim_reset, mai from .models.device_token import DeviceToken from .models.user import User from .proxy import client_address -from .responses import json_error, not_found, parse_uuid +from .responses import json_error, not_found, parse_uuid, too_many from .ratelimit import ( register_by_address, reset_mail_by_account, @@ -196,7 +196,7 @@ def _throttled(retry_after: int): genuinely needs. """ logger.warning("throttled credential attempt from=%s retry_after=%ss", client_address(), retry_after) - return (*json_error("too many attempts — try again shortly", 429), {"Retry-After": str(retry_after)}) + return too_many("too many attempts — try again shortly", retry_after) def _sign_in_block(email: str) -> int | None: diff --git a/src/inkwell/client_dist.py b/src/inkwell/client_dist.py index 185d13b..1587b80 100644 --- a/src/inkwell/client_dist.py +++ b/src/inkwell/client_dist.py @@ -82,7 +82,7 @@ from .auth import login_required from .config import Config from .proxy import is_https from .ratelimit import downloads_by_account -from .responses import json_error +from .responses import json_error, too_many @dataclass(frozen=True) @@ -394,10 +394,7 @@ async def client_download(platform_id: str): account = str(g.user_id) wait = downloads_by_account.retry_after(account) if wait is not None: - response = jsonify({"error": "too many downloads from this account; try again later"}) - response.status_code = 429 - response.headers["Retry-After"] = str(wait) - return response + return too_many("too many downloads from this account; try again later", wait) downloads_by_account.record(account) # From the SAME directory the metadata came from, or a drop-in appearing between diff --git a/src/inkwell/common.py b/src/inkwell/common.py index 1c958e0..b6d72ad 100644 --- a/src/inkwell/common.py +++ b/src/inkwell/common.py @@ -1,6 +1,7 @@ from __future__ import annotations -from datetime import datetime +import asyncio +from datetime import datetime, timedelta # Small, dependency-free value coercions shared across the blueprints. Kept in one # place so the "parse an ISO date" / "is this flag truthy" logic has a single @@ -47,3 +48,28 @@ def normalize_email(raw: str | None) -> str: one from a request goes through here, so "Ana@Example.com " and "ana@example.com" are one account at sign-up, sign-in, reset and invite alike.""" return (raw or "").strip().lower() + + +def detach(running: set[asyncio.Task], coro) -> asyncio.Task: + """Start `coro` without waiting for it, held in `running` until it finishes. + + The event loop keeps only a weak reference to a task, so one nobody holds can be + collected mid-flight. Each caller keeps its own set, which is what its `drain()` + awaits in tests. Raises RuntimeError when there is no running loop.""" + task = asyncio.create_task(coro) + running.add(task) + task.add_done_callback(running.discard) + return task + + +def expired_before(now: datetime, retention_days: int) -> datetime | None: + """The cutoff: anything older than this has expired. `None` = retention is off. + + Kept separate from the query so the window arithmetic — including the two ways + to say "never" (0 and negative, the latter reachable by typing a stray minus in + Settings) — is testable without a database. Trash and the audit log both expire + through it. + """ + if retention_days <= 0: + return None + return now - timedelta(days=retention_days) diff --git a/src/inkwell/groups_api.py b/src/inkwell/groups_api.py index ec7d880..6260919 100644 --- a/src/inkwell/groups_api.py +++ b/src/inkwell/groups_api.py @@ -18,7 +18,8 @@ from .db import session_scope from .models.group import Group, GroupMember from .models.user import User from .responses import json_error, not_found, parse_uuid -from .share_sync import bump_note, group_notes, joined_group, left_group, recipients, revoke +from .serialize import serialize_person +from .share_sync import bump_note, group_notes, joined_group, left_group, recipients, revoke_lost bp = Blueprint("groups", __name__, url_prefix="/api/groups") @@ -49,11 +50,17 @@ async def _serialize_groups(db, groups: list[Group]) -> list[dict]: members: dict = {} for group_id, user in rows: members.setdefault(group_id, []).append( - {"id": str(user.id), "display_name": user.display_name, "email": user.email} + serialize_person(user) ) return [{"id": str(gr.id), "name": gr.name, "members": members.get(gr.id, [])} for gr in groups] +async def _get_group(db, group_id: str) -> Group | None: + """The group a route's path names, or None for a malformed or unknown id.""" + gid = parse_uuid(group_id) + return await db.get(Group, gid) if gid else None + + async def _answer(db, group: Group, status: int = 200): return jsonify((await _serialize_groups(db, [group]))[0]), status @@ -87,9 +94,8 @@ async def rename_group(group_id: str): name = _name(await request.get_json(silent=True) or {}) if name is None: return json_error("give the group a name", 400) - gid = parse_uuid(group_id) async with session_scope() as db: - group = await db.get(Group, gid) if gid else None + group = await _get_group(db, group_id) if group is None: return not_found() if await _name_taken(db, name, but=group.id): @@ -104,9 +110,8 @@ async def rename_group(group_id: str): async def delete_group(group_id: str): """Delete the group and every share made to it. Whoever could see a note only through it loses the note, and their devices are told.""" - gid = parse_uuid(group_id) async with session_scope() as db: - group = await db.get(Group, gid) if gid else None + group = await _get_group(db, group_id) if group is None: return not_found() note_ids = await group_notes(db, group.id) @@ -114,7 +119,7 @@ async def delete_group(group_id: str): await db.delete(group) # members and shares go with it (ON DELETE CASCADE) await db.flush() for nid in note_ids: - await revoke(db, nid, before[nid] - await recipients(db, nid)) + await revoke_lost(db, nid, before[nid]) # The owner's devices learn whether the note is still shared at all. await bump_note(db, nid) await db.commit() @@ -125,10 +130,9 @@ async def delete_group(group_id: str): @require_admin async def add_member(group_id: str): data = await request.get_json(silent=True) or {} - gid = parse_uuid(group_id) uid = parse_uuid(str(data.get("user_id") or "")) async with session_scope() as db: - group = await db.get(Group, gid) if gid else None + group = await _get_group(db, group_id) if group is None: return not_found() if uid is None or await db.get(User, uid) is None: @@ -147,10 +151,9 @@ async def add_member(group_id: str): @bp.delete("//members/") @require_admin async def remove_member(group_id: str, user_id: str): - gid = parse_uuid(group_id) uid = parse_uuid(user_id) async with session_scope() as db: - group = await db.get(Group, gid) if gid else None + group = await _get_group(db, group_id) if group is None or uid is None: return not_found() membership = await db.scalar( diff --git a/src/inkwell/mailer.py b/src/inkwell/mailer.py index 448a2a1..57fd665 100644 --- a/src/inkwell/mailer.py +++ b/src/inkwell/mailer.py @@ -17,6 +17,7 @@ import ssl from dataclasses import dataclass from email.message import EmailMessage +from .common import detach from .settings import get_setting, mail_configured logger = logging.getLogger(__name__) @@ -91,9 +92,7 @@ def send_later(job) -> None: For the forgot-password route, where waiting would also be an oracle: an address with an account would answer as slowly as a mail server, and one without, at once. `job` handles and logs its own failures.""" - task = asyncio.create_task(job) - _running.add(task) - task.add_done_callback(_running.discard) + detach(_running, job) async def drain() -> None: diff --git a/src/inkwell/responses.py b/src/inkwell/responses.py index f464377..fed6301 100644 --- a/src/inkwell/responses.py +++ b/src/inkwell/responses.py @@ -19,6 +19,12 @@ def not_found(): return json_error("not found", 404) +def too_many(message: str, retry_after: int): + """429 with `Retry-After`, the one thing a client that was slowed down needs to + know: how long to wait. Never which limit it hit.""" + return jsonify({"error": message}), 429, {"Retry-After": str(retry_after)} + + def parse_uuid(raw: object) -> uuid.UUID | None: """Parse a path/body UUID, returning None on anything malformed. Pair with not_found() for the ubiquitous 'bad id in the URL → 404' guard.""" diff --git a/src/inkwell/retention.py b/src/inkwell/retention.py index dada529..8dfe0a8 100644 --- a/src/inkwell/retention.py +++ b/src/inkwell/retention.py @@ -22,12 +22,13 @@ from __future__ import annotations import asyncio import logging -from datetime import datetime, timedelta, timezone +from datetime import datetime, timezone from sqlalchemy import delete as sa_delete from sqlalchemy import select from . import audit +from .common import expired_before from .config import Config from .db import session_scope from .models.label import NoteLabel @@ -56,18 +57,6 @@ SWEEP_STARTUP_DELAY_SECONDS = 60 SWEEP_BATCH = 200 -def expired_before(now: datetime, retention_days: int) -> datetime | None: - """The cutoff: trash older than this has expired. `None` = retention is off. - - Kept separate from the query so the window arithmetic — including the two ways - to say "never" (0 and negative, the latter reachable by typing a stray minus in - Settings) — is testable without a database. - """ - if retention_days <= 0: - return None - return now - timedelta(days=retention_days) - - async def purge_note(db, note: Note, edited_at: datetime | None = None) -> None: """Turn a note into a content-less tombstone: delete its children (and the attachment files on disk), clear its content, stamp `purged_at`. diff --git a/src/inkwell/serialize.py b/src/inkwell/serialize.py index e53b862..92857dd 100644 --- a/src/inkwell/serialize.py +++ b/src/inkwell/serialize.py @@ -7,6 +7,7 @@ from __future__ import annotations from .common import iso from .models.label import Label +from .models.user import User def serialize_label(label: Label) -> dict: @@ -24,3 +25,10 @@ def serialize_label_sync(label: Label) -> dict: "purged_at": iso(label.purged_at), "created_at": iso(label.created_at), } + + +def serialize_person(user: User) -> dict: + """Someone as another member sees them: a name and an email, which is what + choosing them takes and all a fellow member needs to know. The share directory, + a share's recipient and a group's members all show people this way.""" + return {"id": str(user.id), "display_name": user.display_name, "email": user.email} diff --git a/src/inkwell/settings.py b/src/inkwell/settings.py index 3d76510..37c26a1 100644 --- a/src/inkwell/settings.py +++ b/src/inkwell/settings.py @@ -3,7 +3,7 @@ from __future__ import annotations import json import secrets from dataclasses import dataclass -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from typing import Any, Literal from .common import coerce_bool @@ -394,6 +394,12 @@ def live(key: str) -> Any: return _live[key] +def apply_session_ttl(app, days) -> None: + """Sessions signed from now on last `days` days. At boot and on every admin save, + so a changed lifetime applies without a restart (rule 25).""" + app.config["PERMANENT_SESSION_LIFETIME"] = timedelta(days=int(days)) + + async def refresh_live(db) -> None: """Re-read the hot settings into the cache. Called at boot and after every save.""" for key in _LIVE_KEYS: diff --git a/src/inkwell/settings_api.py b/src/inkwell/settings_api.py index 2af9434..d0fc08b 100644 --- a/src/inkwell/settings_api.py +++ b/src/inkwell/settings_api.py @@ -1,7 +1,5 @@ from __future__ import annotations -from datetime import timedelta - import logging from quart import Blueprint, current_app, g, jsonify, request @@ -11,7 +9,7 @@ from .db import session_scope from .mailer import mail_settings, send from .responses import json_error from .models.user import User -from .settings import get_admin_settings, get_setting, refresh_live, set_settings, validate_updates +from .settings import apply_session_ttl, get_admin_settings, get_setting, refresh_live, set_settings, validate_updates logger = logging.getLogger(__name__) @@ -50,7 +48,7 @@ async def update_settings(): # Apply the live-tunable knob without a restart (rule 25). if "session_ttl_days" in clean: - current_app.config["PERMANENT_SESSION_LIFETIME"] = timedelta(days=int(clean["session_ttl_days"])) + apply_session_ttl(current_app, clean["session_ttl_days"]) return jsonify({"settings": result}) diff --git a/src/inkwell/share_sync.py b/src/inkwell/share_sync.py index 9d4451c..5652c83 100644 --- a/src/inkwell/share_sync.py +++ b/src/inkwell/share_sync.py @@ -80,6 +80,12 @@ async def granted(db, note_id: uuid.UUID, user_id: uuid.UUID) -> None: # cursor, and losing it leaves a revocation. +async def revoke_lost(db, note_id: uuid.UUID, before: set[uuid.UUID]) -> None: + """After a share or a group goes: tell the devices of whoever could see the note + `before` and no longer can. Someone still reached another way keeps it.""" + await revoke(db, note_id, before - await recipients(db, note_id)) + + async def group_notes(db, group_id: uuid.UUID) -> list[uuid.UUID]: """The notes shared with this group.""" rows = await db.scalars( diff --git a/src/inkwell/shares_api.py b/src/inkwell/shares_api.py index ce44a11..6a9798f 100644 --- a/src/inkwell/shares_api.py +++ b/src/inkwell/shares_api.py @@ -25,15 +25,12 @@ from .models.share import Share from .models.user import User from .notes.helpers import _get_owned from .responses import json_error, not_found, parse_uuid -from .share_sync import bump_note, granted, recipients, revoke +from .serialize import serialize_person +from .share_sync import bump_note, granted, recipients, revoke_lost bp = Blueprint("shares", __name__, url_prefix="/api") -def _member(user: User) -> dict: - return {"id": str(user.id), "display_name": user.display_name, "email": user.email} - - def _group(group: Group, counts: dict) -> dict: return {"id": str(group.id), "name": group.name, "member_count": counts.get(group.id, 0)} @@ -64,7 +61,7 @@ async def directory(): ).all() groups = (await db.scalars(select(Group).order_by(func.lower(Group.name)))).all() counts = await _member_counts(db, [gr.id for gr in groups]) - return jsonify({"members": [_member(u) for u in users], "groups": [_group(gr, counts) for gr in groups]}) + return jsonify({"members": [serialize_person(u) for u in users], "groups": [_group(gr, counts) for gr in groups]}) async def _shares_of(db, note_id: uuid.UUID) -> list[dict]: @@ -83,7 +80,7 @@ async def _shares_of(db, note_id: uuid.UUID) -> list[dict]: return [ { "id": str(share.id), - "member": _member(user) if user is not None else None, + "member": serialize_person(user) if user is not None else None, "group": _group(group, counts) if group is not None else None, "permission": share.permission, "created_at": iso(share.created_at), @@ -167,9 +164,7 @@ async def unshare_note(note_id: str, share_id: str): before = await recipients(db, note.id) await db.delete(share) await db.flush() - # Only the people who can no longer see it: someone still reached through a - # group keeps the note on their devices. - await revoke(db, note.id, before - await recipients(db, note.id)) + await revoke_lost(db, note.id, before) # The owner's devices learn the note is no longer shared. await bump_note(db, note.id) await db.commit() diff --git a/src/inkwell/unfurl_queue.py b/src/inkwell/unfurl_queue.py index 37090d8..8de071f 100644 --- a/src/inkwell/unfurl_queue.py +++ b/src/inkwell/unfurl_queue.py @@ -30,6 +30,7 @@ import uuid from sqlalchemy import select +from .common import detach from .db import session_scope from .models.note import Note from .models.note_link_preview import NoteLinkPreview @@ -118,10 +119,8 @@ def schedule(note_id: uuid.UUID, body: str | None) -> None: if not body or not detect_urls(body): return try: - task = asyncio.create_task(_unfurl_new_urls(note_id, body)) + detach(_running, _unfurl_new_urls(note_id, body)) except RuntimeError: # No running loop — a script or a test calling the write path directly. The # note is saved either way; only the preview is skipped. return - _running.add(task) - task.add_done_callback(_running.discard) diff --git a/tests/test_retention.py b/tests/test_retention.py index 3b68a36..2f06371 100644 --- a/tests/test_retention.py +++ b/tests/test_retention.py @@ -2,11 +2,11 @@ from datetime import datetime, timedelta, timezone import pytest +from inkwell.common import expired_before from inkwell.retention import ( SWEEP_BATCH, SWEEP_INTERVAL_SECONDS, SWEEP_STARTUP_DELAY_SECONDS, - expired_before, sweep_expired_trash, ) from inkwell.settings import REGISTRY, get_public_config, validate_updates