Refactor run_flow_with_caching function in endpoints.py

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-02-27 18:13:29 -03:00
commit feac452f1c

View file

@ -220,7 +220,9 @@ async def preload_flow(
raise HTTPException(status_code=500, detail=str(exc)) from exc raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.post("/run/{flow_id}", response_model=ProcessResponse) @router.post(
"/run/{flow_id}", response_model=RunResponse, response_model_exclude_none=True
)
async def run_flow_with_caching( async def run_flow_with_caching(
session: Annotated[Session, Depends(get_session)], session: Annotated[Session, Depends(get_session)],
flow_id: str, flow_id: str,
@ -235,13 +237,13 @@ async def run_flow_with_caching(
session_data = await session_service.load_session(session_id) session_data = await session_service.load_session(session_id)
graph, artifacts = session_data if session_data else (None, None) graph, artifacts = session_data if session_data else (None, None)
task_result: Any = None task_result: Any = None
task_status = None
if not graph: if not graph:
raise ValueError("Graph not found in the session") raise ValueError("Graph not found in the session")
task_result = await run_graph( task_result = await run_graph(
graph, graph=graph,
session_id, flow_id=flow_id,
inputs, session_id=session_id,
inputs=inputs,
artifacts=artifacts, artifacts=artifacts,
session_service=session_service, session_service=session_service,
) )
@ -262,16 +264,15 @@ async def run_flow_with_caching(
graph_data = flow.data graph_data = flow.data
graph_data = process_tweaks(graph_data, tweaks) graph_data = process_tweaks(graph_data, tweaks)
task_result = await run_graph( task_result = await run_graph(
graph_data, graph=graph_data,
inputs, flow_id=flow_id,
tweaks, session_id=session_id,
session_id, inputs=inputs,
artifacts={},
session_service=session_service, session_service=session_service,
) )
return RunResponse( return RunResponse(outputs=task_result, session_id=session_id)
outputs=task_result, session_id=session_id, status=task_status
)
except sa.exc.StatementError as exc: except sa.exc.StatementError as exc:
# StatementError('(builtins.ValueError) badly formed hexadecimal UUID string') # StatementError('(builtins.ValueError) badly formed hexadecimal UUID string')
if "badly formed hexadecimal UUID string" in str(exc): if "badly formed hexadecimal UUID string" in str(exc):