From be43a199f62ffd7d51cc2cd73b2327dd521bf796 Mon Sep 17 00:00:00 2001 From: yashlamba Date: Wed, 15 Jul 2026 15:49:24 +0200 Subject: [PATCH] fix: move workflow create to service --- app/database/session.py | 2 +- app/main.py | 9 ++--- app/routers/workflows.py | 66 +++--------------------------- app/services/workflows.py | 84 ++++++++++++++++++++++++++++++++++++++- app/temporal.py | 35 ++++++++++++++++ tests/conftest.py | 10 ++--- tests/test_auth.py | 17 ++++++-- 7 files changed, 144 insertions(+), 79 deletions(-) create mode 100644 app/temporal.py diff --git a/app/database/session.py b/app/database/session.py index d18aaa2..680b2e6 100644 --- a/app/database/session.py +++ b/app/database/session.py @@ -38,7 +38,7 @@ def get_engine() -> Engine: @contextmanager def get_session() -> Generator[Session, None, None]: """Provide a database session context manager.""" - with Session(get_engine()) as session: + with Session(get_engine(), expire_on_commit=False) as session: yield session diff --git a/app/main.py b/app/main.py index bf82b66..c5822a3 100644 --- a/app/main.py +++ b/app/main.py @@ -14,13 +14,12 @@ from fastapi.exceptions import RequestValidationError from fastapi.middleware.cors import CORSMiddleware from starlette.exceptions import HTTPException as StarletteHTTPException -from temporalio.client import Client -from temporalio.contrib.pydantic import pydantic_data_converter from .config import get_settings from .database.session import dispose_engine, init_engine from .dependencies import get_current_user from .routers import workflow_feedback, workflows +from .temporal import dispose_temporal_client, init_temporal_client from .tenants import TenantRegistry logger = logging.getLogger(__name__) @@ -32,10 +31,7 @@ async def lifespan(app: FastAPI): settings = get_settings() engine = init_engine() app.state.db_engine = engine - app.state.temporal_client = await Client.connect( - settings.temporal_host, - data_converter=pydantic_data_converter, - ) + await init_temporal_client() # Load tenant registry if not settings.auth_disabled: @@ -46,6 +42,7 @@ async def lifespan(app: FastAPI): app.state.tenant_registry = TenantRegistry() yield + dispose_temporal_client() dispose_engine() diff --git a/app/routers/workflows.py b/app/routers/workflows.py index 524cd65..bf966a2 100644 --- a/app/routers/workflows.py +++ b/app/routers/workflows.py @@ -5,24 +5,19 @@ import asyncio import logging -from datetime import UTC, datetime from typing import Any from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import StreamingResponse -from pydantic import BaseModel, Field, ValidationError +from pydantic import BaseModel, Field from sqlalchemy.exc import SQLAlchemyError from sqlmodel import Session, select -from temporalio.client import Client -from temporalio.common import RetryPolicy from app.auth import AuthContext, decode_access_token -from app.database.models import Workflow, WorkflowStatus +from app.database.models import Workflow from app.database.session import get_db_session from app.dependencies import get_current_user from app.services.workflows import WorkflowService -from app.workflows.registry import get_workflow_spec -from app.workflows.specs import WorkflowContext logger = logging.getLogger(__name__) @@ -43,10 +38,6 @@ class CreateWorkflowRequest(BaseModel): user_id: str | None = None -def _get_temporal_client(request: Request) -> Client: - return request.app.state.temporal_client - - @router.post( "/", response_model=Workflow, @@ -54,65 +45,18 @@ def _get_temporal_client(request: Request) -> Client: ) async def create( body: CreateWorkflowRequest, - request: Request, auth: AuthContext = Depends(get_current_user), session: Session = Depends(get_db_session), ): """Create a new workflow and start the Temporal workflow.""" - try: - spec = get_workflow_spec(body.workflow_type) - except KeyError as exc: - raise HTTPException(status_code=400, detail=str(exc)) - - try: - params = spec.params_model.model_validate(body.params) - except ValidationError as exc: - raise HTTPException(status_code=422, detail=exc.errors()) from exc - - workflow = Workflow( + return await WorkflowService(session).create( workflow_type=body.workflow_type, - status=WorkflowStatus.PROCESSING, - params=params.model_dump(mode="json"), + params=body.params, tenant_id=auth.tenant_id, user_id=body.user_id, + start=True, ) - try: - session.add(workflow) - session.commit() - workflow_id = workflow.public_id - except SQLAlchemyError: - logger.exception("Error creating workflow") - raise HTTPException(status_code=500, detail="Could not create workflow") - - try: - client = _get_temporal_client(request) - await client.start_workflow( - spec.workflow_cls.run, - args=[ - WorkflowContext( - workflow_id=workflow_id, - tenant_id=auth.tenant_id, - user_id=workflow.user_id, - ), - params, - ], - id=f"{spec.id_prefix}-{workflow_id}", - task_queue=spec.task_queue, - retry_policy=RetryPolicy(maximum_attempts=1), - ) - except Exception: - logger.exception("Error starting Temporal workflow") - try: - workflow.status = WorkflowStatus.ERROR - workflow.end_time = datetime.now(UTC) - session.commit() - except SQLAlchemyError: - pass - raise HTTPException(status_code=500, detail="Could not start workflow") - - return workflow - @router.get( "/{workflow_id}", diff --git a/app/services/workflows.py b/app/services/workflows.py index 18e1ee1..86e09b9 100644 --- a/app/services/workflows.py +++ b/app/services/workflows.py @@ -4,19 +4,26 @@ """Workflow database operations and access checks.""" import logging +from datetime import UTC, datetime +from typing import Any from fastapi import HTTPException +from pydantic import ValidationError from sqlalchemy.exc import SQLAlchemyError from sqlmodel import Session, select +from temporalio.common import RetryPolicy from app.auth import AuthContext -from app.database.models import Workflow +from app.database.models import Workflow, WorkflowStatus +from app.temporal import get_temporal_client +from app.workflows.registry import get_workflow_spec +from app.workflows.specs import WorkflowContext logger = logging.getLogger(__name__) class WorkflowService: - """Service for workflow lookup and tenant-scoped access checks.""" + """Service for workflows.""" def __init__(self, session: Session): """Create a workflow service for one database session.""" @@ -76,3 +83,76 @@ def get_tenant_workflow( workflow = self.get_by_public_id(workflow_id) self.verify_tenant_owns_workflow(auth, workflow) return workflow + + async def create( + self, + *, + workflow_type: str, + params: dict[str, Any], + tenant_id: str, + user_id: str | None = None, + start: bool = True, + ) -> Workflow: + """Create a workflow record and optionally start its Temporal workflow. + + The insert commits before the Temporal call so no database transaction + stays open while the external request is in flight. + """ + try: + spec = get_workflow_spec(workflow_type) + except KeyError as exc: + raise HTTPException(status_code=400, detail=str(exc)) + + try: + workflow_params = spec.params_model.model_validate(params) + except ValidationError as exc: + raise HTTPException(status_code=422, detail=exc.errors()) from exc + + workflow = Workflow( + workflow_type=workflow_type, + status=WorkflowStatus.PROCESSING, + params=workflow_params.model_dump(mode="json"), + tenant_id=tenant_id, + user_id=user_id, + ) + workflow_id = workflow.public_id + + try: + self.session.add(workflow) + self.session.commit() + except SQLAlchemyError: + logger.exception("Error creating workflow") + raise HTTPException(status_code=500, detail="Could not create workflow") + + if start: + try: + await get_temporal_client().start_workflow( + spec.workflow_cls.run, + args=[ + WorkflowContext( + workflow_id=workflow_id, + tenant_id=tenant_id, + user_id=user_id, + ), + workflow_params, + ], + id=f"{spec.id_prefix}-{workflow_id}", + task_queue=spec.task_queue, + retry_policy=RetryPolicy(maximum_attempts=1), + ) + except Exception: + logger.exception("Error starting Temporal workflow") + self._mark_workflow_error(workflow) + raise HTTPException(status_code=500, detail="Could not start workflow") + + return workflow + + def _mark_workflow_error(self, workflow: Workflow) -> None: + """Record a failed Temporal start on the workflow row.""" + try: + workflow.status = WorkflowStatus.ERROR + workflow.end_time = datetime.now(UTC) + self.session.add(workflow) + self.session.commit() + except SQLAlchemyError: + logger.exception("Could not mark workflow as errored") diff --git a/app/temporal.py b/app/temporal.py new file mode 100644 index 0000000..729210a --- /dev/null +++ b/app/temporal.py @@ -0,0 +1,35 @@ +# SPDX-FileCopyrightText: 2026 CERN. +# SPDX-License-Identifier: MIT + +"""Application-wide Temporal client, initialized at startup.""" + +from temporalio.client import Client +from temporalio.contrib.pydantic import pydantic_data_converter + +from .config import get_settings + +_client: Client | None = None + + +async def init_temporal_client() -> Client: + """Connect and store the global Temporal client.""" + global _client + settings = get_settings() + _client = await Client.connect( + settings.temporal_host, + data_converter=pydantic_data_converter, + ) + return _client + + +def dispose_temporal_client() -> None: + """Forget the global Temporal client.""" + global _client + _client = None + + +def get_temporal_client() -> Client: + """Return the global Temporal client, if initialized.""" + if _client is None: + raise RuntimeError("Temporal client is not initialized!") + return _client diff --git a/tests/conftest.py b/tests/conftest.py index 02515a3..616178a 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -25,14 +25,14 @@ def db_session(): """Provide a real SQLModel session backed by in-memory SQLite. Creates all tables before the test and drops them after. - Overrides both the FastAPI ``get_session`` dependency and - ``app.state.db_engine`` so that route handlers and the - stream generator use the same test database. + Overrides the FastAPI ``get_session`` dependency and + ``app.state.db_engine`` so that route handlers, services, and the + stream generator all use the same test database. """ SQLModel.metadata.create_all(_test_engine) def _override_get_session(): - with Session(_test_engine) as session: + with Session(_test_engine, expire_on_commit=False) as session: yield session app.dependency_overrides[get_db_session] = _override_get_session @@ -41,7 +41,7 @@ def _override_get_session(): # directly (e.g. the SSE stream generator) uses the test engine. app.state.db_engine = _test_engine - with Session(_test_engine) as session: + with Session(_test_engine, expire_on_commit=False) as session: yield session SQLModel.metadata.drop_all(_test_engine) diff --git a/tests/test_auth.py b/tests/test_auth.py index 0dc3669..2219e3a 100644 --- a/tests/test_auth.py +++ b/tests/test_auth.py @@ -66,7 +66,7 @@ def configure_test_settings(monkeypatch, mocker): get_settings.cache_clear() # Mock Temporal Client - mocker.patch("app.main.Client.connect", return_value=mocker.AsyncMock()) + mocker.patch("app.temporal.Client.connect", return_value=mocker.AsyncMock()) # Patch TenantRegistry.from_file so lifespan doesn't look for a real file mocker.patch( @@ -448,7 +448,10 @@ def test_create_workflow_stamps_tenant_id(client, db_session, mocker): # Mock the temporal client to avoid real connection mock_temporal = mocker.AsyncMock() - mocker.patch.object(client.app.state, "temporal_client", mock_temporal) + mocker.patch( + "app.services.workflows.get_temporal_client", + return_value=mock_temporal, + ) response = client.post( "/workflows/", @@ -491,7 +494,10 @@ def test_create_workflow_rejects_invalid_params(client, db_session, mocker): token = generate_test_token() mock_temporal = mocker.AsyncMock() - mocker.patch.object(client.app.state, "temporal_client", mock_temporal) + mocker.patch( + "app.services.workflows.get_temporal_client", + return_value=mock_temporal, + ) response = client.post( "/workflows/", @@ -512,7 +518,10 @@ def test_create_workflow_rejects_unknown_params(client, db_session, mocker): token = generate_test_token() mock_temporal = mocker.AsyncMock() - mocker.patch.object(client.app.state, "temporal_client", mock_temporal) + mocker.patch( + "app.services.workflows.get_temporal_client", + return_value=mock_temporal, + ) response = client.post( "/workflows/",