from __future__ import annotations from dataclasses import dataclass from datetime import datetime from threading import Lock from sqlalchemy import func from sqlalchemy.orm import Session from govoplan_core.security.time import ensure_aware_utc, utc_now from govoplan_mail.backend.db.models import MailMailboxFolderIndex, MailMailboxMessageIndex from govoplan_mail.backend.sending.imap import ImapFolderListResult, ImapMailboxInfo, ImapMailboxMessageListResult, ImapMailboxMessageSummary MAILBOX_INDEX_TTL_SECONDS = 30 _refresh_lock = Lock() _refreshing_keys: set[tuple[str, str, str]] = set() @dataclass(frozen=True, slots=True) class CachedFolderList: folders: list[ImapMailboxInfo] indexed_at: datetime | None stale: bool @dataclass(frozen=True, slots=True) class CachedMessagePage: folder: str messages: list[ImapMailboxMessageSummary] total_count: int offset: int limit: int uidvalidity: str | None indexed_at: datetime | None stale: bool def clear_mailbox_index(session: Session, *, profile_id: str) -> tuple[int, int]: """Remove every cached mailbox row for a profile in the caller's transaction. System profiles can be used by more than one tenant, so invalidation is intentionally profile-wide rather than limited to the tenant making the transport change. Message rows are removed before their folder metadata. """ deleted_messages = ( session.query(MailMailboxMessageIndex) .filter(MailMailboxMessageIndex.profile_id == profile_id) .delete(synchronize_session=False) ) deleted_folders = ( session.query(MailMailboxFolderIndex) .filter(MailMailboxFolderIndex.profile_id == profile_id) .delete(synchronize_session=False) ) return int(deleted_folders or 0), int(deleted_messages or 0) def begin_mailbox_refresh(tenant_id: str, profile_id: str, folder: str) -> bool: key = (tenant_id, profile_id, folder) with _refresh_lock: if key in _refreshing_keys: return False _refreshing_keys.add(key) return True def finish_mailbox_refresh(tenant_id: str, profile_id: str, folder: str) -> None: with _refresh_lock: _refreshing_keys.discard((tenant_id, profile_id, folder)) def cache_mailbox_folders( session: Session, *, tenant_id: str, profile_id: str, result: ImapFolderListResult, indexed_at: datetime | None = None, ) -> None: indexed_at = indexed_at or utc_now() existing = { row.folder: row for row in session.query(MailMailboxFolderIndex) .filter(MailMailboxFolderIndex.tenant_id == tenant_id, MailMailboxFolderIndex.profile_id == profile_id) .all() } for folder in result.folders: row = existing.get(folder.name) if row is None: row = MailMailboxFolderIndex(tenant_id=tenant_id, profile_id=profile_id, folder=folder.name) row.flags = list(folder.flags or []) row.message_count = folder.message_count row.unseen_count = folder.unseen_count row.indexed_at = indexed_at session.add(row) def cache_mailbox_messages( session: Session, *, tenant_id: str, profile_id: str, result: ImapMailboxMessageListResult, indexed_at: datetime | None = None, ) -> None: indexed_at = indexed_at or utc_now() session.flush() folder_row = ( session.query(MailMailboxFolderIndex) .filter( MailMailboxFolderIndex.tenant_id == tenant_id, MailMailboxFolderIndex.profile_id == profile_id, MailMailboxFolderIndex.folder == result.folder, ) .one_or_none() ) if folder_row is None: folder_row = MailMailboxFolderIndex(tenant_id=tenant_id, profile_id=profile_id, folder=result.folder) folder_row.message_count = result.total_count folder_row.uidvalidity = result.uidvalidity folder_row.message_indexed_at = indexed_at session.add(folder_row) uids = [message.uid for message in result.messages] if result.total_count <= 0: ( session.query(MailMailboxMessageIndex) .filter( MailMailboxMessageIndex.tenant_id == tenant_id, MailMailboxMessageIndex.profile_id == profile_id, MailMailboxMessageIndex.folder == result.folder, ) .delete(synchronize_session=False) ) return window_count = max(0, min(result.limit, result.total_count - result.offset)) if window_count: stale_query = session.query(MailMailboxMessageIndex).filter( MailMailboxMessageIndex.tenant_id == tenant_id, MailMailboxMessageIndex.profile_id == profile_id, MailMailboxMessageIndex.folder == result.folder, MailMailboxMessageIndex.sort_position >= result.offset, MailMailboxMessageIndex.sort_position < result.offset + window_count, ) if uids: stale_query = stale_query.filter(MailMailboxMessageIndex.uid.notin_(uids)) for row in stale_query.all(): session.delete(row) existing = {} if uids: existing = { row.uid: row for row in session.query(MailMailboxMessageIndex) .filter( MailMailboxMessageIndex.tenant_id == tenant_id, MailMailboxMessageIndex.profile_id == profile_id, MailMailboxMessageIndex.folder == result.folder, MailMailboxMessageIndex.uid.in_(uids), ) .all() } for index, message in enumerate(result.messages): row = existing.get(message.uid) if row is None: row = MailMailboxMessageIndex( tenant_id=tenant_id, profile_id=profile_id, folder=result.folder, uid=message.uid, ) row.uid_int = _uid_int(message.uid) row.sort_position = result.offset + index row.subject = message.subject row.from_header = message.from_header row.to_header = message.to_header row.cc_header = message.cc_header row.date = message.date row.message_id = message.message_id row.flags = list(message.flags or []) row.size_bytes = message.size_bytes row.body_preview = message.body_preview row.attachment_count = message.attachment_count row.indexed_at = indexed_at session.add(row) def cached_mailbox_folders( session: Session, *, tenant_id: str, profile_id: str, max_age_seconds: int = MAILBOX_INDEX_TTL_SECONDS, ) -> CachedFolderList | None: rows = ( session.query(MailMailboxFolderIndex) .filter( MailMailboxFolderIndex.tenant_id == tenant_id, MailMailboxFolderIndex.profile_id == profile_id, MailMailboxFolderIndex.indexed_at.isnot(None), ) .order_by(MailMailboxFolderIndex.folder.asc()) .all() ) if not rows: return None indexed_at = max((row.indexed_at for row in rows if row.indexed_at), default=None) return CachedFolderList( folders=[ ImapMailboxInfo(name=row.folder, flags=list(row.flags or []), message_count=row.message_count, unseen_count=row.unseen_count) for row in rows ], indexed_at=indexed_at, stale=_is_stale(indexed_at, max_age_seconds=max_age_seconds), ) def cached_mailbox_message_page( session: Session, *, tenant_id: str, profile_id: str, folder: str, limit: int, offset: int, max_age_seconds: int = MAILBOX_INDEX_TTL_SECONDS, ) -> CachedMessagePage | None: folder_row = ( session.query(MailMailboxFolderIndex) .filter( MailMailboxFolderIndex.tenant_id == tenant_id, MailMailboxFolderIndex.profile_id == profile_id, MailMailboxFolderIndex.folder == folder, ) .one_or_none() ) if folder_row is None or folder_row.message_indexed_at is None: return None total_count = folder_row.message_count if total_count is None: total_count = int( session.query(func.count(MailMailboxMessageIndex.id)) .filter( MailMailboxMessageIndex.tenant_id == tenant_id, MailMailboxMessageIndex.profile_id == profile_id, MailMailboxMessageIndex.folder == folder, ) .scalar() or 0 ) rows = ( session.query(MailMailboxMessageIndex) .filter( MailMailboxMessageIndex.tenant_id == tenant_id, MailMailboxMessageIndex.profile_id == profile_id, MailMailboxMessageIndex.folder == folder, ) .order_by(MailMailboxMessageIndex.sort_position.asc(), MailMailboxMessageIndex.uid_int.desc(), MailMailboxMessageIndex.uid.desc()) .offset(offset) .limit(limit) .all() ) expected_count = max(0, min(limit, total_count - offset)) if len(rows) < expected_count: return None indexed_at = min((row.indexed_at for row in rows if row.indexed_at), default=folder_row.message_indexed_at) return CachedMessagePage( folder=folder, messages=[_message_from_index(row) for row in rows], total_count=total_count, offset=offset, limit=limit, uidvalidity=folder_row.uidvalidity, indexed_at=indexed_at, stale=_is_stale(indexed_at, max_age_seconds=max_age_seconds), ) def _message_from_index(row: MailMailboxMessageIndex) -> ImapMailboxMessageSummary: return ImapMailboxMessageSummary( uid=row.uid, folder=row.folder, subject=row.subject, from_header=row.from_header, to_header=row.to_header, cc_header=row.cc_header, date=row.date, message_id=row.message_id, flags=list(row.flags or []), size_bytes=row.size_bytes, body_preview=row.body_preview, attachment_count=row.attachment_count, ) def _uid_int(uid: str) -> int: try: return int(str(uid)) except ValueError: return 0 def _is_stale(indexed_at: datetime | None, *, max_age_seconds: int) -> bool: aware = ensure_aware_utc(indexed_at) if aware is None: return True return (utc_now() - aware).total_seconds() > max_age_seconds