Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion app/database/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
9 changes: 3 additions & 6 deletions app/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand All @@ -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:
Expand All @@ -46,6 +42,7 @@ async def lifespan(app: FastAPI):
app.state.tenant_registry = TenantRegistry()

yield
dispose_temporal_client()
dispose_engine()


Expand Down
66 changes: 5 additions & 61 deletions app/routers/workflows.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand All @@ -43,76 +38,25 @@ 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,
response_model_exclude={"id"},
)
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}",
Expand Down
84 changes: 82 additions & 2 deletions app/services/workflows.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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")
35 changes: 35 additions & 0 deletions app/temporal.py
Original file line number Diff line number Diff line change
@@ -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
10 changes: 5 additions & 5 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
17 changes: 13 additions & 4 deletions tests/test_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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/",
Expand Down Expand Up @@ -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/",
Expand All @@ -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/",
Expand Down