|
6 | 6 | from fastapi import APIRouter, Depends, HTTPException, Request |
7 | 7 | from pydantic import BaseModel |
8 | 8 | from sqlalchemy.exc import SQLAlchemyError |
| 9 | +from sqlalchemy.orm.exc import NoResultFound |
9 | 10 | from sqlmodel import Session, select |
10 | 11 | from sse_starlette import EventSourceResponse, ServerSentEvent |
11 | 12 | from temporalio.client import Client |
|
14 | 15 | from app.database.models import Workflow, WorkflowStatus |
15 | 16 | from app.database.session import get_session |
16 | 17 | from app.dependencies import get_current_user |
| 18 | +from app.errors import WorkflowEventError |
17 | 19 | from app.workflows.extract_metadata_workflow import ExtractMetadata |
18 | 20 |
|
19 | 21 | router = APIRouter( |
@@ -150,48 +152,43 @@ async def workflow_event(request: Request, workflow_id: str): |
150 | 152 | if await request.is_disconnected(): |
151 | 153 | break |
152 | 154 |
|
153 | | - with Session(request.app.state.db_engine) as session: |
154 | | - try: |
155 | | - workflow = session.exec( |
156 | | - select(Workflow).where(Workflow.public_id == workflow_id) |
157 | | - ).one() |
158 | | - |
159 | | - status = workflow.status |
160 | | - |
161 | | - if status == WorkflowStatus.SUCCESS: |
162 | | - yield ServerSentEvent( |
163 | | - data=json.dumps(workflow.result), event="metadata" |
164 | | - ) |
165 | | - yield ServerSentEvent(data="done", event="end") |
166 | | - break |
167 | | - |
168 | | - if status == WorkflowStatus.ERROR: |
169 | | - yield ServerSentEvent( |
170 | | - # TODO: improve it with a better error message for end users |
171 | | - data="The Temporal Workflow failed", |
172 | | - event="error", |
173 | | - ) |
174 | | - yield ServerSentEvent(data="done", event="end") |
175 | | - break |
176 | | - |
177 | | - except SQLAlchemyError as e: |
178 | | - print("Error in fetching from database (stream_workflow)", e) |
| 155 | + try: |
| 156 | + with Session(request.app.state.db_engine) as session: |
| 157 | + try: |
| 158 | + workflow = session.exec( |
| 159 | + select(Workflow).where(Workflow.public_id == workflow_id) |
| 160 | + ).one() |
| 161 | + except NoResultFound: |
| 162 | + raise WorkflowEventError(error_code="WORKFLOW_NOT_FOUND") |
| 163 | + except SQLAlchemyError as e: |
| 164 | + print("Error in fetching from database (stream_workflow)", e) |
| 165 | + raise WorkflowEventError(error_code="DB_FETCH_FAILED") |
| 166 | + |
| 167 | + status = workflow.status |
| 168 | + |
| 169 | + if status == WorkflowStatus.SUCCESS: |
179 | 170 | yield ServerSentEvent( |
180 | | - data="Failed to read workflow status.", |
181 | | - event="error", |
| 171 | + data=json.dumps(workflow.result), event="metadata" |
182 | 172 | ) |
183 | 173 | yield ServerSentEvent(data="done", event="end") |
184 | 174 | break |
185 | 175 |
|
186 | | - except Exception as e: |
187 | | - print("Error(stream_workflow)", e) |
| 176 | + if status == WorkflowStatus.ERROR: |
188 | 177 | yield ServerSentEvent( |
189 | | - data="An unexpected error occurred while streaming workflow results.", # noqa: E501 |
| 178 | + data=json.dumps({"error_code": "WORKFLOW_FAILED"}), |
190 | 179 | event="error", |
191 | 180 | ) |
192 | 181 | yield ServerSentEvent(data="done", event="end") |
193 | 182 | break |
194 | 183 |
|
| 184 | + except WorkflowEventError as e: |
| 185 | + yield ServerSentEvent( |
| 186 | + data=json.dumps({"error_code": e.error_code}), |
| 187 | + event="error", |
| 188 | + ) |
| 189 | + yield ServerSentEvent(data="done", event="end") |
| 190 | + break |
| 191 | + |
195 | 192 | await asyncio.sleep(STREAM_DELAY) |
196 | 193 |
|
197 | 194 |
|
|
0 commit comments