From de9a62f6323f01687f04852e0232e12c48440b11 Mon Sep 17 00:00:00 2001 From: Anto Subash Date: Thu, 1 Oct 2026 13:33:53 +0200 Subject: [PATCH 1/2] feat(background_tasks): stamp executions with the publishing tenant and retry as it (#371) TaskExecution gets a nullable, indexed tenant_id (NULL = beat / platform publish) stamped in on_task_publish from the message's tenant header or the publisher's current tenant, plus a (tenant_id, status, queued_at) index. Not MultiTenantMixin: beat publishes have no tenant. Retries re-publish under the original row's tenant (all_tenants() for a platform row) and copy tenant_id, so an operator who is a member of another org cannot trip the stamp mismatch. Admin screens stay platform-wide with a tenant column and filter. Claude-Session: https://claude.ai/code/session_01F8RiTBUJQnZmSq56qReZeV --- ...17_background_tasks_execution_tenant_id.py | 43 ++++ .../background_tasks/constants.py | 4 + .../background_tasks/contracts/schemas.py | 2 + .../background_tasks/endpoints/api_admin.py | 15 +- .../background_tasks/endpoints/views.py | 15 +- .../background_tasks/filters.py | 17 +- .../background_tasks/locales/en.json | 12 +- .../background_tasks/models.py | 13 ++ .../background_tasks/pages/Index.tsx | 32 ++- .../pages/components/DetailFacts.tsx | 4 + .../pages/components/ExecutionRow.tsx | 3 + .../pages/components/ExecutionsTable.tsx | 7 +- .../pages/components/TaskFilters.tsx | 43 +++- .../background_tasks/pages/constants.ts | 4 + .../background_tasks/pages/retry.ts | 6 +- .../background_tasks/retry_service.py | 19 +- .../background_tasks/service.py | 38 +++- .../background_tasks/signals.py | 32 +-- .../background_tasks/tenant_context.py | 10 + modules/background_tasks/tests/conftest.py | 2 + .../tests/test_execution_tenancy.py | 194 ++++++++++++++++++ packages/i18n/src/generated-resources.ts | 6 + packages/i18n/src/keys.generated.ts | 6 + 23 files changed, 486 insertions(+), 41 deletions(-) create mode 100644 host/migrations/versions/b5d3f08a6e17_background_tasks_execution_tenant_id.py create mode 100644 modules/background_tasks/tests/test_execution_tenancy.py diff --git a/host/migrations/versions/b5d3f08a6e17_background_tasks_execution_tenant_id.py b/host/migrations/versions/b5d3f08a6e17_background_tasks_execution_tenant_id.py new file mode 100644 index 00000000..d218e4dc --- /dev/null +++ b/host/migrations/versions/b5d3f08a6e17_background_tasks_execution_tenant_id.py @@ -0,0 +1,43 @@ +"""background_tasks_task_execution: add nullable tenant_id + +NULL means a beat / platform publish (no tenant bound), which is also what +every pre-existing row is. The table is deliberately not ``MultiTenantMixin``: +those rows have no tenant and strict mode would raise on their insert. See +GH #371. + +background_tasks has no migration chain of its own (its tables came in with +the host's initial schema), so this extends the mainline head. + +Revision ID: b5d3f08a6e17 +Revises: d4e8a1b6c392 +Create Date: 2026-10-01 13:00:00.000000 +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision: str = "b5d3f08a6e17" +down_revision: str | None = "d4e8a1b6c392" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +_TABLE = "background_tasks_task_execution" +_SINGLE = "ix_background_tasks_task_execution_tenant_id" +_COMPOSITE = "ix_background_tasks_task_execution_tenant_status_queued" + + +def upgrade() -> None: + with op.batch_alter_table(_TABLE) as batch: + batch.add_column(sa.Column("tenant_id", sa.String(length=50), nullable=True)) + batch.create_index(_SINGLE, ["tenant_id"], unique=False) + batch.create_index(_COMPOSITE, ["tenant_id", "status", "queued_at"], unique=False) + + +def downgrade() -> None: + with op.batch_alter_table(_TABLE) as batch: + batch.drop_index(_COMPOSITE) + batch.drop_index(_SINGLE) + batch.drop_column("tenant_id") diff --git a/modules/background_tasks/background_tasks/constants.py b/modules/background_tasks/background_tasks/constants.py index d79fb600..df8799d6 100644 --- a/modules/background_tasks/background_tasks/constants.py +++ b/modules/background_tasks/background_tasks/constants.py @@ -102,3 +102,7 @@ class TaskStatus(StrEnum): # Harmless round-trip task; only wired into CI smoke tests and local dev — # not scheduled, not invoked by other modules. DEMO_ECHO_TASK = "background_tasks.demo_echo" + +TENANT_ID_MAX_LENGTH = 50 +# Filter value meaning "executions with no tenant" (beat / platform publishes). +PLATFORM_TENANT_FILTER = "__platform__" diff --git a/modules/background_tasks/background_tasks/contracts/schemas.py b/modules/background_tasks/background_tasks/contracts/schemas.py index d845a5e4..abbb8169 100644 --- a/modules/background_tasks/background_tasks/contracts/schemas.py +++ b/modules/background_tasks/background_tasks/contracts/schemas.py @@ -28,6 +28,7 @@ class TaskExecutionListItem(SQLModel): task_name: str status: TaskStatus queue: str + tenant_id: str | None = None args: list[Any] = [] kwargs: dict[str, Any] = {} retries: int @@ -49,6 +50,7 @@ class TaskExecutionDetail(SQLModel): task_name: str status: TaskStatus queue: str + tenant_id: str | None = None args: list[Any] = [] kwargs: dict[str, Any] = {} result: dict[str, Any] | None = None diff --git a/modules/background_tasks/background_tasks/endpoints/api_admin.py b/modules/background_tasks/background_tasks/endpoints/api_admin.py index 2997e202..472a1a09 100644 --- a/modules/background_tasks/background_tasks/endpoints/api_admin.py +++ b/modules/background_tasks/background_tasks/endpoints/api_admin.py @@ -37,6 +37,7 @@ async def list_executions( q: str | None = Query(default=None), task_name: str | None = Query(default=None, deprecated=True), queue: str | None = Query(default=None), + tenant_id: str | None = Query(default=None), page: int = Query(default=1, ge=1), per_page: int = Query(default=20, ge=1, le=200), service: BackgroundTaskService = Depends(get_background_task_service), @@ -50,7 +51,12 @@ async def list_executions( existing API callers keep working, and ``q`` wins when both are given. """ return await service.list( - status=status, task_name=q or task_name, queue=queue, page=page, per_page=per_page + status=status, + task_name=q or task_name, + queue=queue, + tenant_id=tenant_id, + page=page, + per_page=per_page, ) @@ -74,15 +80,18 @@ async def retry_failed_executions( status: TaskStatus | None = Query(default=None), task_name: str | None = Query(default=None, alias="q"), queue: str | None = Query(default=None), + tenant_id: str | None = Query(default=None), service: BackgroundTaskService = Depends(get_background_task_service), ) -> RetryFailedResult: """Re-enqueue every failed or stuck execution the current filter can see. - Takes the same three filters as the listing — status, search, queue — so + Takes the same filters as the listing — status, search, queue, tenant — so the sweep covers exactly the rows the operator is looking at and nothing else. Capped per call; the response says how many eligible rows are left. """ - return await service.retry_failed(status=status, task_name=task_name, queue=queue) + return await service.retry_failed( + status=status, task_name=task_name, queue=queue, tenant_id=tenant_id + ) @router.post( diff --git a/modules/background_tasks/background_tasks/endpoints/views.py b/modules/background_tasks/background_tasks/endpoints/views.py index 255ea199..47e019f7 100644 --- a/modules/background_tasks/background_tasks/endpoints/views.py +++ b/modules/background_tasks/background_tasks/endpoints/views.py @@ -12,6 +12,7 @@ from background_tasks.constants import ( PERM_VIEW, + PLATFORM_TENANT_FILTER, TaskStatus, ) from background_tasks.deps import get_background_task_service @@ -53,6 +54,7 @@ async def index( status: TaskStatus | None = Query(default=None), task_name: str = Query(default="", alias="q"), queue: str = Query(default=""), + tenant_id: str = Query(default="", alias="tenant"), page: int = Query(default=1), service: BackgroundTaskService = Depends(get_background_task_service), ) -> InertiaResponse: @@ -60,22 +62,28 @@ async def index( status=status, task_name=task_name or None, queue=queue or None, + tenant_id=tenant_id or None, page=page, per_page=PER_PAGE, ) # A filtered-empty list says nothing about the fleet, so don't pay to poll # it: the screen already knows to blame the filter. - unfiltered_and_empty = response.total == 0 and not task_name and not queue and status is None + unfiltered_and_empty = ( + response.total == 0 and not task_name and not queue and not tenant_id and status is None + ) async def strip_data() -> tuple[dict[str, int], list[str]]: """Everything above the table. Serial: one session, one connection.""" - counts = await service.status_counts(task_name=task_name or None, queue=queue or None) + counts = await service.status_counts( + task_name=task_name or None, queue=queue or None, tenant_id=tenant_id or None + ) # The strip's third tile is a windowed throughput reading rather than a # status total, so it needs its own count — see `success_count_since`. counts["success_24h"] = await service.success_count_since( hours=SUCCESS_WINDOW_HOURS, task_name=task_name or None, queue=queue or None, + tenant_id=tenant_id or None, ) # Deliberately unfiltered: the dropdown is how an operator leaves the # queue they are in, so it keeps offering the ones they are not in. @@ -97,6 +105,8 @@ async def strip_data() -> tuple[dict[str, int], list[str]]: "success_24h": counts["success_24h"], }, "queues": queues, + "tenant_ids": await service.tenant_ids(), + "platform_tenant_value": PLATFORM_TENANT_FILTER, "worker_presence": presence, "pagination": { "page": response.page, @@ -107,6 +117,7 @@ async def strip_data() -> tuple[dict[str, int], list[str]]: "status": status.value if status else "", "task_name": task_name, "queue": queue, + "tenant": tenant_id, }, }, ) diff --git a/modules/background_tasks/background_tasks/filters.py b/modules/background_tasks/background_tasks/filters.py index b41e90aa..68435ebd 100644 --- a/modules/background_tasks/background_tasks/filters.py +++ b/modules/background_tasks/background_tasks/filters.py @@ -11,7 +11,7 @@ from simple_module_db import LIKE_ESCAPE_CHAR, like_contains_pattern from sqlalchemy import ColumnElement, select -from background_tasks.constants import RETRYABLE_STATUSES, TaskStatus +from background_tasks.constants import PLATFORM_TENANT_FILTER, RETRYABLE_STATUSES, TaskStatus from background_tasks.models import TaskExecution Conditions = list[ColumnElement[bool]] @@ -22,9 +22,19 @@ def execution_filters( status: TaskStatus | None = None, task_name: str | None = None, queue: str | None = None, + tenant_id: str | None = None, ) -> Conditions: - """The three axes the executions screen filters on.""" + """The axes the executions screen filters on. + + *tenant_id* is a tenant id, or :data:`PLATFORM_TENANT_FILTER` for + executions published with no tenant. The screen is platform-wide; tenant is + a filter, not a scope. + """ conditions: Conditions = [] + if tenant_id == PLATFORM_TENANT_FILTER: + conditions.append(TaskExecution.tenant_id.is_(None)) + elif tenant_id: + conditions.append(TaskExecution.tenant_id == tenant_id) if status is not None: conditions.append(TaskExecution.status == status) if task_name: @@ -41,6 +51,7 @@ def bulk_retry_conditions( status: TaskStatus | None = None, task_name: str | None = None, queue: str | None = None, + tenant_id: str | None = None, ) -> Conditions | None: """What the bulk sweep may touch, or ``None`` when it may touch nothing. @@ -64,6 +75,6 @@ def bulk_retry_conditions( already_retried = select(child.c.id).where(child.c.retried_from_id == TaskExecution.id).exists() return [ TaskExecution.status.in_(sorted(wanted)), - *execution_filters(task_name=task_name, queue=queue), + *execution_filters(task_name=task_name, queue=queue, tenant_id=tenant_id), ~already_retried, ] diff --git a/modules/background_tasks/background_tasks/locales/en.json b/modules/background_tasks/background_tasks/locales/en.json index 2a95bc52..4d75a9c9 100644 --- a/modules/background_tasks/background_tasks/locales/en.json +++ b/modules/background_tasks/background_tasks/locales/en.json @@ -26,7 +26,11 @@ "stuck": "Stuck", "queue_label": "Queue:", "queue_all": "all", - "queue_aria": "Filter by queue" + "queue_aria": "Filter by queue", + "tenant_label": "Tenant:", + "tenant_all": "all", + "tenant_platform": "Platform", + "tenant_aria": "Filter by tenant" }, "status": { "queued": "queued", @@ -44,7 +48,8 @@ "queued_at": "Queued", "duration": "Duration", "actions": "Actions", - "retry": "Retry" + "retry": "Retry", + "tenant": "Tenant" }, "detail": { "head_title": "Task detail", @@ -69,7 +74,8 @@ "traceback": "Traceback", "copy": "Copy", "copied": "Copied", - "no_traceback": "No traceback recorded." + "no_traceback": "No traceback recorded.", + "tenant": "Tenant" }, "retry_dialog": { "title": "Retry {name}?", diff --git a/modules/background_tasks/background_tasks/models.py b/modules/background_tasks/background_tasks/models.py index 295fa1ce..5d6f8827 100644 --- a/modules/background_tasks/background_tasks/models.py +++ b/modules/background_tasks/background_tasks/models.py @@ -15,6 +15,7 @@ DEFAULT_QUEUE, MODULE_NAME, TABLE_TASK_EXECUTION, + TENANT_ID_MAX_LENGTH, TaskStatus, ) @@ -48,6 +49,12 @@ class TaskExecution(Base, AuditMixin, table=True): # ty: ignore[unsupported-bas ) queue: str = Field(default=DEFAULT_QUEUE, max_length=64) + # The tenant the task was published for; NULL for beat and other platform + # publishes. Deliberately not MultiTenantMixin: those rows have no tenant + # and strict mode would raise on their insert. The admin screens are + # platform-wide and filter on this column. + tenant_id: str | None = Field(default=None, index=True, max_length=TENANT_ID_MAX_LENGTH) + args: list[Any] = Field(default_factory=list, sa_column=Column(JSON)) kwargs: dict[str, Any] = Field(default_factory=dict, sa_column=Column(JSON)) result: dict[str, Any] | None = Field(default=None, sa_column=Column(JSON)) @@ -79,4 +86,10 @@ class TaskExecution(Base, AuditMixin, table=True): # ty: ignore[unsupported-bas "status", "queued_at", ), + Index( + "ix_background_tasks_task_execution_tenant_status_queued", + "tenant_id", + "status", + "queued_at", + ), ) diff --git a/modules/background_tasks/background_tasks/pages/Index.tsx b/modules/background_tasks/background_tasks/pages/Index.tsx index b94061d9..edfc47ac 100644 --- a/modules/background_tasks/background_tasks/pages/Index.tsx +++ b/modules/background_tasks/background_tasks/pages/Index.tsx @@ -12,7 +12,7 @@ import { type StatusCounts, StatusStrip } from './components/StatusStrip'; import { TaskFilters } from './components/TaskFilters'; import { TasksEmptyRow, type WorkerPresence } from './components/TasksEmpty'; import { WorkerHealthBanner } from './components/WorkerHealthBanner'; -import { QUEUE_ALL, STATUS_ALL, TASK_STATUS, VIEW_BASE } from './constants'; +import { QUEUE_ALL, STATUS_ALL, TASK_STATUS, TENANT_ALL, VIEW_BASE } from './constants'; import { type Execution, retryAllFailed, retryExecution } from './retry'; interface Pagination { @@ -25,6 +25,7 @@ interface Filters { status: string; task_name: string; queue: string; + tenant: string; } interface Props { @@ -34,6 +35,10 @@ interface Props { status_counts: StatusCounts; /** Every queue that has run work, for the dropdown. */ queues: string[]; + /** Tenants that have published work, for the dropdown. */ + tenant_ids: string[]; + /** Filter value selecting executions with no tenant. */ + platform_tenant_value: string; /** Null unless the unfiltered list came back empty — see the index view. */ worker_presence: WorkerPresence | null; } @@ -47,6 +52,7 @@ function pushFilters(filters: Filters, page: number): void { if (filters.task_name) params.q = filters.task_name; if (filters.status && filters.status !== STATUS_ALL) params.status = filters.status; if (filters.queue && filters.queue !== QUEUE_ALL) params.queue = filters.queue; + if (filters.tenant && filters.tenant !== TENANT_ALL) params.tenant = filters.tenant; if (page > 1) params.page = String(page); router.get(VIEW_BASE, params, { preserveState: true, preserveScroll: true }); } @@ -58,6 +64,8 @@ function Index() { filters: initialFilters, status_counts: statusCounts, queues, + tenant_ids: tenantIds, + platform_tenant_value: platformTenantValue, worker_presence: workerPresence, } = usePage<{ props: Props }>().props as unknown as Props; @@ -71,16 +79,21 @@ function Index() { const [busy, setBusy] = useState(false); const statusValue = initialFilters.status || STATUS_ALL; const queueValue = initialFilters.queue || QUEUE_ALL; + const tenantValue = initialFilters.tenant || TENANT_ALL; // Derived from the server-confirmed filters, not the live `search` input: // `workerPresence` reflects the last committed request, so mixing it with // unsubmitted local state would flash the wrong empty-state copy during the // debounce window between a keystroke and the resulting navigation. const isFiltered = - !!(initialFilters.task_name ?? '') || statusValue !== STATUS_ALL || queueValue !== QUEUE_ALL; + !!(initialFilters.task_name ?? '') || + statusValue !== STATUS_ALL || + queueValue !== QUEUE_ALL || + tenantValue !== TENANT_ALL; const filters: Filters = { status: statusValue, task_name: search, queue: queueValue, + tenant: tenantValue, }; // Work that is supposed to be moving. Counting `running` too is deliberate: @@ -101,7 +114,7 @@ function Index() { // clear, and an armed flag would swallow the user's next keystroke. if (search !== '') skipNextDebounceRef.current = true; setSearch(''); - pushFilters({ status: STATUS_ALL, task_name: '', queue: QUEUE_ALL }, 1); + pushFilters({ status: STATUS_ALL, task_name: '', queue: QUEUE_ALL, tenant: TENANT_ALL }, 1); } // Debounce search: any change from the server-provided value kicks off a @@ -113,11 +126,15 @@ function Index() { } if (search === (initialFilters.task_name ?? '')) return; const timeout = setTimeout( - () => pushFilters({ status: statusValue, task_name: search, queue: queueValue }, 1), + () => + pushFilters( + { status: statusValue, task_name: search, queue: queueValue, tenant: tenantValue }, + 1, + ), 300, ); return () => clearTimeout(timeout); - }, [search, initialFilters.task_name, statusValue, queueValue]); + }, [search, initialFilters.task_name, statusValue, queueValue, tenantValue]); async function handleRetry() { if (!retryTarget) return; @@ -134,6 +151,7 @@ function Index() { status: statusValue, taskName: initialFilters.task_name ?? '', queue: queueValue, + tenant: tenantValue, }); setBusy(false); setRetryAllOpen(false); @@ -188,6 +206,10 @@ function Index() { queue={queueValue} onQueueChange={(queue) => pushFilters({ ...filters, queue }, 1)} queues={queues ?? []} + tenant={tenantValue} + onTenantChange={(tenant) => pushFilters({ ...filters, tenant }, 1)} + tenantIds={tenantIds ?? []} + platformTenantValue={platformTenantValue} /> {execution.queue} + + {execution.tenant_id ?? t(keys.background_tasks.filters.tenant_platform)} + {ago(execution.queued_at)} diff --git a/modules/background_tasks/background_tasks/pages/components/ExecutionsTable.tsx b/modules/background_tasks/background_tasks/pages/components/ExecutionsTable.tsx index 72d11a5b..7583cd83 100644 --- a/modules/background_tasks/background_tasks/pages/components/ExecutionsTable.tsx +++ b/modules/background_tasks/background_tasks/pages/components/ExecutionsTable.tsx @@ -12,8 +12,8 @@ import type { ReactNode } from 'react'; import type { Execution } from '../retry'; import { ExecutionRow } from './ExecutionRow'; -/** Task, Status, Queue, Queued, Duration, Actions. */ -export const COLUMN_COUNT = 6; +/** Task, Status, Queue, Tenant, Queued, Duration, Actions. */ +export const COLUMN_COUNT = 7; // Same header treatment as the other admin tables (users, audit log, flags). const TH = 'text-[11px] font-semibold uppercase tracking-[0.08em] text-muted-foreground'; @@ -63,6 +63,9 @@ export function ExecutionsTable({ + diff --git a/modules/background_tasks/background_tasks/pages/components/TaskFilters.tsx b/modules/background_tasks/background_tasks/pages/components/TaskFilters.tsx index 2e6bacc8..0dc251c7 100644 --- a/modules/background_tasks/background_tasks/pages/components/TaskFilters.tsx +++ b/modules/background_tasks/background_tasks/pages/components/TaskFilters.tsx @@ -9,7 +9,13 @@ import { SelectValue, } from '@simple-module-py/ui/components/ui/select'; import { Search } from 'lucide-react'; -import { QUEUE_ALL, SEGMENT_LABEL_KEY, SEGMENT_STATUSES, STATUS_ALL } from '../constants'; +import { + QUEUE_ALL, + SEGMENT_LABEL_KEY, + SEGMENT_STATUSES, + STATUS_ALL, + TENANT_ALL, +} from '../constants'; interface Props { search: string; @@ -22,6 +28,13 @@ interface Props { onQueueChange: (next: string) => void; /** Every queue that has run work, for the dropdown. */ queues: string[]; + /** Active tenant filter, or `TENANT_ALL`. */ + tenant: string; + onTenantChange: (next: string) => void; + /** Tenants that have published work, for the dropdown. */ + tenantIds: string[]; + /** Filter value selecting executions with no tenant (platform publishes). */ + platformTenantValue: string; } /** @@ -40,6 +53,10 @@ export function TaskFilters({ queue, onQueueChange, queues, + tenant, + onTenantChange, + tenantIds, + platformTenantValue, }: Props) { const { t } = useT(); const options = SEGMENT_STATUSES.map((value) => ({ @@ -87,6 +104,30 @@ export function TaskFilters({ ))} + ); diff --git a/modules/background_tasks/background_tasks/pages/constants.ts b/modules/background_tasks/background_tasks/pages/constants.ts index 86dd42da..48c766d3 100644 --- a/modules/background_tasks/background_tasks/pages/constants.ts +++ b/modules/background_tasks/background_tasks/pages/constants.ts @@ -95,6 +95,7 @@ export interface TaskDetail { task_name: string; status: TaskStatus; queue: string; + tenant_id: string | null; args: unknown[]; kwargs: Record; result: Record | null; @@ -192,6 +193,9 @@ export const STATUS_ALL = 'all'; */ export const QUEUE_ALL = '__all__'; +/** Sentinel for "no tenant filter" in the tenant dropdown. */ +export const TENANT_ALL = '__all__'; + /** * The four the segmented control offers: everything, plus the three states an * operator triages. `pending`, `success`, `revoked` and `retrying` stay diff --git a/modules/background_tasks/background_tasks/pages/retry.ts b/modules/background_tasks/background_tasks/pages/retry.ts index 22047723..ce4bbe14 100644 --- a/modules/background_tasks/background_tasks/pages/retry.ts +++ b/modules/background_tasks/background_tasks/pages/retry.ts @@ -4,7 +4,7 @@ import { keys, t } from '@simple-module-py/i18n'; import { toast } from 'sonner'; -import { API_BASE, QUEUE_ALL, STATUS_ALL, type TaskStatus } from './constants'; +import { API_BASE, QUEUE_ALL, STATUS_ALL, type TaskStatus, TENANT_ALL } from './constants'; export interface Execution { id: string; @@ -12,6 +12,8 @@ export interface Execution { task_name: string; status: TaskStatus; queue: string; + /** Tenant that published the task; null for beat / platform publishes. */ + tenant_id: string | null; // Shown by the retry confirm before it re-enqueues them — see the list schema. args: unknown[]; kwargs: Record; @@ -63,11 +65,13 @@ export async function retryAllFailed(filters: { status: string; taskName: string; queue: string; + tenant: string; }): Promise { const params = new URLSearchParams(); if (filters.status && filters.status !== STATUS_ALL) params.set('status', filters.status); if (filters.taskName) params.set('q', filters.taskName); if (filters.queue && filters.queue !== QUEUE_ALL) params.set('queue', filters.queue); + if (filters.tenant && filters.tenant !== TENANT_ALL) params.set('tenant_id', filters.tenant); const query = params.toString(); try { const res = await fetch(`${API_BASE}/executions/retry-failed${query ? `?${query}` : ''}`, { diff --git a/modules/background_tasks/background_tasks/retry_service.py b/modules/background_tasks/background_tasks/retry_service.py index 0f029d2a..96321f03 100644 --- a/modules/background_tasks/background_tasks/retry_service.py +++ b/modules/background_tasks/background_tasks/retry_service.py @@ -14,7 +14,7 @@ from typing import TYPE_CHECKING, Any from simple_module_core.events import EventBus -from simple_module_db import RequestSession +from simple_module_db import RequestSession, all_tenants, tenant_context from sqlalchemy import Select, func, select from background_tasks.constants import RETRY_ALL_BATCH, TaskStatus @@ -95,6 +95,7 @@ async def retry_failed( status: TaskStatus | None = None, task_name: str | None = None, queue: str | None = None, + tenant_id: str | None = None, limit: int | None = None, ) -> RetryFailedResult: """Re-enqueue the retryable executions the current view can see. @@ -120,7 +121,9 @@ async def retry_failed( # Read at call time, not bound as a default argument, so the cap has # exactly one source of truth that tests and operators can move. batch = RETRY_ALL_BATCH if limit is None else limit - conditions = bulk_retry_conditions(status=status, task_name=task_name, queue=queue) + conditions = bulk_retry_conditions( + status=status, task_name=task_name, queue=queue, tenant_id=tenant_id + ) if conditions is None: return RetryFailedResult(queued=0, remaining=0) @@ -179,7 +182,16 @@ def _publish(self, row: TaskExecution, task_id: str | None = None) -> str: } if task_id is not None: options["task_id"] = task_id - return self.celery.send_task(row.task_name, **options).id + # Publish as the row's own tenant, never the operator's: a platform + # admin who is also a member of some org must not re-run another + # tenant's job stamped with theirs (``stamp_tenant`` would refuse the + # mismatch). A row with no tenant is a platform publish, so it goes out + # with none bound. + if row.tenant_id: + with tenant_context(row.tenant_id): + return self.celery.send_task(row.task_name, **options).id + with all_tenants(): + return self.celery.send_task(row.task_name, **options).id def _new_attempt(self, row: TaskExecution, celery_task_id: str) -> TaskExecution: """The row recording a fresh attempt at *row*. Not yet flushed.""" @@ -188,6 +200,7 @@ def _new_attempt(self, row: TaskExecution, celery_task_id: str) -> TaskExecution task_name=row.task_name, status=TaskStatus.PENDING, queue=row.queue, + tenant_id=row.tenant_id, args=list(row.args or []), kwargs=dict(row.kwargs or {}), retried_from_id=row.id, diff --git a/modules/background_tasks/background_tasks/service.py b/modules/background_tasks/background_tasks/service.py index 9f60fcc7..a0946ed7 100644 --- a/modules/background_tasks/background_tasks/service.py +++ b/modules/background_tasks/background_tasks/service.py @@ -58,6 +58,7 @@ async def list( status: TaskStatus | None = None, task_name: str | None = None, queue: str | None = None, + tenant_id: str | None = None, page: int = 1, per_page: int = 20, ) -> TaskExecutionListResponse: @@ -65,7 +66,9 @@ async def list( page = max(page, 1) per_page = max(1, min(per_page, 200)) - conditions = execution_filters(status=status, task_name=task_name, queue=queue) + conditions = execution_filters( + status=status, task_name=task_name, queue=queue, tenant_id=tenant_id + ) # A window-function total lets us fetch the page and the count in a # single round trip. SQLAlchemy's ``AsyncSession`` serialises calls on @@ -106,10 +109,15 @@ async def fetch(p: int) -> tuple[Sequence[Row[tuple[TaskExecution, int]]], int]: status=status, task_name=task_name, queue=queue, + tenant_id=tenant_id, ) async def status_counts( - self, *, task_name: str | None = None, queue: str | None = None + self, + *, + task_name: str | None = None, + queue: str | None = None, + tenant_id: str | None = None, ) -> dict[str, int]: """Count executions per status for the ops strip above the table. @@ -122,7 +130,7 @@ async def status_counts( """ query = ( select(TaskExecution.status, func.count().label("n")) - .where(*execution_filters(task_name=task_name, queue=queue)) + .where(*execution_filters(task_name=task_name, queue=queue, tenant_id=tenant_id)) .group_by(TaskExecution.status) ) @@ -137,6 +145,7 @@ async def success_count_since( hours: int = 24, task_name: str | None = None, queue: str | None = None, + tenant_id: str | None = None, ) -> int: """Successes that *finished* inside the window. @@ -152,7 +161,12 @@ async def success_count_since( select(func.count()) .select_from(TaskExecution) .where( - *execution_filters(status=TaskStatus.SUCCESS, task_name=task_name, queue=queue), + *execution_filters( + status=TaskStatus.SUCCESS, + task_name=task_name, + queue=queue, + tenant_id=tenant_id, + ), TaskExecution.finished_at.is_not(None), TaskExecution.finished_at >= cutoff, ) @@ -170,6 +184,19 @@ async def queues(self) -> QueueNames: query = select(TaskExecution.queue).distinct().order_by(TaskExecution.queue) return [q for q in (await self.db.execute(query)).scalars() if q] + async def tenant_ids(self) -> QueueNames: + """Every tenant that has published work, for the tenant dropdown. + + Platform (NULL) executions are offered as a fixed option by the screen. + """ + query = ( + select(TaskExecution.tenant_id) + .where(TaskExecution.tenant_id.is_not(None)) + .distinct() + .order_by(TaskExecution.tenant_id) + ) + return [t for t in (await self.db.execute(query)).scalars() if t] + async def get(self, execution_id: uuid.UUID) -> TaskExecutionDetail | None: row = await self.db.get(TaskExecution, execution_id) if row is None: @@ -204,9 +231,10 @@ async def retry_failed( status: TaskStatus | None = None, task_name: str | None = None, queue: str | None = None, + tenant_id: str | None = None, limit: int | None = None, ) -> RetryFailedResult: """Re-enqueue the retryable executions the current view can see.""" return await self.retries.retry_failed( - status=status, task_name=task_name, queue=queue, limit=limit + status=status, task_name=task_name, queue=queue, tenant_id=tenant_id, limit=limit ) diff --git a/modules/background_tasks/background_tasks/signals.py b/modules/background_tasks/background_tasks/signals.py index 72c2e094..5bdfaac3 100644 --- a/modules/background_tasks/background_tasks/signals.py +++ b/modules/background_tasks/background_tasks/signals.py @@ -41,7 +41,12 @@ from background_tasks.log_context import signal_task_finished, signal_task_started from background_tasks.models import TaskExecution from background_tasks.sync_db import sync_session -from background_tasks.tenant_context import release_tenant, restore_tenant, stamp_tenant +from background_tasks.tenant_context import ( + published_tenant, + release_tenant, + restore_tenant, + stamp_tenant, +) logger = logging.getLogger(__name__) @@ -138,18 +143,19 @@ def on_task_publish( args_in, kwargs_in = body[0], body[1] args, kwargs = coerce_args_kwargs(args_in, kwargs_in) - _apply( - "on_task_publish", - celery_task_id=task_id, - defaults={ - "task_name": task_name, - "status": TaskStatus.PENDING, - "queue": routing_key or DEFAULT_QUEUE, - "args": args, - "kwargs": kwargs, - "queued_at": now_utc(), - }, - ) + defaults: dict[str, Any] = { + "task_name": task_name, + "status": TaskStatus.PENDING, + "queue": routing_key or DEFAULT_QUEUE, + "args": args, + "kwargs": kwargs, + "queued_at": now_utc(), + } + # Omitted when none, so a platform publish never blanks a stamped row. + if tenant_id := published_tenant(headers): + defaults["tenant_id"] = tenant_id + + _apply("on_task_publish", celery_task_id=task_id, defaults=defaults) # ── Execution lifecycle ───────────────────────────────────────── diff --git a/modules/background_tasks/background_tasks/tenant_context.py b/modules/background_tasks/background_tasks/tenant_context.py index 3bc01d55..25137f61 100644 --- a/modules/background_tasks/background_tasks/tenant_context.py +++ b/modules/background_tasks/background_tasks/tenant_context.py @@ -55,6 +55,15 @@ def stamp_tenant(headers: dict[str, Any] | None) -> None: headers[TENANT_HEADER] = tenant_id +def published_tenant(headers: dict[str, Any] | None) -> str | None: + """The tenant an outgoing message is for: its header, else the publisher's. + + Call after :func:`stamp_tenant`. The header may also have been set by + platform code naming a tenant while none is bound. + """ + return (headers or {}).get(TENANT_HEADER) or current_tenant_id.get() + + def _tenant_of(task: Any) -> str | None: request = getattr(task, "request", None) if request is None: @@ -97,6 +106,7 @@ def release_tenant(*, task_id: str | None) -> None: __all__ = [ "TENANT_HEADER", + "published_tenant", "release_tenant", "restore_tenant", "set_default_tenant", diff --git a/modules/background_tasks/tests/conftest.py b/modules/background_tasks/tests/conftest.py index 65526503..fb38837a 100644 --- a/modules/background_tasks/tests/conftest.py +++ b/modules/background_tasks/tests/conftest.py @@ -70,6 +70,7 @@ async def _seed( status: TaskStatus = TaskStatus.FAILED, queue: str = "default", queued_at: datetime | None = None, + tenant_id: str | None = None, ) -> TaskExecution: row = TaskExecution( celery_task_id=str(uuid.uuid4()), @@ -79,6 +80,7 @@ async def _seed( args=[], kwargs={}, queued_at=queued_at or datetime.now(UTC), + tenant_id=tenant_id, ) async with app.state.sm.db.session_factory() as session: session.add(row) diff --git a/modules/background_tasks/tests/test_execution_tenancy.py b/modules/background_tasks/tests/test_execution_tenancy.py new file mode 100644 index 00000000..df527138 --- /dev/null +++ b/modules/background_tasks/tests/test_execution_tenancy.py @@ -0,0 +1,194 @@ +"""TaskExecution carries the tenant it was published for (#371). + +The table is not tenant-scoped (beat publishes have no tenant); the admin +screens are platform-wide and filter on the column. Retries re-publish as the +original row's tenant, not the operator's. +""" + +from __future__ import annotations + +import uuid +from pathlib import Path +from unittest.mock import MagicMock + +import httpx +import pytest +from background_tasks import sync_db +from background_tasks.constants import PLATFORM_TENANT_FILTER, TaskStatus +from background_tasks.models import TaskExecution +from background_tasks.retry_service import RetryCoordinator +from background_tasks.signals import on_task_publish +from background_tasks.tenant_context import TENANT_HEADER +from simple_module_core.tenancy import TenantRole +from simple_module_db import current_tenant_id, tenant_context +from sqlalchemy import select + +ADMIN = "/api/background_tasks/admin" + + +def _publish(headers: dict | None = None) -> str: + task_id = str(uuid.uuid4()) + on_task_publish( + sender="demo.echo", + headers={"id": task_id, "task": "demo.echo", **(headers or {})}, + body=([], {}, {}), + ) + return task_id + + +def _tenant_of(task_id: str) -> str | None: + with sync_db.get_sync_session_factory()() as session: + return session.execute( + select(TaskExecution.tenant_id).where(TaskExecution.celery_task_id == task_id) + ).scalar_one() + + +class TestPublishStamp: + def test_publish_under_a_tenant_stamps_it(self, sync_sqlite: Path): + with tenant_context("acme"): + task_id = _publish() + assert _tenant_of(task_id) == "acme" + + def test_beat_or_platform_publish_stays_null(self, sync_sqlite: Path): + assert _tenant_of(_publish()) is None + + def test_platform_code_may_name_the_tenant_in_the_header(self, sync_sqlite: Path): + assert _tenant_of(_publish({TENANT_HEADER: "globex"})) == "globex" + + def test_a_later_signal_does_not_blank_the_tenant(self, sync_sqlite: Path): + with tenant_context("acme"): + task_id = _publish() + # Republished with no tenant (e.g. the retry's send under no context). + on_task_publish( + sender="demo.echo", headers={"id": task_id, "task": "demo.echo"}, body=([], {}, {}) + ) + assert _tenant_of(task_id) == "acme" + + +def _coordinator() -> tuple[RetryCoordinator, list[str | None]]: + seen: list[str | None] = [] + celery = MagicMock(name="Celery") + + def send_task(*_a, **_k): + seen.append(current_tenant_id.get()) + return MagicMock(id="mocked-id") + + celery.send_task.side_effect = send_task + return RetryCoordinator(MagicMock(), celery, MagicMock()), seen + + +def _row(tenant_id: str | None) -> TaskExecution: + return TaskExecution(task_name="demo.echo", queue="default", tenant_id=tenant_id) + + +class TestRetryPublish: + def test_publishes_as_the_row_tenant_not_the_operators(self): + coordinator, seen = _coordinator() + with tenant_context("operators-org"): + coordinator._publish(_row("acme")) + assert seen == ["acme"] + + def test_a_platform_row_is_published_with_no_tenant(self): + coordinator, seen = _coordinator() + with tenant_context("operators-org"): + coordinator._publish(_row(None)) + assert seen == [None] + + def test_the_new_attempt_copies_the_tenant(self): + coordinator, _ = _coordinator() + assert coordinator._new_attempt(_row("acme"), "new-id").tenant_id == "acme" + assert coordinator._new_attempt(_row(None), "new-id").tenant_id is None + + +@pytest.mark.usefixtures("_stub_celery") +class TestRetryEndpoints: + @pytest.fixture + def seen(self, app) -> list[str | None]: + seen: list[str | None] = [] + + def send_task(*_a, **_k): + seen.append(current_tenant_id.get()) + return MagicMock(id=str(uuid.uuid4())) + + app.state.background_tasks.celery.send_task.side_effect = send_task + return seen + + async def test_single_retry_runs_as_the_original_tenant( + self, seed_execution, execution_rows, authenticated_client: httpx.AsyncClient, seen + ): + original = await seed_execution(status=TaskStatus.FAILED, tenant_id="acme") + + resp = await authenticated_client.post(f"{ADMIN}/executions/{original.id}/retry") + + assert resp.status_code == 200, resp.text + assert resp.json()["tenant_id"] == "acme" + assert seen == ["acme"] + assert sorted(r.tenant_id for r in await execution_rows()) == ["acme", "acme"] + + async def test_bulk_retry_runs_each_row_as_its_own_tenant( + self, seed_execution, execution_rows, authenticated_client: httpx.AsyncClient, seen + ): + await seed_execution(status=TaskStatus.FAILED, tenant_id="acme") + await seed_execution(status=TaskStatus.FAILED, tenant_id="globex") + await seed_execution(status=TaskStatus.FAILED) + + resp = await authenticated_client.post(f"{ADMIN}/executions/retry-failed") + + assert resp.json() == {"queued": 3, "remaining": 0} + assert sorted(seen, key=str) == [None, "acme", "globex"] + children = [r for r in await execution_rows() if r.retried_from_id] + assert sorted((r.tenant_id or "") for r in children) == ["", "acme", "globex"] + + async def test_bulk_retry_can_be_narrowed_to_a_tenant( + self, seed_execution, authenticated_client: httpx.AsyncClient, seen + ): + await seed_execution(status=TaskStatus.FAILED, tenant_id="acme") + await seed_execution(status=TaskStatus.FAILED, tenant_id="globex") + + resp = await authenticated_client.post( + f"{ADMIN}/executions/retry-failed", params={"tenant_id": "acme"} + ) + + assert resp.json() == {"queued": 1, "remaining": 0} + assert seen == ["acme"] + + +class TestTenantFilter: + async def test_list_filters_by_tenant_and_platform( + self, seed_execution, authenticated_client: httpx.AsyncClient + ): + await seed_execution(task_name="a", tenant_id="acme") + await seed_execution(task_name="g", tenant_id="globex") + await seed_execution(task_name="p") + + async def names(**params: str) -> set[str]: + resp = await authenticated_client.get(f"{ADMIN}/executions", params=params) + return {i["task_name"] for i in resp.json()["items"]} + + assert await names() == {"a", "g", "p"} + assert await names(tenant_id="acme") == {"a"} + assert await names(tenant_id=PLATFORM_TENANT_FILTER) == {"p"} + + async def test_index_view_exposes_tenants_and_echoes_the_filter( + self, seed_execution, authenticated_client: httpx.AsyncClient + ): + await seed_execution(task_name="a", tenant_id="acme") + await seed_execution(task_name="g", tenant_id="globex") + + resp = await authenticated_client.get( + "/admin/background-tasks/", + params={"tenant": "acme"}, + headers={"X-Inertia": "true", "Accept": "application/json"}, + ) + + props = resp.json()["props"] + assert props["tenant_ids"] == ["acme", "globex"] + assert props["filters"]["tenant"] == "acme" + assert [e["task_name"] for e in props["executions"]] == ["a"] + assert props["pagination"]["total"] == 1 + + +@pytest.mark.parametrize("role", list(TenantRole)) +async def test_tenant_members_cannot_read_executions(tenant_client, role): + async with tenant_client(role) as m: + assert (await m.client.get(f"{ADMIN}/executions")).status_code == 403 diff --git a/packages/i18n/src/generated-resources.ts b/packages/i18n/src/generated-resources.ts index 32067c57..c7f94503 100644 --- a/packages/i18n/src/generated-resources.ts +++ b/packages/i18n/src/generated-resources.ts @@ -76,6 +76,7 @@ export default { 'background_tasks.detail.retried_from': '', 'background_tasks.detail.retry_button': '', 'background_tasks.detail.started_at': '', + 'background_tasks.detail.tenant': '', 'background_tasks.detail.traceback': '', 'background_tasks.detail.worker': '', 'background_tasks.filters.all': '', @@ -86,6 +87,10 @@ export default { 'background_tasks.filters.running': '', 'background_tasks.filters.status_label': '', 'background_tasks.filters.stuck': '', + 'background_tasks.filters.tenant_all': '', + 'background_tasks.filters.tenant_aria': '', + 'background_tasks.filters.tenant_label': '', + 'background_tasks.filters.tenant_platform': '', 'background_tasks.index.description': '', 'background_tasks.index.next': '', 'background_tasks.index.previous': '', @@ -134,6 +139,7 @@ export default { 'background_tasks.table.retry': '', 'background_tasks.table.status': '', 'background_tasks.table.task': '', + 'background_tasks.table.tenant': '', 'background_tasks.tasks_empty.broker_unreachable_description': '', 'background_tasks.tasks_empty.broker_unreachable_title': '', 'background_tasks.tasks_empty.clear_filters': '', diff --git a/packages/i18n/src/keys.generated.ts b/packages/i18n/src/keys.generated.ts index 9ec68d82..2db30f53 100644 --- a/packages/i18n/src/keys.generated.ts +++ b/packages/i18n/src/keys.generated.ts @@ -101,6 +101,7 @@ export const keys = { retried_from: 'background_tasks.detail.retried_from', retry_button: 'background_tasks.detail.retry_button', started_at: 'background_tasks.detail.started_at', + tenant: 'background_tasks.detail.tenant', traceback: 'background_tasks.detail.traceback', worker: 'background_tasks.detail.worker', }, @@ -113,6 +114,10 @@ export const keys = { running: 'background_tasks.filters.running', status_label: 'background_tasks.filters.status_label', stuck: 'background_tasks.filters.stuck', + tenant_all: 'background_tasks.filters.tenant_all', + tenant_aria: 'background_tasks.filters.tenant_aria', + tenant_label: 'background_tasks.filters.tenant_label', + tenant_platform: 'background_tasks.filters.tenant_platform', }, index: { description: 'background_tasks.index.description', @@ -176,6 +181,7 @@ export const keys = { retry: 'background_tasks.table.retry', status: 'background_tasks.table.status', task: 'background_tasks.table.task', + tenant: 'background_tasks.table.tenant', }, tasks_empty: { broker_unreachable_description: 'background_tasks.tasks_empty.broker_unreachable_description', From d17da7f3d2c77242aaa6f3b3712cb0fd7520ad39 Mon Sep 17 00:00:00 2001 From: Anto Subash Date: Thu, 1 Oct 2026 15:19:29 +0200 Subject: [PATCH 2/2] fix(background_tasks): stamp the message tenant worker-side and honour default_tenant (review of #371) - prerun/postrun and the terminal signals pass the header's tenant into the upsert; never blank an existing stamp (platform messages stay NULL). - set_database_url takes default_tenant and builds EngineTenancy with it, so all_tenants() inserts in a worker land in the install's tenant (mirrors DatabaseState.default_tenant_id). Module startup and run_worker.py pass it. - Move the event-bus bridge to _signal_bus.py to keep signals.py under the 300-line cap (bind/unbind re-exported). Claude-Session: https://claude.ai/code/session_01F8RiTBUJQnZmSq56qReZeV --- docs/framework/multi-tenancy.md | 11 ++- .../background_tasks/_signal_bus.py | 61 ++++++++++++ .../background_tasks/module.py | 5 +- .../background_tasks/signals.py | 88 +++++++---------- .../background_tasks/sync_db.py | 26 ++++- .../background_tasks/tenant_context.py | 24 ++++- .../tests/test_signal_tenant_stamp.py | 99 +++++++++++++++++++ .../background_tasks/tests/test_signals.py | 4 +- .../tests/test_worker_tenancy.py | 51 +++++++++- scripts/run_worker.py | 7 +- 10 files changed, 305 insertions(+), 71 deletions(-) create mode 100644 modules/background_tasks/background_tasks/_signal_bus.py create mode 100644 modules/background_tasks/tests/test_signal_tenant_stamp.py diff --git a/docs/framework/multi-tenancy.md b/docs/framework/multi-tenancy.md index a4dc08fc..274e5e53 100644 --- a/docs/framework/multi-tenancy.md +++ b/docs/framework/multi-tenancy.md @@ -164,9 +164,16 @@ bound may name one. Beat tasks have no request: wrap cross-tenant work in The worker never builds the app, so `background_tasks.sync_db` attaches the same listeners to its own session class and reads `multi_tenant` from the host settings (`scripts/run_worker.py`): task bodies get the same fail-closed rules -as request code. A process that talks to the DB some other way must do the +as request code. It also passes `default_tenant`, so an unbound insert in a +task body (an `all_tenants()` block) lands in that tenant exactly as it would +in the web process. A process that talks to the DB some other way must do the same — `attach_session_listeners(MySession)` plus -`bind_engine_policy(engine, EngineTenancy(tenant_strict=...))`. +`bind_engine_policy(engine, EngineTenancy(tenant_strict=..., default_tenant_id=...))`. + +The worker-side signals (prerun, success, failure, retry, revoked) stamp the +`TaskExecution` row with the tenant on the message header themselves — the +publish signal may never have written the row — and a signal whose message +carries no tenant never blanks a row already stamped. ## Resolution diff --git a/modules/background_tasks/background_tasks/_signal_bus.py b/modules/background_tasks/background_tasks/_signal_bus.py new file mode 100644 index 00000000..fb8bd3e7 --- /dev/null +++ b/modules/background_tasks/background_tasks/_signal_bus.py @@ -0,0 +1,61 @@ +"""Bridge from Celery's sync signal thread to the app's async event bus. + +Split out of :mod:`.signals` to keep it under the 300-line cap; ``signals`` +re-exports :func:`bind_event_bus` / :func:`unbind_event_bus`. +""" + +from __future__ import annotations + +import asyncio +import logging +from typing import Any + +from simple_module_core.events import EventBus + +logger = logging.getLogger("background_tasks.signals") + + +_bus: EventBus | None = None +_loop: asyncio.AbstractEventLoop | None = None + + +def bind_event_bus(bus: EventBus, loop: asyncio.AbstractEventLoop) -> None: + """Bind an event bus + its running loop so signals can publish events. + + Signals fire on the Celery sync thread; `run_coroutine_threadsafe` + bridges back to ``loop`` so handlers run on the API event loop + regardless of which thread triggered the signal. + """ + global _bus, _loop + _bus = bus + _loop = loop + + +def unbind_event_bus() -> None: + """Drop the bound bus — called from ``on_shutdown`` so tests stay isolated.""" + global _bus, _loop + _bus = None + _loop = None + + +def publish_from_signal(event: Any) -> None: + """Dispatch ``event`` onto the bound bus without blocking the signal thread.""" + if _bus is None or _loop is None: + return + try: + future = asyncio.run_coroutine_threadsafe(_bus.publish(event), _loop) + except RuntimeError: + # Loop has stopped (shutdown race). The DB row is already written. + logger.debug("Event bus loop is not running; skipping %s", type(event).__name__) + return + # Surface subscriber exceptions — run_coroutine_threadsafe otherwise only + # logs them when the Future is GC'd, which happens far from the failure. + future.add_done_callback(_log_publish_failure) + + +def _log_publish_failure(future: asyncio.Future[Any]) -> None: + if future.cancelled(): + return + exc = future.exception() + if exc is not None: + logger.error("Event publish raised: %s", exc, exc_info=exc) diff --git a/modules/background_tasks/background_tasks/module.py b/modules/background_tasks/background_tasks/module.py index 4255a68e..11ef3cb2 100644 --- a/modules/background_tasks/background_tasks/module.py +++ b/modules/background_tasks/background_tasks/module.py @@ -161,8 +161,9 @@ async def on_startup(self, app: FastAPI) -> None: # SQLite default and silently drop ``TaskExecution`` rows. host = app.state.sm.settings multi = bool(getattr(host, "multi_tenant", False)) - set_database_url(host.database_url, tenant_strict=multi) - set_default_tenant(None if multi else getattr(host, "default_tenant", "") or None) + default_tenant = None if multi else getattr(host, "default_tenant", "") or None + set_database_url(host.database_url, tenant_strict=multi, default_tenant=default_tenant) + set_default_tenant(default_tenant) # build_celery imports `signals` for side effects and runs # `autodiscover_tasks` across every installed module. services.celery = build_celery(services.settings) diff --git a/modules/background_tasks/background_tasks/signals.py b/modules/background_tasks/background_tasks/signals.py index 5bdfaac3..bc1ae92f 100644 --- a/modules/background_tasks/background_tasks/signals.py +++ b/modules/background_tasks/background_tasks/signals.py @@ -19,14 +19,13 @@ from __future__ import annotations -import asyncio import logging from collections.abc import Callable from typing import Any from celery import signals -from simple_module_core.events import EventBus +from background_tasks._signal_bus import bind_event_bus, publish_from_signal, unbind_event_bus from background_tasks._signal_support import ( coerce_args_kwargs, jsonable_result, @@ -42,6 +41,7 @@ from background_tasks.models import TaskExecution from background_tasks.sync_db import sync_session from background_tasks.tenant_context import ( + message_tenant, published_tenant, release_tenant, restore_tenant, @@ -51,64 +51,24 @@ logger = logging.getLogger(__name__) -_bus: EventBus | None = None -_loop: asyncio.AbstractEventLoop | None = None - - -def bind_event_bus(bus: EventBus, loop: asyncio.AbstractEventLoop) -> None: - """Bind an event bus + its running loop so signals can publish events. - - Signals fire on the Celery sync thread; `run_coroutine_threadsafe` - bridges back to ``loop`` so handlers run on the API event loop - regardless of which thread triggered the signal. - """ - global _bus, _loop - _bus = bus - _loop = loop - - -def unbind_event_bus() -> None: - """Drop the bound bus — called from ``on_shutdown`` so tests stay isolated.""" - global _bus, _loop - _bus = None - _loop = None - - -def _publish_from_signal(event: Any) -> None: - """Dispatch ``event`` onto the bound bus without blocking the signal thread.""" - if _bus is None or _loop is None: - return - try: - future = asyncio.run_coroutine_threadsafe(_bus.publish(event), _loop) - except RuntimeError: - # Loop has stopped (shutdown race). The DB row is already written. - logger.debug("Event bus loop is not running; skipping %s", type(event).__name__) - return - # Surface subscriber exceptions — run_coroutine_threadsafe otherwise only - # logs them when the Future is GC'd, which happens far from the failure. - future.add_done_callback(_log_publish_failure) - - -def _log_publish_failure(future: asyncio.Future[Any]) -> None: - if future.cancelled(): - return - exc = future.exception() - if exc is not None: - logger.error("Event publish raised: %s", exc, exc_info=exc) - - def _apply( handler: str, *, celery_task_id: str | None, defaults: dict[str, Any], after: Callable[[TaskExecution], None] | None = None, + tenant_id: str | None = None, ) -> TaskExecution | None: """Open a sync session, upsert by celery_task_id, log on failure. Returns the upserted row (or ``None`` if the session raised) so callers can read DB-assigned fields like ``id`` without a second round trip. + + ``tenant_id`` (the message's) is written only when known, so a signal + whose message carries no tenant never blanks a row already stamped. """ + if tenant_id: + defaults = {**defaults, "tenant_id": tenant_id} try: with sync_session() as session: row = upsert_by_celery_id(session, celery_task_id=celery_task_id, defaults=defaults) @@ -151,11 +111,12 @@ def on_task_publish( "kwargs": kwargs, "queued_at": now_utc(), } - # Omitted when none, so a platform publish never blanks a stamped row. - if tenant_id := published_tenant(headers): - defaults["tenant_id"] = tenant_id - - _apply("on_task_publish", celery_task_id=task_id, defaults=defaults) + _apply( + "on_task_publish", + celery_task_id=task_id, + defaults=defaults, + tenant_id=published_tenant(headers), + ) # ── Execution lifecycle ───────────────────────────────────────── @@ -185,6 +146,7 @@ def on_task_prerun( "started_at": now, "heartbeat_at": now, }, + tenant_id=message_tenant(task=task), ) signal_task_started(task_id=task_id, task_name=name) restore_tenant(task_id=task_id, task=task) @@ -207,6 +169,7 @@ def on_task_postrun( "on_task_postrun", celery_task_id=task_id, defaults={"task_name": task_name_of(sender, task), "heartbeat_at": now_utc()}, + tenant_id=message_tenant(task=task), ) release_tenant(task_id=task_id) signal_task_finished(task_id=task_id) @@ -225,6 +188,7 @@ def on_task_success(sender: Any = None, result: Any = None, **_k: Any) -> None: "traceback": None, "exception_type": None, }, + tenant_id=message_tenant(task=sender), ) @@ -249,10 +213,11 @@ def on_task_failure( "exception_type": exception_type, "finished_at": now_utc(), }, + tenant_id=message_tenant(task=sender), ) if row is not None: - _publish_from_signal( + publish_from_signal( TaskFailed( task_execution_id=row.id, task_name=task_name, @@ -284,6 +249,7 @@ def _stamp_reason(row: TaskExecution) -> None: "heartbeat_at": now_utc(), }, after=_stamp_reason, + tenant_id=message_tenant(request=request), ) @@ -297,4 +263,18 @@ def on_task_revoked(sender: Any = None, request: Any = None, **_k: Any) -> None: "status": TaskStatus.REVOKED, "finished_at": now_utc(), }, + tenant_id=message_tenant(request=request), ) + + +__all__ = [ + "bind_event_bus", + "on_task_failure", + "on_task_postrun", + "on_task_prerun", + "on_task_publish", + "on_task_retry", + "on_task_revoked", + "on_task_success", + "unbind_event_bus", +] diff --git a/modules/background_tasks/background_tasks/sync_db.py b/modules/background_tasks/background_tasks/sync_db.py index 35cda138..f23cd4a8 100644 --- a/modules/background_tasks/background_tasks/sync_db.py +++ b/modules/background_tasks/background_tasks/sync_db.py @@ -18,6 +18,7 @@ from collections.abc import Iterator from contextlib import contextmanager +from simple_module_db import DEFAULT_TENANT_ID from simple_module_db.listeners import attach_session_listeners from simple_module_db.query_filter import EngineTenancy, bind_engine_policy from sqlalchemy import create_engine @@ -30,6 +31,7 @@ _session_factory: sessionmaker[Session] | None = None _url_override: str | None = None _tenant_strict: bool = False +_default_tenant: str | None = None class WorkerSession(Session): @@ -47,7 +49,9 @@ def _sync_url(async_url: str) -> str: return async_url.replace("+aiosqlite", "").replace("+asyncpg", "+psycopg2") -def set_database_url(url: str | None, *, tenant_strict: bool = False) -> None: +def set_database_url( + url: str | None, *, tenant_strict: bool = False, default_tenant: str | None = None +) -> None: """Pin the URL used to build the sync engine. The web process loads ``.env`` via pydantic-settings, but those values @@ -60,12 +64,20 @@ def set_database_url(url: str | None, *, tenant_strict: bool = False) -> None: ``tenant_strict`` mirrors the host's ``multi_tenant``: task bodies get the same fail-closed tenant rules as request code. + + ``default_tenant`` mirrors the host's ``default_tenant`` (ignored when + strict): an insert with no tenant bound — an ``all_tenants()`` block in a + task body — is stamped with it, as ``DatabaseState.default_tenant_id`` does + in the web process, instead of landing in ``DEFAULT_TENANT_ID`` where the + install's own scoped reads never see it. """ - global _url_override, _engine, _session_factory, _tenant_strict - if _url_override == url and _tenant_strict == tenant_strict: + global _url_override, _engine, _session_factory, _tenant_strict, _default_tenant + default_tenant = None if tenant_strict else (default_tenant or None) + if (_url_override, _tenant_strict, _default_tenant) == (url, tenant_strict, default_tenant): return _url_override = url _tenant_strict = tenant_strict + _default_tenant = default_tenant if _engine is not None: _engine.dispose() _engine = None @@ -86,7 +98,10 @@ def _build_engine() -> Engine: # no tenant filter at all: the tenant restored around a task body would # scope nothing (#371). attach_session_listeners(WorkerSession) - bind_engine_policy(engine, EngineTenancy(tenant_strict=_tenant_strict)) + policy = EngineTenancy( + tenant_strict=_tenant_strict, default_tenant_id=_default_tenant or DEFAULT_TENANT_ID + ) + bind_engine_policy(engine, policy) return engine @@ -106,13 +121,14 @@ def dispose_sync_engine() -> None: restarts within one process (test runners, uvicorn dev reload) don't accumulate engines against the old DB URL. """ - global _engine, _session_factory, _url_override, _tenant_strict + global _engine, _session_factory, _url_override, _tenant_strict, _default_tenant if _engine is not None: _engine.dispose() _engine = None _session_factory = None _url_override = None _tenant_strict = False + _default_tenant = None @contextmanager diff --git a/modules/background_tasks/background_tasks/tenant_context.py b/modules/background_tasks/background_tasks/tenant_context.py index 25137f61..e32db02c 100644 --- a/modules/background_tasks/background_tasks/tenant_context.py +++ b/modules/background_tasks/background_tasks/tenant_context.py @@ -65,12 +65,17 @@ def published_tenant(headers: dict[str, Any] | None) -> str | None: def _tenant_of(task: Any) -> str | None: - request = getattr(task, "request", None) + return _tenant_of_request(getattr(task, "request", None)) + + +def _tenant_of_request(request: Any) -> str | None: if request is None: return None value = getattr(request, TENANT_HEADER, None) - if value is None and isinstance(getattr(request, "headers", None), dict): - value = request.headers.get(TENANT_HEADER) + # A worker ``Request`` (revoked signal) keeps headers in a dict instead. + for attr in ("headers", "request_dict"): + if value is None and isinstance(getattr(request, attr, None), dict): + value = getattr(request, attr).get(TENANT_HEADER) if not value: return None if not is_valid_tenant_id(value): @@ -81,6 +86,18 @@ def _tenant_of(task: Any) -> str | None: return str(value) +def message_tenant(*, task: Any = None, request: Any = None) -> str | None: + """The tenant named on a running task's message header, if any. + + Signal handlers stamp it onto the ``TaskExecution`` row: the publish signal + may never have written the row (it fires in the publisher's process, which + can have a different DB, or none bound), so the worker-side upsert must + carry the tenant itself. The configured default tenant is deliberately not + returned — a platform message stays a platform row. + """ + return _tenant_of_request(request) if request is not None else _tenant_of(task) + + def restore_tenant(*, task_id: str | None, task: Any) -> None: """Enter the message's tenant for the task body (prerun).""" if not task_id: @@ -106,6 +123,7 @@ def release_tenant(*, task_id: str | None) -> None: __all__ = [ "TENANT_HEADER", + "message_tenant", "published_tenant", "release_tenant", "restore_tenant", diff --git a/modules/background_tasks/tests/test_signal_tenant_stamp.py b/modules/background_tasks/tests/test_signal_tenant_stamp.py new file mode 100644 index 00000000..757e3256 --- /dev/null +++ b/modules/background_tasks/tests/test_signal_tenant_stamp.py @@ -0,0 +1,99 @@ +"""Worker-side signals stamp the message's tenant themselves (#371). + +The publish signal fires in the publisher's process and may never have +written the row (another DB, or a publisher with no ``background_tasks`` +loaded), so prerun and the terminal handlers carry the header's tenant into +the upsert — and never blank one already stamped. +""" + +from __future__ import annotations + +import uuid +from pathlib import Path +from types import SimpleNamespace + +from background_tasks import sync_db +from background_tasks.models import TaskExecution +from background_tasks.signals import ( + on_task_failure, + on_task_postrun, + on_task_prerun, + on_task_publish, + on_task_retry, + on_task_revoked, + on_task_success, +) +from background_tasks.tenant_context import TENANT_HEADER +from simple_module_db import tenant_context +from sqlalchemy import select + + +def _tenant_of(task_id: str) -> str | None: + with sync_db.get_sync_session_factory()() as session: + return session.execute( + select(TaskExecution.tenant_id).where(TaskExecution.celery_task_id == task_id) + ).scalar_one() + + +def _task(tenant: str | None, task_id: str | None = None) -> SimpleNamespace: + headers = {TENANT_HEADER: tenant} if tenant else {} + return SimpleNamespace(name="demo.echo", request=SimpleNamespace(id=task_id, **headers)) + + +def _run(task_id: str, tenant: str | None) -> None: + task = _task(tenant, task_id) + on_task_prerun(sender=task, task_id=task_id, task=task) + on_task_postrun(sender=task, task_id=task_id, task=task) + + +def test_prerun_creates_the_row_with_the_header_tenant(sync_sqlite: Path): + task_id = str(uuid.uuid4()) + task = _task("acme", task_id) + on_task_prerun(sender=task, task_id=task_id, task=task) + try: + assert _tenant_of(task_id) == "acme" + finally: + on_task_postrun(sender=task, task_id=task_id, task=task) + + +def test_success_and_failure_stamp_a_row_they_create(sync_sqlite: Path): + ok_id, failed_id = str(uuid.uuid4()), str(uuid.uuid4()) + on_task_success(sender=_task("globex", ok_id), result=None) + on_task_failure(sender=_task("globex"), task_id=failed_id, exception=RuntimeError("x")) + assert _tenant_of(ok_id) == "globex" + assert _tenant_of(failed_id) == "globex" + + +def test_retry_and_revoke_read_the_request_header(sync_sqlite: Path): + retry_id, revoked_id = str(uuid.uuid4()), str(uuid.uuid4()) + on_task_retry( + sender=_task(None), + request=SimpleNamespace(id=retry_id, retries=1, **{TENANT_HEADER: "acme"}), + ) + # A worker ``Request`` keeps its headers in ``request_dict``. + on_task_revoked( + sender=_task(None), + request=SimpleNamespace(id=revoked_id, request_dict={TENANT_HEADER: "acme"}), + ) + assert _tenant_of(retry_id) == "acme" + assert _tenant_of(revoked_id) == "acme" + + +def test_a_headerless_signal_never_blanks_a_stamp(sync_sqlite: Path): + task_id = str(uuid.uuid4()) + with tenant_context("acme"): + on_task_publish( + sender="demo.echo", headers={"id": task_id, "task": "demo.echo"}, body=([], {}, {}) + ) + _run(task_id, None) + assert _tenant_of(task_id) == "acme" + + +def test_a_platform_message_stays_a_platform_row_on_a_default_tenant_install( + sync_sqlite: Path, +): + # The default tenant is what the body *binds*, not what the row records. + sync_db.set_database_url(sync_db._resolve_url(), default_tenant="main") + task_id = str(uuid.uuid4()) + _run(task_id, None) + assert _tenant_of(task_id) is None diff --git a/modules/background_tasks/tests/test_signals.py b/modules/background_tasks/tests/test_signals.py index 15972837..88dd92e3 100644 --- a/modules/background_tasks/tests/test_signals.py +++ b/modules/background_tasks/tests/test_signals.py @@ -13,7 +13,7 @@ from pathlib import Path import pytest -from background_tasks import signals as bg_signals +from background_tasks import _signal_bus as bg_signals_bus from background_tasks import sync_db from background_tasks._signal_support import upsert_by_celery_id from background_tasks.constants import TaskStatus @@ -238,7 +238,7 @@ class _Einfo: def test_no_op_when_bus_unbound(self, sync_sqlite: Path): """Without a bound bus (the standalone-worker case) the signal must still record the failure row and not raise.""" - assert bg_signals._bus is None + assert bg_signals_bus._bus is None task_id = str(uuid.uuid4()) on_task_publish( sender="demo.boom", diff --git a/modules/background_tasks/tests/test_worker_tenancy.py b/modules/background_tasks/tests/test_worker_tenancy.py index db9bfc54..58011950 100644 --- a/modules/background_tasks/tests/test_worker_tenancy.py +++ b/modules/background_tasks/tests/test_worker_tenancy.py @@ -7,7 +7,13 @@ import pytest from background_tasks import sync_db from background_tasks.tenant_context import TENANT_HEADER, release_tenant, restore_tenant -from simple_module_db import MultiTenantMixin, TenantIsolationError, create_module_base +from simple_module_db import ( + DEFAULT_TENANT_ID, + MultiTenantMixin, + TenantIsolationError, + all_tenants, + create_module_base, +) from sqlalchemy import create_engine, select from sqlmodel import Field @@ -56,3 +62,46 @@ def test_task_body_cannot_write_another_tenant(worker_db): def test_task_without_a_tenant_fails_closed(worker_db): with pytest.raises(TenantIsolationError), sync_db.sync_session() as s: s.execute(select(_Job)) + + +def _fresh(tmp_path, name: str, **policy) -> None: + url = f"sqlite:///{tmp_path}/{name}.db" + _Base.metadata.create_all(create_engine(url)) + sync_db.set_database_url(url, **policy) + + +def test_unbound_insert_lands_in_the_default_tenant(tmp_path): + """A ``default_tenant`` install's worker stamps an ``all_tenants()`` insert + with it — as ``DatabaseState.default_tenant_id`` does in the web process.""" + _fresh(tmp_path, "default", tenant_strict=False, default_tenant="main") + try: + with all_tenants(), sync_db.sync_session() as s: + s.add(_Job(name="swept")) + with sync_db.sync_session() as s: + assert s.scalars(select(_Job.tenant_id)).all() == ["main"] + finally: + sync_db.dispose_sync_engine() + + +def test_without_default_tenant_the_constant_is_used(tmp_path): + _fresh(tmp_path, "constant", tenant_strict=False) + try: + with all_tenants(), sync_db.sync_session() as s: + s.add(_Job(name="swept")) + with sync_db.sync_session() as s: + assert s.scalars(select(_Job.tenant_id)).all() == [DEFAULT_TENANT_ID] + finally: + sync_db.dispose_sync_engine() + + +def test_default_tenant_is_ignored_when_strict(tmp_path): + _fresh(tmp_path, "strict", tenant_strict=True, default_tenant="main") + try: + from simple_module_db.query_filter import _engine_policy + + sync_db.get_sync_session_factory() + policy = _engine_policy[sync_db._engine] + assert policy.tenant_strict + assert policy.default_tenant_id == DEFAULT_TENANT_ID + finally: + sync_db.dispose_sync_engine() diff --git a/scripts/run_worker.py b/scripts/run_worker.py index 7e4f58a2..d66b5bd7 100644 --- a/scripts/run_worker.py +++ b/scripts/run_worker.py @@ -41,8 +41,11 @@ # host/main.py does (env → DB → default): task bodies must hit the same # database with the same fail-closed tenant rules as request code. _host = merge_host_settings() -set_database_url(_host.database_url, tenant_strict=_host.multi_tenant) -set_default_tenant(None if _host.multi_tenant else _host.default_tenant or None) +_default_tenant = None if _host.multi_tenant else _host.default_tenant or None +set_database_url( + _host.database_url, tenant_strict=_host.multi_tenant, default_tenant=_default_tenant +) +set_default_tenant(_default_tenant) # Module-level name ``celery`` is what ``celery -A scripts.run_worker:celery`` # looks for. Keep it stable.