Source code for litestar_queues.backends.advanced_alchemy.backend

"""Advanced Alchemy queue backend."""

from contextlib import asynccontextmanager
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, cast
from uuid import UUID

from advanced_alchemy.exceptions import IntegrityError as AdvancedAlchemyIntegrityError
from sqlalchemy import delete, text, update
from sqlalchemy import inspect as sqlalchemy_inspect
from sqlalchemy.engine import make_url
from sqlalchemy.exc import IntegrityError as SQLAlchemyIntegrityError

from litestar_queues.backends.advanced_alchemy._notifications import (
    SUPPORTED_NOTIFY_DRIVERS,
    create_notification_listener,
)
from litestar_queues.backends.advanced_alchemy.config import SQLAlchemyBackendConfig
from litestar_queues.backends.advanced_alchemy.event_log import AdvancedAlchemyQueueEventLog
from litestar_queues.backends.advanced_alchemy.mixins import (
    QueueEventHistoryModelMixin,
    QueueMaintenanceModelMixin,
    QueueTaskModelMixin,
    QueueTaskReservationModelMixin,
)
from litestar_queues.backends.advanced_alchemy.service import (
    QueueEventLogService,
    QueueTaskReservationService,
    QueueTaskService,
)
from litestar_queues.backends.base import BaseQueueBackend
from litestar_queues.exceptions import QueueConfigurationError
from litestar_queues.models import HeartbeatTouchResult, QueueBackendCapabilities, TaskReservation
from litestar_queues.observability import create_observability_runtime

if TYPE_CHECKING:
    from collections.abc import AsyncIterator, Mapping, Sequence
    from datetime import timedelta

    from sqlalchemy.ext.asyncio import AsyncSession

    from litestar_queues.backends.advanced_alchemy._notifications import NotificationListener
    from litestar_queues.config import QueueConfig
    from litestar_queues.events import EventHistoryConfig
    from litestar_queues.models import (
        HeartbeatTouch,
        QueuedTaskRecord,
        QueueStatistics,
        StaleTaskRecoveryResult,
        TaskRequest,
    )
    from litestar_queues.observability import QueueObservabilityRuntimeProtocol

__all__ = ("SQLAlchemyBackend",)

_POSTGRES_NOTIFY_BACKEND = "postgres-listen-notify"
_POSTGRES_NOTIFY_PAYLOAD = "tasks"


def _utc_now() -> "datetime":
    return datetime.now(timezone.utc)


[docs] class SQLAlchemyBackend(BaseQueueBackend): """SQLAlchemy queue backend using Advanced Alchemy services.""" _model_class: "type[QueueTaskModelMixin]" _service_class: 'type["QueueTaskService"]' _event_history_model_class: "type[QueueEventHistoryModelMixin]" _event_log_service_class: 'type["QueueEventLogService"]' _maintenance_model_class: "type[QueueMaintenanceModelMixin]" _task_reservation_model_class: "type[QueueTaskReservationModelMixin]" _task_reservation_service_class: 'type["QueueTaskReservationService"]' __slots__ = ( "_event_history_model_class", "_event_log", "_event_log_service_class", "_event_poll_interval", "_heartbeat_session_maker", "_maintenance_model_class", "_model_class", "_notification_listener", "_notifications", "_observability_runtime", "_opened", "_service_class", "_sqlalchemy_config", "_task_reservation_model_class", "_task_reservation_service_class", "_wakeup_channel", )
[docs] def __init__( self, config: "QueueConfig | None" = None, *, backend_config: "SQLAlchemyBackendConfig | None" = None ) -> "None": super().__init__(config=config) backend_config = backend_config or SQLAlchemyBackendConfig() self._sqlalchemy_config = backend_config.sqlalchemy_config self._heartbeat_session_maker = backend_config.heartbeat_session_maker self._model_class, self._service_class = self._resolve_model_classes(backend_config.model_class) self._event_history_model_class, self._event_log_service_class = self._resolve_event_history_model_classes( backend_config.event_history_model_class ) self._maintenance_model_class = self._resolve_maintenance_model_class(backend_config.maintenance_model_class) self._task_reservation_model_class, self._task_reservation_service_class = ( self._resolve_task_reservation_model_classes(backend_config.task_reservation_model_class) ) self._notifications = backend_config.worker_wakeups self._wakeup_channel = ( backend_config.wakeup_channel if backend_config.wakeup_channel is not None else config.names.database_channel("tasks") if config is not None else "litestar_queues_tasks" ) self._event_poll_interval = backend_config.wakeup_poll_interval self._notification_listener: "NotificationListener | None" = None self._observability_runtime: "QueueObservabilityRuntimeProtocol | None" = None self._event_log: "AdvancedAlchemyQueueEventLog | None" = None self._opened = False
@property def capabilities(self) -> "QueueBackendCapabilities": """Backend behavior capabilities.""" notifications_enabled = self._notifications_supported() return QueueBackendCapabilities( supports_worker_wakeups=notifications_enabled, wakeup_backend=_POSTGRES_NOTIFY_BACKEND if notifications_enabled else None, wakeups_durable=False, supports_maintenance=True, )
[docs] async def open(self) -> "bool": """Open Advanced Alchemy resources. Returns: True when resources are ready. """ if self._opened: return True self._ensure_configured() self._opened = True return True
[docs] async def close(self) -> "None": """Close backend-owned resources.""" if self._notification_listener is not None: await self._notification_listener.close() self._notification_listener = None if self._event_log is not None: await self._event_log.flush_events() self._opened = False
[docs] def get_event_log(self, config: "EventHistoryConfig") -> "AdvancedAlchemyQueueEventLog": """Return Advanced Alchemy-managed queue event history.""" if self._event_log is None: self._event_log = AdvancedAlchemyQueueEventLog( config=config, service_factory=self._event_log_service, transaction_factory=self._event_log_operation, runtime_logger=self._logger, ) return self._event_log
[docs] async def enqueue( self, task_name: "str", *, args: "tuple[Any, ...]" = (), kwargs: "dict[str, Any] | None" = None, queue: "str" = "default", priority: "int" = 0, max_retries: "int" = 0, scheduled_at: "datetime | None" = None, expires_at: "datetime | None" = None, key: "str | None" = None, execution_backend: "str" = "local", execution_profile: "str | None" = None, metadata: "dict[str, Any] | None" = None, id: "UUID | None" = None, # noqa: A002 ) -> "QueuedTaskRecord": try: async with self._operation() as service: record = await service.enqueue( task_name, args=args, kwargs=dict(kwargs or {}), queue=queue, priority=priority, max_retries=max_retries, scheduled_at=scheduled_at, expires_at=expires_at, key=key, execution_backend=execution_backend, execution_profile=execution_profile, metadata=dict(metadata or {}), id=id, ) except (AdvancedAlchemyIntegrityError, SQLAlchemyIntegrityError): if key is None: raise async with self._service() as service: existing = await service.get_task_by_key(key) if existing is None: raise record = existing await self.notify_new_task(record) return record
[docs] async def enqueue_many(self, requests: "Sequence[TaskRequest]") -> "list[QueuedTaskRecord]": """Persist multiple queued tasks in one Advanced Alchemy operation. Returns: Queue task records in input order. """ if not requests: return [] async with self._operation() as service: records = await service.enqueue_many(requests) self._increment_queue_metric("enqueue", float(len(records))) await self.notify_new_tasks(records) self._record_enqueue_batch(len(requests)) return records
[docs] async def get_task(self, task_id: "UUID") -> "QueuedTaskRecord | None": async with self._service() as service: return await service.get_task(task_id)
[docs] async def get_task_by_key(self, key: "str") -> "QueuedTaskRecord | None": async with self._service() as service: return await service.get_task_by_key(key)
[docs] async def list_pending( self, *, limit: "int" = 1, queue: "str | None" = None, execution_backend: "str | None" = None ) -> "list[QueuedTaskRecord]": async with self._service() as service: return await service.list_pending(limit=limit, queue=queue, execution_backend=execution_backend)
[docs] async def time_until_next_due(self, *, queues: "tuple[str, ...]" = ()) -> "float | None": """Return seconds until the earliest not-yet-due pending/scheduled record. Returns: Seconds until the next due record, or ``None`` when there is no upcoming scheduled work. """ async with self._service() as service: next_at = await service.next_scheduled_at(queues=queues) if next_at is None: return None return max((next_at - _utc_now()).total_seconds(), 0.0)
[docs] async def claim_task( self, task_id: "UUID", *, expected_retry_count: "int | None" = None, expected_execution_ref: "str | None" = None ) -> "QueuedTaskRecord | None": async with self._operation() as service: return await service.claim_task( task_id, expected_retry_count=expected_retry_count, expected_execution_ref=expected_execution_ref )
[docs] async def claim_task_with_expired( self, task_id: "UUID", *, expected_retry_count: "int | None" = None, expected_execution_ref: "str | None" = None ) -> "tuple[QueuedTaskRecord | None, QueuedTaskRecord | None]": async with self._operation() as service: return await service.claim_task_with_expired( task_id, expected_retry_count=expected_retry_count, expected_execution_ref=expected_execution_ref )
[docs] async def claim_next( self, *, queues: "tuple[str, ...]" = (), execution_backend: "str | None" = None ) -> "QueuedTaskRecord | None": async with self._operation() as service: for queue in queues or (None,): claimed = await service.claim_next(queue=queue, execution_backend=execution_backend) if claimed is not None: return claimed return None
[docs] async def claim_many( self, *, limit: "int", queues: "tuple[str, ...]" = (), execution_backend: "str | None" = None, queue_limits: "Mapping[str, int] | None" = None, ) -> "list[QueuedTaskRecord]": """Claim up to ``limit`` due tasks across the requested queues. Returns: Claimed task records. """ if queue_limits is not None: return await super().claim_many( limit=limit, queues=queues, execution_backend=execution_backend, queue_limits=queue_limits ) if limit <= 0: return [] records: "list[QueuedTaskRecord]" = [] async with self._operation() as service: for queue in queues or (None,): if len(records) >= limit: break remaining = limit - len(records) claimed_records = await service.claim_many( limit=remaining, queue=queue, execution_backend=execution_backend ) records.extend(claimed_records) self._increment_queue_metric("claim", float(len(records))) return records
[docs] async def claim_many_with_expired( self, *, limit: "int", queues: "tuple[str, ...]" = (), execution_backend: "str | None" = None, queue_limits: "Mapping[str, int] | None" = None, ) -> "tuple[list[QueuedTaskRecord], list[QueuedTaskRecord]]": """Claim records and report expiry transitions from the same transactions.""" if queue_limits is not None: return await super().claim_many_with_expired( limit=limit, queues=queues, execution_backend=execution_backend, queue_limits=queue_limits ) if limit <= 0: return [], [] records: "list[QueuedTaskRecord]" = [] expired: "list[QueuedTaskRecord]" = [] async with self._operation() as service: for queue in queues or (None,): if len(records) >= limit: break claimed_records, expired_records = await service.claim_many_with_expired( limit=limit - len(records), queue=queue, execution_backend=execution_backend ) records.extend(claimed_records) expired.extend(expired_records) self._increment_queue_metric("claim", float(len(records))) unique_expired = {record.id: record for record in expired} return records, list(unique_expired.values())
[docs] async def complete_task( self, task_id: "UUID", *, result: "Any" = None, expected_retry_count: "int | None" = None ) -> "QueuedTaskRecord | None": async with self._operation() as service: return await service.complete_task(task_id, result=result, expected_retry_count=expected_retry_count)
[docs] async def fail_task( self, task_id: "UUID", error: "str", *, retry: "bool" = True, expected_retry_count: "int | None" = None, retry_at: "datetime | None" = None, queued_at: "datetime | None" = None, ) -> "QueuedTaskRecord | None": async with self._operation() as service: return await service.fail_task( task_id, error, retry=retry, expected_retry_count=expected_retry_count, retry_at=retry_at, queued_at=queued_at, )
[docs] async def assign_worker( self, task_id: "UUID", *, worker_id: "str", expected_retry_count: "int" ) -> "QueuedTaskRecord | None": async with self._operation() as service: return await service.assign_worker(task_id, worker_id=worker_id, expected_retry_count=expected_retry_count)
[docs] async def interrupt_task( self, task_id: "UUID", *, expected_retry_count: "int", worker_id: "str", queued_at: "datetime" ) -> "QueuedTaskRecord | None": async with self._operation() as service: return await service.interrupt_task( task_id, expected_retry_count=expected_retry_count, worker_id=worker_id, queued_at=queued_at )
[docs] async def cancel_task( self, task_id: "UUID", *, include_running: "bool" = False, expected_retry_count: "int | None" = None ) -> "bool": async with self._operation() as service: return await service.cancel_task( task_id, include_running=include_running, expected_retry_count=expected_retry_count )
[docs] async def cancel_tasks( self, *, task_name: "str | None" = None, queue: "str | None" = None, kwargs: "Mapping[str, Any] | None" = None, metadata: "Mapping[str, Any] | None" = None, include_running: "bool" = False, ) -> "int": async with self._operation() as service: return await service.cancel_tasks( task_name=task_name, queue=queue, kwargs=kwargs, metadata=metadata, include_running=include_running )
[docs] async def touch_heartbeats(self, touches: "Sequence[HeartbeatTouch]") -> "HeartbeatTouchResult": if not touches: return HeartbeatTouchResult() async with self._heartbeat_operation() as service: return await service.touch_heartbeats(touches)
[docs] async def null_heartbeats(self, task_ids: "list[UUID]", *, expected_retry_count: "int | None" = None) -> "None": async with self._heartbeat_operation() as service: await service.null_heartbeats(task_ids, expected_retry_count=expected_retry_count)
[docs] async def requeue_stale_running( self, *, stale_after: "timedelta", limit: "int | None" = None ) -> "StaleTaskRecoveryResult": async with self._operation() as service: return await service.requeue_stale_running( stale_after=stale_after, limit=limit, priority_policy=self._stale_requeue_priority_policy() )
[docs] async def set_execution_ref( self, task_id: "UUID", execution_backend: "str", execution_ref: "str", *, execution_profile: "str | None" = None ) -> "QueuedTaskRecord | None": async with self._operation() as service: return await service.set_execution_ref( task_id, execution_backend, execution_ref, execution_profile=execution_profile )
[docs] async def reserve_external_dispatch( self, task_id: "UUID", execution_backend: "str", reservation_ref: "str", *, execution_profile: "str | None" = None, expected_retry_count: "int | None" = None, ) -> "QueuedTaskRecord | None": async with self._operation() as service: return await service.reserve_external_dispatch( task_id, execution_backend, reservation_ref, execution_profile=execution_profile, expected_retry_count=expected_retry_count, )
[docs] async def clear_execution_ref( self, task_id: "UUID", expected_retry_count: "int", expected_execution_ref: "str" ) -> "QueuedTaskRecord | None": async with self._operation() as service: record = await service.clear_execution_ref(task_id, expected_retry_count, expected_execution_ref) if record is not None: await self.notify_new_task(record) return record
[docs] async def replace_execution_ref( self, task_id: "UUID", expected_retry_count: "int", expected_execution_ref: "str", execution_ref: "str" ) -> "QueuedTaskRecord | None": async with self._operation() as service: return await service.replace_execution_ref( task_id, expected_retry_count, expected_execution_ref, execution_ref )
[docs] async def release_external_dispatch( self, task_id: "UUID", reservation_ref: "str", execution_backend: "str", *, execution_profile: "str | None" = None, ) -> "QueuedTaskRecord | None": async with self._operation() as service: record = await service.release_external_dispatch( task_id, reservation_ref, execution_backend, execution_profile=execution_profile ) if record is not None: await self.notify_new_task(record) return record
[docs] async def finalize_external_dispatch( self, task_id: "UUID", reservation_ref: "str", execution_backend: "str", execution_ref: "str", *, execution_profile: "str | None" = None, ) -> "QueuedTaskRecord | None": async with self._operation() as service: return await service.finalize_external_dispatch( task_id, reservation_ref, execution_backend, execution_ref, execution_profile=execution_profile )
[docs] async def set_execution_backend( self, task_id: "UUID", execution_backend: "str", *, execution_profile: "str | None" = None ) -> "QueuedTaskRecord | None": async with self._operation() as service: record = await service.set_execution_backend( task_id, execution_backend, execution_profile=execution_profile ) if record is not None: await self.notify_new_task(record) return record
[docs] async def notify_new_task(self, record: "QueuedTaskRecord") -> "None": """Publish a PostgreSQL worker wakeup marker when enabled.""" if not self._notifications_supported() or record.status not in {"pending", "scheduled"} or not record.is_due: return await self._send_notification_marker() self._increment_queue_metric("notify") self._record_wakeup_emitted()
[docs] async def notify_new_tasks(self, records: "Sequence[QueuedTaskRecord]") -> "None": """Coalesce a batch of task records into at most one wakeup marker.""" due = tuple(record for record in records if record.status in {"pending", "scheduled"} and record.is_due) if due: await self.notify_new_task(due[0]) self._record_wakeup_coalesced(len(due) - 1)
[docs] async def wait_for_wakeups(self, timeout: "float | None" = None) -> "bool": """Wait for a PostgreSQL worker wakeup marker when configured. Returns: True when a wakeup marker or due-row reconciliation is observed. """ if not self._notifications_supported(): return await super().wait_for_wakeups(timeout=timeout) listener = self._get_notification_listener() await listener.start() if await self._has_due_tasks(): self._increment_queue_metric("poll_fallback") return True wait_timeout = self._event_poll_interval if self._event_poll_interval is not None else timeout notified = await listener.wait(wait_timeout) if notified: self._increment_queue_metric("listener_wakeup") return bool(notified)
[docs] async def list_running_external(self, *, limit: "int | None" = None) -> "list[QueuedTaskRecord]": async with self._service() as service: return await service.list_running_external(limit=limit)
[docs] async def get_statistics(self, *, queue: "str | None" = None) -> "QueueStatistics": async with self._service() as service: return await service.get_statistics(queue=queue)
[docs] async def expire_overdue(self, *, limit: "int | None" = None) -> "list[QueuedTaskRecord]": async with self._operation() as service: return await service.expire_overdue(limit=limit)
[docs] async def list_completed_by_task( self, task_name: "str", *, since: "datetime | None" = None, limit: "int" = 10 ) -> "list[QueuedTaskRecord]": async with self._service() as service: return await service.list_completed_by_task(task_name, since=since, limit=limit)
[docs] async def cleanup_terminal(self, before: "datetime", *, limit: "int | None" = None) -> "int": async with self._operation() as service: return await service.cleanup_terminal(before, limit=limit)
[docs] async def acquire_maintenance(self, name: "str", token: "str", *, ttl: "timedelta") -> "bool": """Acquire maintenance ownership via a portable compare-and-set. Updates an existing expired row for ``name`` to this token; if no row was expired, inserts a fresh one inside a savepoint so a uniqueness race (live ownership held elsewhere) is treated as a denied acquisition. Returns: True when maintenance ownership is held under ``token``. """ model = cast("Any", self._maintenance_model_class) now = _utc_now() new_expiry = now + ttl async with self._session() as session, session.begin(): result = cast( "Any", await session.execute( update(model) .where(model.name == name, model.expires_at <= now) .values(token=token, expires_at=new_expiry) .execution_options(synchronize_session=False) ), ) if result.rowcount == 1: return True try: async with session.begin_nested(): session.add(model(name=name, token=token, expires_at=new_expiry)) await session.flush() except SQLAlchemyIntegrityError: return False return True
[docs] async def release_maintenance(self, name: "str", token: "str") -> "bool": """Release maintenance ownership only when ``token`` matches the holder. Returns: True when ownership held under ``token`` was deleted. """ model = cast("Any", self._maintenance_model_class) async with self._session() as session, session.begin(): result = cast( "Any", await session.execute( delete(model) .where(model.name == name, model.token == token) .execution_options(synchronize_session=False) ), ) return bool(result.rowcount == 1)
def _ensure_configured(self) -> "None": if self._sqlalchemy_config is None: msg = "SQLAlchemyBackend requires sqlalchemy_config." raise QueueConfigurationError(msg) def _ensure_opened(self) -> "None": if not self._opened: msg = "SQLAlchemyBackend.open() must be called before using the backend." raise RuntimeError(msg) def _driver_name(self) -> "str | None": sqlalchemy_config = self._sqlalchemy_config if sqlalchemy_config is None or sqlalchemy_config.connection_string is None: return None return make_url(sqlalchemy_config.connection_string).drivername def _notifications_supported(self) -> "bool": return self._notifications and self._driver_name() in SUPPORTED_NOTIFY_DRIVERS def _get_notification_listener(self) -> "NotificationListener": if self._notification_listener is None: self._notification_listener = self._create_notification_listener() return self._notification_listener def _create_notification_listener(self) -> "NotificationListener": sqlalchemy_config = self._sqlalchemy_config if sqlalchemy_config is None or sqlalchemy_config.connection_string is None: msg = "SQLAlchemyBackend requires sqlalchemy_config for PostgreSQL notifications." raise QueueConfigurationError(msg) return create_notification_listener( connection_string=sqlalchemy_config.connection_string, channel=self._wakeup_channel ) async def _send_notification_marker(self) -> "None": sqlalchemy_config = self._sqlalchemy_config if sqlalchemy_config is None: msg = "SQLAlchemyBackend requires sqlalchemy_config for PostgreSQL notifications." raise QueueConfigurationError(msg) engine = sqlalchemy_config.get_engine() async with engine.begin() as connection: await connection.execute( text("SELECT pg_notify(:channel, :payload)"), {"channel": self._wakeup_channel, "payload": _POSTGRES_NOTIFY_PAYLOAD}, ) async def _has_due_tasks(self) -> "bool": async with self._service() as service: return bool(await service.list_pending(limit=1, queue=None, execution_backend=None)) def _increment_queue_metric(self, name: "str", amount: "float" = 1.0) -> "None": if amount == 0 or self.config is None or self.config.observability is None: return if self._observability_runtime is None: self._observability_runtime = create_observability_runtime( self.config.observability, namespace=self.config.names ) self._observability_runtime.record_counter( f"litestar_queues.queue.{name}", int(amount), attributes={"messaging.system": self.config.names.root, "backend": "advanced-alchemy"}, ) def _resolve_model_classes( self, model_class: "type[object] | None" ) -> 'tuple[type[QueueTaskModelMixin], type["QueueTaskService"]]': if model_class is None: msg = "SQLAlchemyBackendConfig.model_class must inherit QueueTaskModelMixin." raise QueueConfigurationError(msg) try: valid_model = issubclass(model_class, QueueTaskModelMixin) except TypeError: valid_model = False if not valid_model: msg = "SQLAlchemyBackendConfig.model_class must inherit QueueTaskModelMixin." raise QueueConfigurationError(msg) if "__tablename__" not in model_class.__dict__: msg = "SQLAlchemyBackendConfig.model_class must declare __tablename__." raise QueueConfigurationError(msg) typed_model = cast("type[QueueTaskModelMixin]", model_class) mapper = cast("Any", sqlalchemy_inspect(typed_model)) missing_columns = { "id", "created_at", "queued_at", "task_name", "args_json", "kwargs_json", "queue", "execution_backend", "execution_profile", "execution_ref", "worker_id", "status", "priority", "max_retries", "retry_count", "scheduled_at", "started_at", "completed_at", "heartbeat_at", "result_json", "error", "task_key", "metadata_json", } - {property_.key for property_ in mapper.column_attrs} if missing_columns: columns = ", ".join(sorted(missing_columns)) msg = f"SQLAlchemyBackendConfig.model_class is missing queue columns: {columns}." raise QueueConfigurationError(msg) return typed_model, QueueTaskService.for_model(typed_model) def _resolve_event_history_model_classes( self, model_class: "type[object] | None" ) -> 'tuple[type[QueueEventHistoryModelMixin], type["QueueEventLogService"]]': if model_class is None: msg = "SQLAlchemyBackendConfig.event_history_model_class must inherit QueueEventHistoryModelMixin." raise QueueConfigurationError(msg) try: valid_model = issubclass(model_class, QueueEventHistoryModelMixin) except TypeError: valid_model = False if not valid_model: msg = "SQLAlchemyBackendConfig.event_history_model_class must inherit QueueEventHistoryModelMixin." raise QueueConfigurationError(msg) if "__tablename__" not in model_class.__dict__: msg = "SQLAlchemyBackendConfig.event_history_model_class must declare __tablename__." raise QueueConfigurationError(msg) typed_model = cast("type[QueueEventHistoryModelMixin]", model_class) mapper = cast("Any", sqlalchemy_inspect(typed_model)) missing_columns = { "created_at", "event_id", "event_type", "task_id", "task_name", "queue", "worker_id", "execution_backend", "execution_profile", "level", "message", "detail_json", "progress_current", "progress_total", "progress_percent", "sequence", "occurred_at", } - {property_.key for property_ in mapper.column_attrs} if missing_columns: columns = ", ".join(sorted(missing_columns)) msg = f"SQLAlchemyBackendConfig.event_history_model_class is missing event-log columns: {columns}." raise QueueConfigurationError(msg) return typed_model, QueueEventLogService.for_model(typed_model) def _resolve_maintenance_model_class( self, model_class: "type[object] | None" ) -> "type[QueueMaintenanceModelMixin]": if model_class is None: msg = "SQLAlchemyBackendConfig.maintenance_model_class must inherit QueueMaintenanceModelMixin." raise QueueConfigurationError(msg) try: valid_model = issubclass(model_class, QueueMaintenanceModelMixin) except TypeError: valid_model = False if not valid_model: msg = "SQLAlchemyBackendConfig.maintenance_model_class must inherit QueueMaintenanceModelMixin." raise QueueConfigurationError(msg) if "__tablename__" not in model_class.__dict__: msg = "SQLAlchemyBackendConfig.maintenance_model_class must declare __tablename__." raise QueueConfigurationError(msg) typed_model = cast("type[QueueMaintenanceModelMixin]", model_class) mapper = cast("Any", sqlalchemy_inspect(typed_model)) missing_columns = {"name", "token", "expires_at"} - {property_.key for property_ in mapper.column_attrs} if missing_columns: columns = ", ".join(sorted(missing_columns)) msg = f"SQLAlchemyBackendConfig.maintenance_model_class is missing coordination columns: {columns}." raise QueueConfigurationError(msg) return typed_model def _resolve_task_reservation_model_classes( self, model_class: "type[object] | None" ) -> 'tuple[type[QueueTaskReservationModelMixin], type["QueueTaskReservationService"]]': if model_class is None: msg = "SQLAlchemyBackendConfig.task_reservation_model_class must inherit QueueTaskReservationModelMixin." raise QueueConfigurationError(msg) try: valid_model = issubclass(model_class, QueueTaskReservationModelMixin) except TypeError: valid_model = False if not valid_model: msg = "SQLAlchemyBackendConfig.task_reservation_model_class must inherit QueueTaskReservationModelMixin." raise QueueConfigurationError(msg) if "__tablename__" not in model_class.__dict__: msg = "SQLAlchemyBackendConfig.task_reservation_model_class must declare __tablename__." raise QueueConfigurationError(msg) typed_model = cast("type[QueueTaskReservationModelMixin]", model_class) mapper = cast("Any", sqlalchemy_inspect(typed_model)) missing_columns = {"id", "created_at", "identity_key", "task_id", "task_name"} - { property_.key for property_ in mapper.column_attrs } if missing_columns: columns = ", ".join(sorted(missing_columns)) msg = f"SQLAlchemyBackendConfig.task_reservation_model_class is missing reservation columns: {columns}." raise QueueConfigurationError(msg) return typed_model, QueueTaskReservationService.for_model(typed_model)
[docs] async def reserve_identity(self, key: "str", *, task_id: "UUID", task_name: "str") -> "TaskReservation | None": """Reserve a forever identity via select-then-insert with an integrity fallback. The reservation table's unique ``identity_key`` column is the atomicity arbiter: a losing concurrent insert surfaces an integrity error and the loser re-reads the winning owner. The reservation table is separate from the task table and terminal cleanup never touches it. Returns: ``None`` when this caller won the reservation; otherwise the existing owner reservation. """ try: async with self._task_reservation_operation() as service: existing = await service.reserve(key, task_id=task_id, task_name=task_name) if existing is not None: return self._reservation_from_model(existing) except (AdvancedAlchemyIntegrityError, SQLAlchemyIntegrityError): owner = await self.has_identity(key) if owner is not None: return owner raise else: return None
[docs] async def has_identity(self, key: "str") -> "TaskReservation | None": """Return the reservation owning a reserved forever identity, if any.""" async with self._task_reservation_service() as service: model = await service.get_owner(key) return self._reservation_from_model(model) if model is not None else None
[docs] async def reset_identity(self, key: "str", *, expected_task_id: "UUID | None" = None) -> "bool": """Delete a forever identity reservation via atomic compare-and-delete. Args: key: The exact effective identity key. expected_task_id: Optional task owner required for deletion. Returns: ``True`` when a reservation was removed. """ async with self._task_reservation_operation() as service: return await service.delete_by_key(key, expected_task_id=expected_task_id)
def _reservation_from_model(self, model: "Any") -> "TaskReservation": created_at = model.created_at if created_at.tzinfo is None: created_at = created_at.replace(tzinfo=timezone.utc) return TaskReservation( key=str(model.identity_key), task_id=UUID(str(model.task_id)), task_name=str(model.task_name), created_at=created_at.astimezone(timezone.utc), ) def _event_history_enabled(self) -> "bool": events_config = self.config.events if self.config is not None else None return events_config is not None and events_config.history is not None @asynccontextmanager async def _session(self) -> "AsyncIterator[AsyncSession]": self._ensure_configured() sqlalchemy_config = self._sqlalchemy_config if sqlalchemy_config is None: msg = "SQLAlchemyBackend requires sqlalchemy_config." raise QueueConfigurationError(msg) session_maker = sqlalchemy_config.create_session_maker() async with session_maker() as session: yield session @asynccontextmanager async def _service(self) -> 'AsyncIterator["QueueTaskService"]': self._ensure_opened() async with self._session() as session: yield self._service_class(session=session) @asynccontextmanager async def _operation(self) -> 'AsyncIterator["QueueTaskService"]': self._ensure_opened() async with self._session() as session, session.begin(): yield self._service_class(session=session) @asynccontextmanager async def _event_log_service(self) -> 'AsyncIterator["QueueEventLogService"]': self._ensure_opened() async with self._session() as session: yield self._event_log_service_class(session=session) @asynccontextmanager async def _event_log_operation(self) -> 'AsyncIterator["QueueEventLogService"]': self._ensure_opened() async with self._session() as session, session.begin(): yield self._event_log_service_class(session=session) @asynccontextmanager async def _task_reservation_service(self) -> 'AsyncIterator["QueueTaskReservationService"]': self._ensure_opened() async with self._session() as session: yield self._task_reservation_service_class(session=session) @asynccontextmanager async def _task_reservation_operation(self) -> 'AsyncIterator["QueueTaskReservationService"]': self._ensure_opened() async with self._session() as session, session.begin(): yield self._task_reservation_service_class(session=session) @asynccontextmanager async def _heartbeat_operation(self) -> 'AsyncIterator["QueueTaskService"]': """Yield a ``QueueTaskService`` bound to the dedicated heartbeat session maker. Falls back to :meth:`_operation` when ``heartbeat_session_maker`` is not configured. The dedicated engine is supplied and owned by the adopter; :meth:`close` does not dispose it. Yields: Queue task service bound to the heartbeat or default operation. """ self._ensure_opened() if self._heartbeat_session_maker is None: async with self._operation() as service: yield service else: async with self._heartbeat_session_maker() as session, session.begin(): yield self._service_class(session=session)