diff --git a/src/backend/base/langflow/graph/utils.py b/src/backend/base/langflow/graph/utils.py index 4ad1b5235..b1a69efc1 100644 --- a/src/backend/base/langflow/graph/utils.py +++ b/src/backend/base/langflow/graph/utils.py @@ -18,7 +18,7 @@ from langflow.services.database.models.transactions.crud import log_transaction from langflow.services.database.models.transactions.model import TransactionBase from langflow.services.database.models.vertex_builds.crud import log_vertex_build as crud_log_vertex_build from langflow.services.database.models.vertex_builds.model import VertexBuildBase -from langflow.services.database.utils import session_getter +from langflow.services.database.utils import async_session_getter from langflow.services.deps import get_db_service, get_settings_service if TYPE_CHECKING: @@ -157,14 +157,14 @@ async def log_transaction( error=error, flow_id=flow_id if isinstance(flow_id, UUID) else UUID(flow_id), ) - with session_getter(get_db_service()) as session: - inserted = crud_log_transaction(session, transaction) + async with async_session_getter(get_db_service()) as session: + inserted = await crud_log_transaction(session, transaction) logger.debug(f"Logged transaction: {inserted.id}") except Exception: # noqa: BLE001 logger.exception("Error logging transaction") -def log_vertex_build( +async def log_vertex_build( *, flow_id: str, vertex_id: str, @@ -186,8 +186,8 @@ def log_vertex_build( # ugly hack to get the model dump with weird datatypes artifacts=json.loads(json.dumps(artifacts, default=str)), ) - with session_getter(get_db_service()) as session: - inserted = crud_log_vertex_build(session, vertex_build) + async with async_session_getter(get_db_service()) as session: + inserted = await crud_log_vertex_build(session, vertex_build) logger.debug(f"Logged vertex build: {inserted.build_id}") except Exception: # noqa: BLE001 logger.exception("Error logging vertex build") diff --git a/src/backend/base/langflow/graph/vertex/types.py b/src/backend/base/langflow/graph/vertex/types.py index 696d47ed7..d57039939 100644 --- a/src/backend/base/langflow/graph/vertex/types.py +++ b/src/backend/base/langflow/graph/vertex/types.py @@ -435,7 +435,7 @@ class InterfaceVertex(ComponentVertex): and hasattr(self.custom_component, "store_message") ): self.custom_component.store_message(message) - log_vertex_build( + await log_vertex_build( flow_id=self.graph.flow_id, vertex_id=self.id, valid=True, diff --git a/src/backend/base/langflow/helpers/flow.py b/src/backend/base/langflow/helpers/flow.py index 675e14c5c..73296913a 100644 --- a/src/backend/base/langflow/helpers/flow.py +++ b/src/backend/base/langflow/helpers/flow.py @@ -10,7 +10,7 @@ from sqlmodel import select from langflow.schema.schema import INPUT_FIELD_NAME from langflow.services.database.models.flow import Flow from langflow.services.database.models.flow.model import FlowRead -from langflow.services.deps import get_settings_service, session_scope +from langflow.services.deps import async_session_scope, get_settings_service, session_scope if TYPE_CHECKING: from collections.abc import Awaitable, Callable @@ -53,13 +53,13 @@ async def load_flow( msg = "Flow ID or Flow Name is required" raise ValueError(msg) if not flow_id and flow_name: - flow_id = find_flow(flow_name, user_id) + flow_id = await find_flow(flow_name, user_id) if not flow_id: msg = f"Flow {flow_name} not found" raise ValueError(msg) - with session_scope() as session: - graph_data = flow.data if (flow := session.get(Flow, flow_id)) else None + async with async_session_scope() as session: + graph_data = flow.data if (flow := await session.get(Flow, flow_id)) else None if not graph_data: msg = f"Flow {flow_id} not found" raise ValueError(msg) @@ -68,9 +68,10 @@ async def load_flow( return Graph.from_payload(graph_data, flow_id=flow_id, user_id=user_id) -def find_flow(flow_name: str, user_id: str) -> str | None: - with session_scope() as session: - flow = session.exec(select(Flow).where(Flow.name == flow_name).where(Flow.user_id == user_id)).first() +async def find_flow(flow_name: str, user_id: str) -> str | None: + async with async_session_scope() as session: + stmt = select(Flow).where(Flow.name == flow_name).where(Flow.user_id == user_id) + flow = (await session.exec(stmt)).first() return flow.id if flow else None @@ -273,18 +274,18 @@ def get_arg_names(inputs: list[Vertex]) -> list[dict[str, str]]: ] -def get_flow_by_id_or_endpoint_name(flow_id_or_name: str, user_id: UUID | None = None) -> FlowRead | None: - with session_scope() as session: +async def get_flow_by_id_or_endpoint_name(flow_id_or_name: str, user_id: UUID | None = None) -> FlowRead | None: + async with async_session_scope() as session: endpoint_name = None try: flow_id = UUID(flow_id_or_name) - flow = session.get(Flow, flow_id) + flow = await session.get(Flow, flow_id) except ValueError: endpoint_name = flow_id_or_name stmt = select(Flow).where(Flow.endpoint_name == endpoint_name) if user_id: stmt = stmt.where(Flow.user_id == user_id) - flow = session.exec(stmt).first() + flow = (await session.exec(stmt)).first() if flow is None: raise HTTPException(status_code=404, detail=f"Flow identifier {flow_id_or_name} not found") return FlowRead.model_validate(flow, from_attributes=True) diff --git a/src/backend/base/langflow/helpers/user.py b/src/backend/base/langflow/helpers/user.py index 9f3488033..e5b956b59 100644 --- a/src/backend/base/langflow/helpers/user.py +++ b/src/backend/base/langflow/helpers/user.py @@ -8,19 +8,19 @@ from langflow.services.database.models.user.model import User, UserRead from langflow.services.deps import get_db_service -def get_user_by_flow_id_or_endpoint_name(flow_id_or_name: str) -> UserRead | None: - with get_db_service().with_session() as session: +async def get_user_by_flow_id_or_endpoint_name(flow_id_or_name: str) -> UserRead | None: + async with get_db_service().with_async_session() as session: try: flow_id = UUID(flow_id_or_name) - flow = session.get(Flow, flow_id) + flow = await session.get(Flow, flow_id) except ValueError: stmt = select(Flow).where(Flow.endpoint_name == flow_id_or_name) - flow = session.exec(stmt).first() + flow = (await session.exec(stmt)).first() if flow is None: raise HTTPException(status_code=404, detail=f"Flow identifier {flow_id_or_name} not found") - user = session.get(User, flow.user_id) + user = await session.get(User, flow.user_id) if user is None: raise HTTPException(status_code=404, detail=f"User for flow {flow_id_or_name} not found") diff --git a/src/backend/base/langflow/services/database/models/transactions/crud.py b/src/backend/base/langflow/services/database/models/transactions/crud.py index 006e6d9bb..3b590d89b 100644 --- a/src/backend/base/langflow/services/database/models/transactions/crud.py +++ b/src/backend/base/langflow/services/database/models/transactions/crud.py @@ -1,7 +1,7 @@ from uuid import UUID from sqlalchemy.exc import IntegrityError -from sqlmodel import Session, col, select +from sqlmodel import col, select from sqlmodel.ext.asyncio.session import AsyncSession from langflow.services.database.models.transactions.model import TransactionBase, TransactionTable @@ -21,12 +21,13 @@ async def get_transactions_by_flow_id( return list(transactions) -def log_transaction(db: Session, transaction: TransactionBase) -> TransactionTable: +async def log_transaction(db: AsyncSession, transaction: TransactionBase) -> TransactionTable: table = TransactionTable(**transaction.model_dump()) db.add(table) try: - db.commit() + await db.commit() + await db.refresh(table) except IntegrityError: - db.rollback() + await db.rollback() raise return table diff --git a/src/backend/base/langflow/services/database/models/vertex_builds/crud.py b/src/backend/base/langflow/services/database/models/vertex_builds/crud.py index 6bd068e96..213ae8502 100644 --- a/src/backend/base/langflow/services/database/models/vertex_builds/crud.py +++ b/src/backend/base/langflow/services/database/models/vertex_builds/crud.py @@ -1,7 +1,7 @@ from uuid import UUID from sqlalchemy.exc import IntegrityError -from sqlmodel import Session, col, delete, select +from sqlmodel import col, delete, select from sqlmodel.ext.asyncio.session import AsyncSession from langflow.services.database.models.vertex_builds.model import VertexBuildBase, VertexBuildTable @@ -21,13 +21,14 @@ async def get_vertex_builds_by_flow_id( return list(builds) -def log_vertex_build(db: Session, vertex_build: VertexBuildBase) -> VertexBuildTable: +async def log_vertex_build(db: AsyncSession, vertex_build: VertexBuildBase) -> VertexBuildTable: table = VertexBuildTable(**vertex_build.model_dump()) db.add(table) try: - db.commit() + await db.commit() + await db.refresh(table) except IntegrityError: - db.rollback() + await db.rollback() raise return table diff --git a/src/backend/base/langflow/services/socket/utils.py b/src/backend/base/langflow/services/socket/utils.py index 95d3a90ba..b8f2ed7d7 100644 --- a/src/backend/base/langflow/services/socket/utils.py +++ b/src/backend/base/langflow/services/socket/utils.py @@ -88,7 +88,7 @@ async def build_vertex( result_dict = ResultDataResponse(results={}) artifacts = {} await set_cache(flow_id, graph) - log_vertex_build( + await log_vertex_build( flow_id=flow_id, vertex_id=vertex_id, valid=valid, diff --git a/src/backend/tests/unit/test_database.py b/src/backend/tests/unit/test_database.py index 77b195904..5d2e2bc61 100644 --- a/src/backend/tests/unit/test_database.py +++ b/src/backend/tests/unit/test_database.py @@ -307,7 +307,7 @@ async def test_delete_flows_with_transaction_and_build(client: AsyncClient, logg "vertex_id": "vid", "flow_id": flow_id, } - log_vertex_build( + await log_vertex_build( flow_id=build["flow_id"], vertex_id=build["vertex_id"], valid=build["valid"], @@ -376,7 +376,7 @@ async def test_delete_folder_with_flows_with_transaction_and_build(client: Async "vertex_id": "vid", "flow_id": flow_id, } - log_vertex_build( + await log_vertex_build( flow_id=build["flow_id"], vertex_id=build["vertex_id"], valid=build["valid"],