celery-dag 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (64) hide show
  1. celery_dag/__init__.py +8 -0
  2. celery_dag/api/__init__.py +15 -0
  3. celery_dag/api/app.py +88 -0
  4. celery_dag/api/routes/__init__.py +2 -0
  5. celery_dag/api/routes/dlq.py +58 -0
  6. celery_dag/api/routes/metrics.py +31 -0
  7. celery_dag/api/routes/tasks.py +19 -0
  8. celery_dag/api/routes/triggers.py +89 -0
  9. celery_dag/api/routes/workflows.py +336 -0
  10. celery_dag/api/schemas.py +132 -0
  11. celery_dag/celery_app/__init__.py +5 -0
  12. celery_dag/celery_app/app.py +54 -0
  13. celery_dag/celery_app/tasks.py +453 -0
  14. celery_dag/config/__init__.py +6 -0
  15. celery_dag/config/settings.py +178 -0
  16. celery_dag/core/__init__.py +2 -0
  17. celery_dag/core/branching.py +48 -0
  18. celery_dag/core/dag.py +273 -0
  19. celery_dag/core/exceptions.py +69 -0
  20. celery_dag/core/interfaces.py +39 -0
  21. celery_dag/engine/__init__.py +23 -0
  22. celery_dag/engine/cancellation.py +90 -0
  23. celery_dag/engine/dispatcher.py +200 -0
  24. celery_dag/engine/executor.py +104 -0
  25. celery_dag/engine/scheduler.py +11 -0
  26. celery_dag/engine/state_machine.py +59 -0
  27. celery_dag/engine/trigger_service.py +98 -0
  28. celery_dag/engine/workflow_control.py +56 -0
  29. celery_dag/example_tasks.py +80 -0
  30. celery_dag/middleware/__init__.py +2 -0
  31. celery_dag/middleware/cancellation_token.py +178 -0
  32. celery_dag/middleware/retry_policy.py +106 -0
  33. celery_dag/migrations/env.py +48 -0
  34. celery_dag/migrations/versions/0001_initial_schema.py +57 -0
  35. celery_dag/migrations/versions/0002_dispatch_outbox.py +48 -0
  36. celery_dag/migrations/versions/0003_cancellation_tokens.py +36 -0
  37. celery_dag/models/__init__.py +9 -0
  38. celery_dag/models/base.py +42 -0
  39. celery_dag/models/dispatch_outbox.py +58 -0
  40. celery_dag/models/enums.py +60 -0
  41. celery_dag/models/task_run.py +48 -0
  42. celery_dag/models/trigger.py +28 -0
  43. celery_dag/models/workflow.py +48 -0
  44. celery_dag/observability/__init__.py +6 -0
  45. celery_dag/observability/dag_visualizer.py +87 -0
  46. celery_dag/observability/logging.py +125 -0
  47. celery_dag/persistence/__init__.py +6 -0
  48. celery_dag/persistence/artifact_store.py +684 -0
  49. celery_dag/persistence/database.py +39 -0
  50. celery_dag/persistence/repositories/__init__.py +15 -0
  51. celery_dag/persistence/repositories/base.py +19 -0
  52. celery_dag/persistence/repositories/dispatch_outbox_repository.py +68 -0
  53. celery_dag/persistence/repositories/task_run_repository.py +57 -0
  54. celery_dag/persistence/repositories/trigger_repository.py +38 -0
  55. celery_dag/persistence/repositories/workflow_repository.py +96 -0
  56. celery_dag/persistence/unit_of_work.py +41 -0
  57. celery_dag/registry/__init__.py +6 -0
  58. celery_dag/registry/plugin_loader.py +23 -0
  59. celery_dag/registry/task_registry.py +60 -0
  60. celery_dag-1.0.0.dist-info/METADATA +297 -0
  61. celery_dag-1.0.0.dist-info/RECORD +64 -0
  62. celery_dag-1.0.0.dist-info/WHEEL +5 -0
  63. celery_dag-1.0.0.dist-info/licenses/LICENSE +177 -0
  64. celery_dag-1.0.0.dist-info/top_level.txt +1 -0
celery_dag/__init__.py ADDED
@@ -0,0 +1,8 @@
1
+ from celery_dag.core.dag import DAGBuilder
2
+ from celery_dag.engine.executor import DAGExecutor
3
+ from celery_dag.persistence.database import create_tables
4
+ from celery_dag.registry import task_hub
5
+
6
+ __version__ = "1.0.0"
7
+ __all__ = ["DAGBuilder", "DAGExecutor", "create_tables", "task_hub"]
8
+
@@ -0,0 +1,15 @@
1
+ """FastAPI REST API package for the DAG orchestrator."""
2
+
3
+ try:
4
+ import fastapi # noqa: F401
5
+ except ImportError as exc:
6
+ raise ImportError(
7
+ "The REST API components of celery_dag require FastAPI and Uvicorn. "
8
+ "Install them using: pip install 'celery-dag[api]'"
9
+ ) from exc
10
+
11
+ from celery_dag.api.app import create_api_app
12
+
13
+ __all__ = ["create_api_app"]
14
+
15
+
celery_dag/api/app.py ADDED
@@ -0,0 +1,88 @@
1
+ """FastAPI application setup."""
2
+ from __future__ import annotations
3
+
4
+ from contextlib import asynccontextmanager
5
+ from typing import AsyncGenerator
6
+
7
+ from fastapi import FastAPI
8
+ from fastapi.middleware.cors import CORSMiddleware
9
+
10
+ from redis import Redis
11
+
12
+ from celery_dag.api.routes.dlq import router as dlq_router
13
+ from celery_dag.api.routes.metrics import router as metrics_router
14
+ from celery_dag.api.routes.tasks import router as tasks_router
15
+ from celery_dag.api.routes.triggers import router as triggers_router
16
+ from celery_dag.api.routes.workflows import router as workflows_router
17
+ from celery_dag.config.settings import Settings, get_settings
18
+ from celery_dag.observability.logging import setup_logging
19
+ from celery_dag.persistence.unit_of_work import UnitOfWork
20
+
21
+
22
+ @asynccontextmanager
23
+ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
24
+ """Lifespan handler for FastAPI application startup and shutdown."""
25
+ setup_logging()
26
+ yield
27
+
28
+
29
+ def create_api_app(settings: Settings | None = None) -> FastAPI:
30
+ settings = settings or get_settings()
31
+ app = FastAPI(
32
+ title="Celery DAG Orchestrator API",
33
+ description="Production-grade resilient DAG Workflow Orchestrator REST API",
34
+ version="0.1.0",
35
+ lifespan=lifespan,
36
+ )
37
+
38
+ app.add_middleware(
39
+ CORSMiddleware,
40
+ allow_origins=["*"],
41
+ allow_credentials=True,
42
+ allow_methods=["*"],
43
+ allow_headers=["*"],
44
+ )
45
+
46
+ @app.get("/healthz", tags=["Health"], summary="Comprehensive service readiness & liveness probe")
47
+ def health_check() -> dict[str, Any]:
48
+ checks: dict[str, str] = {}
49
+ is_healthy = True
50
+
51
+ # Database Check
52
+ try:
53
+ with UnitOfWork() as uow:
54
+ uow.workflows.running_ids(limit=1)
55
+ checks["database"] = "ok"
56
+ except Exception as exc:
57
+ checks["database"] = f"error: {exc}"
58
+ is_healthy = False
59
+
60
+ # Redis Check
61
+ try:
62
+ r = Redis.from_url(settings.redis_url, socket_connect_timeout=1.0)
63
+ r.ping()
64
+ r.close()
65
+ checks["redis"] = "ok"
66
+ except Exception as exc:
67
+ checks["redis"] = f"error: {exc}"
68
+ is_healthy = False
69
+
70
+ status_str = "healthy" if is_healthy else "degraded"
71
+ return {
72
+ "status": status_str,
73
+ "checks": checks,
74
+ "environment": settings.environment,
75
+ "app_name": settings.app_name,
76
+ }
77
+
78
+ app.include_router(workflows_router, prefix="/api/v1")
79
+ app.include_router(dlq_router, prefix="/api/v1")
80
+ app.include_router(tasks_router, prefix="/api/v1")
81
+ app.include_router(triggers_router, prefix="/api/v1")
82
+ app.include_router(metrics_router, prefix="/api/v1")
83
+
84
+ return app
85
+
86
+
87
+ app = create_api_app()
88
+
@@ -0,0 +1,2 @@
1
+ """Routes package for REST API."""
2
+
@@ -0,0 +1,58 @@
1
+ """Dead-Letter Queue (DLQ) management endpoints."""
2
+ from __future__ import annotations
3
+
4
+ from uuid import UUID
5
+
6
+ from fastapi import APIRouter, HTTPException, Query, status
7
+
8
+ from celery_dag.api.schemas import DLQRecordResponse, ReplayDLQResponse
9
+ from celery_dag.engine.dispatcher import DAGDispatcher
10
+ from celery_dag.persistence.unit_of_work import UnitOfWork
11
+
12
+ router = APIRouter(prefix="/dlq", tags=["Dead-Letter Queue"])
13
+
14
+
15
+ @router.get(
16
+ "",
17
+ response_model=list[DLQRecordResponse],
18
+ summary="List dead-lettered outbox records",
19
+ )
20
+ def list_dead_letters(
21
+ limit: int = Query(default=50, ge=1, le=200),
22
+ ) -> list[DLQRecordResponse]:
23
+ with UnitOfWork() as uow:
24
+ records = uow.dispatch_outbox.dead_letters(limit=limit)
25
+ return [
26
+ DLQRecordResponse(
27
+ id=rec.id,
28
+ workflow_run_id=rec.workflow_run_id,
29
+ task_run_id=rec.task_run_id,
30
+ task_name=rec.task_name,
31
+ payload=rec.payload,
32
+ publish_attempts=rec.publish_attempts,
33
+ last_error=rec.last_error,
34
+ created_at=rec.created_at,
35
+ )
36
+ for rec in records
37
+ ]
38
+
39
+
40
+ @router.post(
41
+ "/{outbox_id}/replay",
42
+ response_model=ReplayDLQResponse,
43
+ summary="Replay a dead-lettered message and attempt immediate re-publication",
44
+ )
45
+ def replay_dead_letter(outbox_id: UUID) -> ReplayDLQResponse:
46
+ dispatcher = DAGDispatcher(uow_factory=UnitOfWork)
47
+ success = dispatcher.replay_dead_letter(outbox_id)
48
+ if success:
49
+ return ReplayDLQResponse(
50
+ outbox_id=outbox_id,
51
+ replayed=True,
52
+ message="Outbox record successfully replayed and published.",
53
+ )
54
+ raise HTTPException(
55
+ status_code=status.HTTP_404_NOT_FOUND,
56
+ detail=f"Dead-letter record {outbox_id} not found or is not dead-lettered.",
57
+ )
58
+
@@ -0,0 +1,31 @@
1
+ """Metrics API endpoints for KEDA and Prometheus monitoring."""
2
+ from __future__ import annotations
3
+
4
+ from fastapi import APIRouter
5
+ from pydantic import BaseModel
6
+
7
+ from celery_dag.models.enums import TaskStatus, WorkflowStatus
8
+ from celery_dag.persistence.unit_of_work import UnitOfWork
9
+
10
+ router = APIRouter(prefix="/metrics", tags=["Metrics"])
11
+
12
+
13
+ class SystemMetricsResponse(BaseModel):
14
+ running_workflows_count: int
15
+ pending_outbox_count: int
16
+ dead_letters_count: int
17
+
18
+
19
+ @router.get("", response_model=SystemMetricsResponse, summary="Get system metrics for KEDA/Prometheus autoscaling")
20
+ def get_metrics() -> SystemMetricsResponse:
21
+ with UnitOfWork() as uow:
22
+ running_ids = uow.workflows.running_ids(limit=10000)
23
+ pending_outbox = uow.dispatch_outbox.pending(limit=10000)
24
+ dead_letters = uow.dispatch_outbox.dead_letters(limit=10000)
25
+
26
+ return SystemMetricsResponse(
27
+ running_workflows_count=len(running_ids),
28
+ pending_outbox_count=len(pending_outbox),
29
+ dead_letters_count=len(dead_letters),
30
+ )
31
+
@@ -0,0 +1,19 @@
1
+ """Registered tasks endpoint."""
2
+ from __future__ import annotations
3
+
4
+ from fastapi import APIRouter
5
+
6
+ from celery_dag.api.schemas import RegisteredTasksResponse
7
+ from celery_dag.registry.task_registry import task_hub
8
+
9
+ router = APIRouter(prefix="/tasks", tags=["Tasks"])
10
+
11
+
12
+ @router.get(
13
+ "",
14
+ response_model=RegisteredTasksResponse,
15
+ summary="List all registered task callables available for DAG nodes",
16
+ )
17
+ def list_registered_tasks() -> RegisteredTasksResponse:
18
+ return RegisteredTasksResponse(tasks=list(task_hub.names()))
19
+
@@ -0,0 +1,89 @@
1
+ """Workflow triggers API endpoints."""
2
+ from __future__ import annotations
3
+
4
+ from typing import Any
5
+ from uuid import UUID
6
+
7
+ from fastapi import APIRouter, HTTPException, Query, status
8
+ from pydantic import BaseModel, Field
9
+
10
+ from celery_dag.engine.trigger_service import TriggerService
11
+ from celery_dag.persistence.unit_of_work import UnitOfWork
12
+
13
+ router = APIRouter(prefix="/triggers", tags=["Triggers"])
14
+
15
+
16
+ class CreateTriggerRequest(BaseModel):
17
+ name: str = Field(..., description="Unique trigger name.")
18
+ trigger_type: str = Field(..., description="CRON or WEBHOOK")
19
+ cron_expression: str | None = Field(default=None, description="Standard 5-part cron expression for CRON trigger.")
20
+ dag_definition: dict[str, Any] = Field(..., description="Complete DAG definition dictionary.")
21
+
22
+
23
+ class TriggerResponse(BaseModel):
24
+ id: UUID
25
+ name: str
26
+ trigger_type: str
27
+ cron_expression: str | None = None
28
+ is_active: bool
29
+ last_triggered_at: Any | None = None
30
+
31
+
32
+ class FireWebhookResponse(BaseModel):
33
+ trigger_name: str
34
+ workflow_run_id: UUID
35
+ message: str
36
+
37
+
38
+ @router.post("", response_model=TriggerResponse, status_code=status.HTTP_201_CREATED, summary="Create a Workflow Trigger")
39
+ def create_trigger(request: CreateTriggerRequest) -> TriggerResponse:
40
+ try:
41
+ trigger = TriggerService().create_trigger(
42
+ name=request.name,
43
+ trigger_type=request.trigger_type,
44
+ dag_definition=request.dag_definition,
45
+ cron_expression=request.cron_expression,
46
+ )
47
+ return TriggerResponse(
48
+ id=trigger.id,
49
+ name=trigger.name,
50
+ trigger_type=trigger.trigger_type,
51
+ cron_expression=trigger.cron_expression,
52
+ is_active=trigger.is_active,
53
+ last_triggered_at=trigger.last_triggered_at,
54
+ )
55
+ except Exception as exc:
56
+ raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
57
+
58
+
59
+ @router.get("", response_model=list[TriggerResponse], summary="List all Workflow Triggers")
60
+ def list_triggers(limit: int = Query(default=50, ge=1, le=200)) -> list[TriggerResponse]:
61
+ with UnitOfWork() as uow:
62
+ triggers = uow.triggers.list_triggers(limit=limit)
63
+ return [
64
+ TriggerResponse(
65
+ id=tr.id,
66
+ name=tr.name,
67
+ trigger_type=tr.trigger_type,
68
+ cron_expression=tr.cron_expression,
69
+ is_active=tr.is_active,
70
+ last_triggered_at=tr.last_triggered_at,
71
+ )
72
+ for tr in triggers
73
+ ]
74
+
75
+
76
+ @router.post("/{trigger_name}/event", response_model=FireWebhookResponse, summary="Fire a Webhook trigger")
77
+ def fire_webhook(trigger_name: str, event_payload: dict[str, Any] | None = None) -> FireWebhookResponse:
78
+ run_id = TriggerService().fire_webhook_trigger(trigger_name, event_payload=event_payload)
79
+ if run_id:
80
+ return FireWebhookResponse(
81
+ trigger_name=trigger_name,
82
+ workflow_run_id=run_id,
83
+ message="Webhook trigger fired successfully; workflow run initiated.",
84
+ )
85
+ raise HTTPException(
86
+ status_code=status.HTTP_404_NOT_FOUND,
87
+ detail=f"Active webhook trigger {trigger_name!r} not found.",
88
+ )
89
+
@@ -0,0 +1,336 @@
1
+ """Workflow management endpoints."""
2
+ from __future__ import annotations
3
+
4
+ from datetime import datetime, timezone
5
+ from uuid import UUID
6
+
7
+ from fastapi import APIRouter, HTTPException, Query, status
8
+
9
+ from celery_dag.api.schemas import (
10
+ CancelWorkflowRequest,
11
+ CancelWorkflowResponse,
12
+ CleanupWorkflowsResponse,
13
+ SubmitWorkflowRequest,
14
+ SubmitWorkflowResponse,
15
+ TaskRunResponse,
16
+ WorkflowControlActionResponse,
17
+ WorkflowRunDetailResponse,
18
+ WorkflowRunResponse,
19
+ )
20
+ from celery_dag.core.dag import DAGBuilder
21
+ from celery_dag.core.exceptions import (
22
+ DAGValidationError,
23
+ DuplicateNodeError,
24
+ UnknownNodeError,
25
+ WorkflowNotFoundError,
26
+ )
27
+ from celery_dag.engine.cancellation import cancel_workflow
28
+ from celery_dag.engine.executor import DAGExecutor
29
+ from celery_dag.engine.workflow_control import pause_workflow, resume_workflow
30
+ from celery_dag.middleware.retry_policy import RetryPolicy
31
+ from celery_dag.models.enums import WorkflowStatus
32
+ from celery_dag.persistence.unit_of_work import UnitOfWork
33
+
34
+ router = APIRouter(prefix="/workflows", tags=["Workflows"])
35
+
36
+
37
+ @router.post(
38
+ "/submit",
39
+ response_model=SubmitWorkflowResponse,
40
+ status_code=status.HTTP_201_CREATED,
41
+ summary="Submit a new DAG workflow run",
42
+ )
43
+ def submit_workflow(request: SubmitWorkflowRequest) -> SubmitWorkflowResponse:
44
+ try:
45
+ builder = DAGBuilder(
46
+ dag_id=request.dag_id,
47
+ version=request.version,
48
+ timeout=request.timeout,
49
+ metadata=request.metadata,
50
+ )
51
+ for node in request.nodes:
52
+ builder.add_node(
53
+ node_id=node.node_id,
54
+ task_name=node.task_name,
55
+ depends_on=node.dependencies,
56
+ retry_policy=RetryPolicy(
57
+ max_retries=node.retry_policy.max_retries,
58
+ delay=node.retry_policy.delay,
59
+ backoff=node.retry_policy.backoff,
60
+ max_delay=node.retry_policy.max_delay,
61
+ jitter=node.retry_policy.jitter,
62
+ ),
63
+ timeout=node.timeout,
64
+ priority=node.priority,
65
+ queue=node.queue,
66
+ payload=node.payload,
67
+ )
68
+ dag = builder.build()
69
+ run_id = DAGExecutor().submit(dag)
70
+ return SubmitWorkflowResponse(
71
+ workflow_run_id=run_id,
72
+ dag_id=dag.dag_id,
73
+ version=dag.version,
74
+ status=WorkflowStatus.RUNNING if dag.nodes else WorkflowStatus.COMPLETED,
75
+ submitted_at=datetime.now(timezone.utc),
76
+ )
77
+ except (DAGValidationError, DuplicateNodeError, UnknownNodeError) as exc:
78
+ raise HTTPException(
79
+ status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)
80
+ ) from exc
81
+
82
+
83
+ @router.get(
84
+ "",
85
+ response_model=list[WorkflowRunResponse],
86
+ summary="List workflow runs with filtering and pagination",
87
+ )
88
+ def list_workflows(
89
+ status_filter: WorkflowStatus | None = Query(default=None, alias="status"),
90
+ dag_id: str | None = Query(default=None, alias="dag_id"),
91
+ limit: int = Query(default=50, ge=1, le=200),
92
+ offset: int = Query(default=0, ge=0),
93
+ ) -> list[WorkflowRunResponse]:
94
+ with UnitOfWork() as uow:
95
+ runs = uow.workflows.list(
96
+ status=status_filter, dag_id=dag_id, limit=limit, offset=offset
97
+ )
98
+ return [
99
+ WorkflowRunResponse(
100
+ id=run.id,
101
+ dag_id=run.dag_id,
102
+ version=run.version,
103
+ status=run.status,
104
+ timeout_seconds=run.timeout_seconds,
105
+ run_metadata=run.run_metadata,
106
+ error=run.error,
107
+ created_at=run.created_at,
108
+ started_at=run.started_at,
109
+ finished_at=run.finished_at,
110
+ cancelled_at=run.cancelled_at,
111
+ cancellation_reason=run.cancellation_reason,
112
+ )
113
+ for run in runs
114
+ ]
115
+
116
+
117
+ @router.post(
118
+ "/cleanup",
119
+ response_model=CleanupWorkflowsResponse,
120
+ summary="Purge completed/failed/cancelled workflows past retention period",
121
+ )
122
+ def cleanup_workflows(
123
+ retention_days: int = Query(default=30, ge=1, le=365, description="Retention window in days."),
124
+ ) -> CleanupWorkflowsResponse:
125
+ with UnitOfWork() as uow:
126
+ count = uow.workflows.delete_expired_workflows(retention_days=retention_days)
127
+ uow.commit()
128
+ return CleanupWorkflowsResponse(
129
+ deleted_count=count,
130
+ retention_days=retention_days,
131
+ message=f"Purged {count} expired workflows older than {retention_days} days.",
132
+ )
133
+
134
+
135
+ @router.get(
136
+ "/{run_id}",
137
+ response_model=WorkflowRunDetailResponse,
138
+ summary="Get detailed workflow run info and all task runs",
139
+ )
140
+ def get_workflow(run_id: UUID) -> WorkflowRunDetailResponse:
141
+ with UnitOfWork() as uow:
142
+ try:
143
+ workflow = uow.workflows.get(run_id)
144
+ task_runs = uow.task_runs.list_for_workflow(run_id)
145
+ return WorkflowRunDetailResponse(
146
+ id=workflow.id,
147
+ dag_id=workflow.dag_id,
148
+ version=workflow.version,
149
+ status=workflow.status,
150
+ timeout_seconds=workflow.timeout_seconds,
151
+ run_metadata=workflow.run_metadata,
152
+ error=workflow.error,
153
+ created_at=workflow.created_at,
154
+ started_at=workflow.started_at,
155
+ finished_at=workflow.finished_at,
156
+ cancelled_at=workflow.cancelled_at,
157
+ cancellation_reason=workflow.cancellation_reason,
158
+ dag_definition=workflow.dag_definition,
159
+ task_runs=[
160
+ TaskRunResponse(
161
+ id=tr.id,
162
+ workflow_run_id=tr.workflow_run_id,
163
+ node_id=tr.node_id,
164
+ task_name=tr.task_name,
165
+ status=tr.status,
166
+ attempt=tr.attempt,
167
+ input_data=tr.input_data,
168
+ result=tr.result,
169
+ error=tr.error,
170
+ started_at=tr.started_at,
171
+ finished_at=tr.finished_at,
172
+ )
173
+ for tr in task_runs
174
+ ],
175
+ )
176
+ except WorkflowNotFoundError as exc:
177
+ raise HTTPException(
178
+ status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)
179
+ ) from exc
180
+
181
+
182
+ @router.post(
183
+ "/{run_id}/pause",
184
+ response_model=WorkflowControlActionResponse,
185
+ summary="Pause execution of a running workflow",
186
+ )
187
+ def pause_workflow_endpoint(run_id: UUID) -> WorkflowControlActionResponse:
188
+ success = pause_workflow(run_id)
189
+ if success:
190
+ return WorkflowControlActionResponse(
191
+ workflow_run_id=run_id,
192
+ action="pause",
193
+ success=True,
194
+ message="Workflow paused successfully. New nodes will not be queued until resumed.",
195
+ )
196
+ return WorkflowControlActionResponse(
197
+ workflow_run_id=run_id,
198
+ action="pause",
199
+ success=False,
200
+ message="Workflow is not currently in RUNNING status.",
201
+ )
202
+
203
+
204
+ @router.post(
205
+ "/{run_id}/resume",
206
+ response_model=WorkflowControlActionResponse,
207
+ summary="Resume execution of a paused workflow",
208
+ )
209
+ def resume_workflow_endpoint(run_id: UUID) -> WorkflowControlActionResponse:
210
+ success = resume_workflow(run_id)
211
+ if success:
212
+ return WorkflowControlActionResponse(
213
+ workflow_run_id=run_id,
214
+ action="resume",
215
+ success=True,
216
+ message="Workflow resumed successfully. Ready nodes have been scheduled.",
217
+ )
218
+ return WorkflowControlActionResponse(
219
+ workflow_run_id=run_id,
220
+ action="resume",
221
+ success=False,
222
+ message="Workflow is not currently in PAUSED status.",
223
+ )
224
+
225
+
226
+ @router.post(
227
+ "/{run_id}/cancel",
228
+ response_model=CancelWorkflowResponse,
229
+ summary="Request cancellation of a running workflow",
230
+ )
231
+ def request_cancellation(
232
+ run_id: UUID, request: CancelWorkflowRequest | None = None
233
+ ) -> CancelWorkflowResponse:
234
+ reason = request.reason if request else None
235
+ cancelled = cancel_workflow(run_id, reason=reason)
236
+ if cancelled:
237
+ return CancelWorkflowResponse(
238
+ workflow_run_id=run_id,
239
+ cancelled=True,
240
+ message="Workflow cancellation initiated successfully.",
241
+ )
242
+ return CancelWorkflowResponse(
243
+ workflow_run_id=run_id,
244
+ cancelled=False,
245
+ message="Workflow was already cancelled or in terminal state.",
246
+ )
247
+
248
+
249
+ @router.get(
250
+ "/{run_id}/tasks/{node_id}",
251
+ response_model=TaskRunResponse,
252
+ summary="Get status and details of a specific task node",
253
+ )
254
+ def get_task_node(run_id: UUID, node_id: str) -> TaskRunResponse:
255
+ with UnitOfWork() as uow:
256
+ try:
257
+ tr = uow.task_runs.get_for_node(run_id, node_id)
258
+ return TaskRunResponse(
259
+ id=tr.id,
260
+ workflow_run_id=tr.workflow_run_id,
261
+ node_id=tr.node_id,
262
+ task_name=tr.task_name,
263
+ status=tr.status,
264
+ attempt=tr.attempt,
265
+ input_data=tr.input_data,
266
+ result=tr.result,
267
+ error=tr.error,
268
+ started_at=tr.started_at,
269
+ finished_at=tr.finished_at,
270
+ )
271
+ except Exception as exc:
272
+ raise HTTPException(
273
+ status_code=status.HTTP_404_NOT_FOUND,
274
+ detail=f"Task node {node_id!r} not found for workflow {run_id}",
275
+ ) from exc
276
+
277
+
278
+ @router.post(
279
+ "/{run_id}/resume-failed",
280
+ response_model=WorkflowControlActionResponse,
281
+ summary="Resume execution of a FAILED or CANCELLED workflow from point of failure",
282
+ )
283
+ def resume_failed_workflow(run_id: UUID) -> WorkflowControlActionResponse:
284
+ try:
285
+ executor = DAGExecutor()
286
+ executor.resume_from_failure(run_id)
287
+ return WorkflowControlActionResponse(
288
+ workflow_run_id=run_id,
289
+ action="resume-failed",
290
+ success=True,
291
+ message="Workflow resumed successfully from point of failure.",
292
+ )
293
+ except ValueError as exc:
294
+ return WorkflowControlActionResponse(
295
+ workflow_run_id=run_id,
296
+ action="resume-failed",
297
+ success=False,
298
+ message=str(exc),
299
+ )
300
+ except Exception as exc:
301
+ raise HTTPException(
302
+ status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
303
+ detail=f"Failed to resume workflow: {exc}",
304
+ ) from exc
305
+
306
+
307
+ @router.get(
308
+ "/{run_id}/visualization",
309
+ summary="Generate live Mermaid diagram or Cytoscape JSON visualization of DAG run",
310
+ )
311
+ def get_workflow_visualization(
312
+ run_id: UUID,
313
+ format: str = Query(default="mermaid", description="Format: 'mermaid' or 'json'"),
314
+ ) -> Any:
315
+ from celery_dag.observability.dag_visualizer import (
316
+ generate_cytoscape_json,
317
+ generate_mermaid_diagram,
318
+ )
319
+
320
+ with UnitOfWork() as uow:
321
+ try:
322
+ workflow = uow.workflows.get(run_id)
323
+ task_runs = uow.task_runs.list_for_workflow(run_id)
324
+ dag = DAG.from_dict(workflow.dag_definition)
325
+ except Exception as exc:
326
+ raise HTTPException(
327
+ status_code=status.HTTP_404_NOT_FOUND,
328
+ detail=f"Workflow run not found: {run_id}",
329
+ ) from exc
330
+
331
+ if format.lower() == "json":
332
+ return generate_cytoscape_json(dag, task_runs)
333
+
334
+ mermaid_diagram = generate_mermaid_diagram(dag, task_runs)
335
+ return {"run_id": run_id, "format": "mermaid", "diagram": mermaid_diagram}
336
+