"""Durable, account-scoped unfiled issue drafts and ordered evidence.""" import base64 import binascii import json import os import re import sqlite3 from datetime import datetime from pathlib import Path from cryptography.exceptions import InvalidTag from cryptography.hazmat.primitives.ciphers.aead import AESGCM from src.private_state import connect_private_sqlite _DRAFT_ID = re.compile(r"^[A-Za-z0-9_-]{1,100}$") _REPOSITORY = re.compile(r"^[A-Za-z0-9_.-]+/[A-Za-z0-9_.-]+$") _CONTENT_TYPES = {"image/png", "image/jpeg", "image/webp"} _KEY_ID = re.compile(r"^[A-Za-z0-9_-]{1,32}$") class UnfiledDraftConflict(ValueError): """Raised when a client attempts to replace a stale draft collection.""" def __init__(self, snapshot: dict): super().__init__("unfiled drafts changed on another device") self.snapshot = snapshot class UnfiledDraftEncryptionError(RuntimeError): """Raised when private draft state cannot be authenticated and decrypted.""" def decode_unfiled_draft_encryption_key(encoded: str) -> bytes: """Decode the independently injected 256-bit draft-encryption key.""" try: key = base64.b64decode(encoded, validate=True) except (binascii.Error, ValueError) as error: raise UnfiledDraftEncryptionError( "unfiled draft encryption key is invalid" ) from error if len(key) != 32: raise UnfiledDraftEncryptionError( "unfiled draft encryption key must decode to exactly 32 bytes" ) return key def decode_unfiled_draft_encryption_keyring( encoded: str, active_key_id: str ) -> tuple[dict[str, bytes], str]: """Decode a bounded JSON object of named keys without exposing key material.""" def unique_object(pairs): value = {} for key, item in pairs: if key in value: raise ValueError("duplicate key id") value[key] = item return value try: raw = json.loads(encoded, object_pairs_hook=unique_object) if not isinstance(raw, dict) or not raw or len(raw) > 8: raise ValueError("invalid keyring size") if active_key_id not in raw: raise ValueError("active key is unavailable") keys = {} for key_id, value in raw.items(): if not isinstance(key_id, str) or not _KEY_ID.fullmatch(key_id): raise ValueError("invalid key id") if not isinstance(value, str): raise ValueError("invalid key value") keys[key_id] = decode_unfiled_draft_encryption_key(value) if len(set(keys.values())) != len(keys): raise ValueError("duplicate encryption key") return keys, active_key_id except (json.JSONDecodeError, TypeError, ValueError, UnfiledDraftEncryptionError) as error: raise UnfiledDraftEncryptionError( "unfiled draft encryption keyring is invalid" ) from error class UnfiledDraftStore: def __init__( self, path: str | Path, *, encryption_key: bytes | None = None, encryption_keys: dict[str, bytes] | None = None, active_key_id: str | None = None, limit: int = 20, max_total_bytes: int = 12 * 1024 * 1024, timeout: float = 1.0, ): self.path = Path(path) if encryption_keys is None: if not isinstance(encryption_key, bytes) or len(encryption_key) != 32: raise UnfiledDraftEncryptionError( "unfiled draft encryption requires exactly 32 key bytes" ) self._ciphers = {"legacy": AESGCM(encryption_key)} self._active_key_id = None else: if ( not encryption_keys or len(encryption_keys) > 8 or active_key_id not in encryption_keys or any( not isinstance(key_id, str) or not _KEY_ID.fullmatch(key_id) or not isinstance(key, bytes) or len(key) != 32 for key_id, key in encryption_keys.items() ) ): raise UnfiledDraftEncryptionError("unfiled draft encryption keyring is invalid") self._ciphers = { key_id: AESGCM(key) for key_id, key in encryption_keys.items() } self._active_key_id = active_key_id self.limit = limit self.max_total_bytes = max_total_bytes self.timeout = timeout self._initialize() def _initialize(self) -> None: with self._connect() as connection: connection.execute("PRAGMA journal_mode=WAL") connection.execute( """CREATE TABLE IF NOT EXISTS unfiled_drafts ( login TEXT PRIMARY KEY, revision INTEGER NOT NULL, drafts TEXT NOT NULL )""" ) def _connect(self) -> sqlite3.Connection: return connect_private_sqlite(self.path, timeout=self.timeout) @staticmethod def _login(login: str) -> str: normalized = login.strip().lower() if not normalized: raise ValueError("login is required") return normalized @staticmethod def _aad(login: str, revision: int) -> bytes: return f"stackchain:unfiled-drafts:v1\0{login}\0{revision}".encode() @staticmethod def _v2_aad(login: str, revision: int, key_id: str) -> bytes: return f"stackchain:unfiled-drafts:v2\0{key_id}\0{login}\0{revision}".encode() def _encrypt(self, login: str, revision: int, drafts: list[dict]) -> str: plaintext = json.dumps(drafts, separators=(",", ":")).encode() nonce = os.urandom(12) if self._active_key_id is None: sealed = nonce + self._ciphers["legacy"].encrypt( nonce, plaintext, self._aad(login, revision) ) return "v1:" + base64.urlsafe_b64encode(sealed).decode() key_id = self._active_key_id sealed = nonce + self._ciphers[key_id].encrypt( nonce, plaintext, self._v2_aad(login, revision, key_id) ) return f"v2:{key_id}:" + base64.urlsafe_b64encode(sealed).decode() def _decrypt(self, login: str, revision: int, payload: str) -> list[dict]: try: if payload.startswith("v1:"): encoded = payload[3:] cipher = self._ciphers.get("legacy") aad = self._aad(login, revision) elif payload.startswith("v2:"): _version, key_id, encoded = payload.split(":", 2) cipher = self._ciphers.get(key_id) aad = self._v2_aad(login, revision, key_id) else: raise ValueError("encrypted payload version is invalid") if cipher is None: raise ValueError("encrypted payload key is unavailable") sealed = base64.b64decode(encoded, altchars=b"-_", validate=True) if len(sealed) < 12 + 16: raise ValueError("encrypted payload is too short") plaintext = cipher.decrypt(sealed[:12], sealed[12:], aad) drafts = json.loads(plaintext) if not isinstance(drafts, list): raise ValueError("decrypted payload is not a collection") return drafts except (binascii.Error, InvalidTag, UnicodeDecodeError, ValueError, json.JSONDecodeError) as error: raise UnfiledDraftEncryptionError( "unfiled drafts could not be decrypted" ) from error def _snapshot(self, login: str, row) -> tuple[dict, bool]: if row is None: return {"revision": 0, "drafts": []}, False revision, payload = int(row[0]), row[1] if payload.startswith(("v1:", "v2:")): drafts = self._decrypt(login, revision, payload) active_prefix = ( "v1:" if self._active_key_id is None else f"v2:{self._active_key_id}:" ) return {"revision": revision, "drafts": drafts}, not payload.startswith(active_prefix) try: drafts = json.loads(payload) except (TypeError, json.JSONDecodeError) as error: raise UnfiledDraftEncryptionError( "legacy unfiled drafts could not be decoded" ) from error if not isinstance(drafts, list): raise UnfiledDraftEncryptionError( "legacy unfiled drafts could not be decoded" ) return {"revision": revision, "drafts": drafts}, True def get(self, login: str) -> dict: login = self._login(login) with self._connect() as connection: connection.execute("BEGIN IMMEDIATE") row = connection.execute( "SELECT revision, drafts FROM unfiled_drafts WHERE login = ?", (login,), ).fetchone() snapshot, legacy = self._snapshot(login, row) if legacy: connection.execute( "UPDATE unfiled_drafts SET drafts = ? WHERE login = ?", ( self._encrypt(login, snapshot["revision"], snapshot["drafts"]), login, ), ) return snapshot def rewrap_all(self) -> dict[str, int]: """Rewrap every readable row and return aggregate counts only.""" with self._connect() as connection: rows = connection.execute( "SELECT login, drafts FROM unfiled_drafts ORDER BY login" ).fetchall() result = {"total": len(rows), "migrated": 0, "current": 0, "failed": 0} active_prefix = ( "v1:" if self._active_key_id is None else f"v2:{self._active_key_id}:" ) for login, payload in rows: was_current = payload.startswith(active_prefix) try: self.get(login) except UnfiledDraftEncryptionError: result["failed"] += 1 else: result["current" if was_current else "migrated"] += 1 return result @staticmethod def _filing_plan(raw: object) -> dict | None: if raw is None: return None if not isinstance(raw, dict): raise ValueError("filing plan is invalid") repository = raw.get("repository") if not isinstance(repository, str) or not _REPOSITORY.fullmatch(repository): raise ValueError("filing repository is invalid") label_ids = raw.get("label_ids", []) if not isinstance(label_ids, list) or len(label_ids) > 20: raise ValueError("filing labels are invalid") if any(not isinstance(item, int) or isinstance(item, bool) or item < 1 for item in label_ids): raise ValueError("filing labels are invalid") if len(set(label_ids)) != len(label_ids): raise ValueError("filing labels are invalid") plan = {"repository": repository, "label_ids": label_ids} for key in ("milestone_id", "estimate_minutes"): value = raw.get(key) if value is None: continue upper = 1440 if key == "estimate_minutes" else None lower = 5 if key == "estimate_minutes" else 1 if ( not isinstance(value, int) or isinstance(value, bool) or value < lower or (upper is not None and value > upper) ): raise ValueError(f"filing {key.replace('_', ' ')} is invalid") plan[key] = value due_date = raw.get("due_date") if due_date is not None: try: datetime.strptime(due_date, "%Y-%m-%d") except (TypeError, ValueError) as error: raise ValueError("filing due date is invalid") from error plan["due_date"] = due_date for key, limit in { "template_name": 80, "template_id": 80, "captured_body": 10_000, "assignee_name": 255, }.items(): value = raw.get(key) if value is not None: if not isinstance(value, str) or not value.strip() or len(value) > limit: raise ValueError(f"filing {key.replace('_', ' ')} is invalid") plan[key] = value unassigned = raw.get("unassigned", False) if not isinstance(unassigned, bool): raise ValueError("filing owner intent is invalid") assignee = raw.get("assignee") if unassigned and assignee is not None: raise ValueError("filing owner intent is invalid") if unassigned: plan["unassigned"] = True elif assignee is not None: if not isinstance(assignee, str) or not re.fullmatch(r"[A-Za-z0-9_.-]+", assignee): raise ValueError("filing assignee is invalid") plan["assignee"] = assignee completion_intent = raw.get("completion_intent") if completion_intent is not None: if completion_intent not in {"create", "create-and-start"}: raise ValueError("filing completion intent is invalid") plan["completion_intent"] = completion_intent return plan def _normalize(self, drafts: list[dict]) -> list[dict]: if not isinstance(drafts, list): raise ValueError("drafts must be a list") if len(drafts) > self.limit: raise ValueError(f"unfiled drafts are limited to {self.limit}") normalized = [] seen = set() decoded_total = 0 for raw in drafts: if not isinstance(raw, dict): raise ValueError("draft must be an object") draft_id = raw.get("id") if not isinstance(draft_id, str) or not _DRAFT_ID.fullmatch(draft_id): raise ValueError("draft id is invalid") if draft_id in seen: raise ValueError("draft ids must be unique") title = raw.get("title") body = raw.get("body", "") saved_at = raw.get("saved_at") if not isinstance(title, str) or len(title.strip()) > 255: raise ValueError("title is invalid") if not isinstance(body, str) or len(body) > 10_000: raise ValueError("body is invalid") if not isinstance(saved_at, int) or isinstance(saved_at, bool) or saved_at < 0: raise ValueError("saved_at is invalid") filing_plan = self._filing_plan(raw.get("filing_plan")) blockers = raw.get("blockers", []) if not isinstance(blockers, list) or len(blockers) > 5: raise ValueError("blockers are invalid") clean_blockers = [] for blocker in blockers: repository = blocker.get("repository") if isinstance(blocker, dict) else None number = blocker.get("number") if isinstance(blocker, dict) else None blocker_title = blocker.get("title", "") if isinstance(blocker, dict) else "" if not isinstance(repository, str) or not _REPOSITORY.fullmatch(repository): raise ValueError("blocker repository is invalid") if not isinstance(number, int) or isinstance(number, bool) or number < 1: raise ValueError("blocker number is invalid") if not isinstance(blocker_title, str) or len(blocker_title) > 255: raise ValueError("blocker title is invalid") clean_blockers.append({"repository": repository, "number": number, "title": blocker_title}) evidence = raw.get("evidence", []) if not isinstance(evidence, list) or len(evidence) > 5: raise ValueError("evidence is invalid") clean_evidence = [] for item in evidence: if not isinstance(item, dict): raise ValueError("evidence is invalid") filename = item.get("filename") content_type = item.get("content_type") note = item.get("note", "") data = item.get("data") if not isinstance(filename, str) or not filename or len(filename) > 255: raise ValueError("evidence filename is invalid") if content_type not in _CONTENT_TYPES: raise ValueError("evidence content type is invalid") if not isinstance(note, str) or len(note) > 240: raise ValueError("evidence note is invalid") if not isinstance(data, str): raise ValueError("evidence data is invalid") try: decoded_total += len(base64.b64decode(data, validate=True)) except (binascii.Error, ValueError) as error: raise ValueError("evidence data must be valid base64") from error if decoded_total > self.max_total_bytes: raise ValueError("synchronized evidence is too large") clean_evidence.append({ "filename": filename, "content_type": content_type, **({"note": note} if note else {}), "data": data, }) if not title.strip() and not clean_evidence: raise ValueError("title or evidence is required") seen.add(draft_id) normalized.append({ "id": draft_id, "title": title.strip(), "body": body, "saved_at": saved_at, **({"filing_plan": filing_plan} if filing_plan else {}), **({"blockers": clean_blockers} if clean_blockers else {}), **({"evidence": clean_evidence} if clean_evidence else {}), }) return normalized def replace(self, login: str, expected_revision: int, drafts: list[dict]) -> dict: login = self._login(login) if not isinstance(expected_revision, int) or isinstance(expected_revision, bool) or expected_revision < 0: raise ValueError("revision is invalid") normalized = self._normalize(drafts) with self._connect() as connection: connection.execute("BEGIN IMMEDIATE") row = connection.execute( "SELECT revision, drafts FROM unfiled_drafts WHERE login = ?", (login,) ).fetchone() current, _legacy = self._snapshot(login, row) if current["revision"] != expected_revision: raise UnfiledDraftConflict(current) revision = expected_revision + 1 serialized = self._encrypt(login, revision, normalized) connection.execute( "INSERT INTO unfiled_drafts(login, revision, drafts) VALUES (?, ?, ?) " "ON CONFLICT(login) DO UPDATE SET revision=excluded.revision, drafts=excluded.drafts", (login, revision, serialized), ) return {"revision": revision, "drafts": normalized}