🔧 fix(endpoints.py): import missing dependencies to improve code readability and maintainability

🔧 fix(endpoints.py): generate session ID if it is None to ensure a valid session ID is used
🔧 fix(endpoints.py): add backend information to the ProcessResponse to provide additional context
🔧 fix(schemas.py): add backend field to the ProcessResponse schema to match the changes in the endpoints.py file
This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-09-23 23:28:56 -03:00
commit 8760cf7f9a
2 changed files with 18 additions and 2 deletions

View file

@ -7,7 +7,12 @@ from langflow.services.cache.utils import save_uploaded_file
from langflow.services.database.models.flow import Flow from langflow.services.database.models.flow import Flow
from langflow.processing.process import process_graph_cached, process_tweaks from langflow.processing.process import process_graph_cached, process_tweaks
from langflow.services.database.models.user.user import User from langflow.services.database.models.user.user import User
from langflow.services.getters import get_settings_service, get_task_service from langflow.services.getters import (
get_cache_service,
get_session_service,
get_settings_service,
get_task_service,
)
from loguru import logger from loguru import logger
from fastapi import APIRouter, Depends, HTTPException, UploadFile, Body, status from fastapi import APIRouter, Depends, HTTPException, UploadFile, Body, status
import sqlalchemy as sa import sqlalchemy as sa
@ -122,6 +127,11 @@ async def process_flow(
session_id = result.session_id session_id = result.session_id
else: else:
if session_id is None:
# Generate a session ID
session_id = get_session_service().generate_key(
session_id=session_id, data_graph=graph_data
)
task_id, task = await task_service.launch_task( task_id, task = await task_service.launch_task(
process_graph_cached_task process_graph_cached_task
if task_service.use_celery if task_service.use_celery
@ -132,7 +142,12 @@ async def process_flow(
session_id, session_id,
) )
task_result = task.status task_result = task.status
return ProcessResponse(result=task_result, id=task_id, session_id=session_id) return ProcessResponse(
result=task_result,
id=task_id,
session_id=session_id,
backend=str(type(task_service.backend)),
)
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):

View file

@ -53,6 +53,7 @@ class ProcessResponse(BaseModel):
result: Any result: Any
id: Optional[str] = None id: Optional[str] = None
session_id: Optional[str] = None session_id: Optional[str] = None
backend: str = None
# TaskStatusResponse( # TaskStatusResponse(