fix: Use AsyncSession in crud log and find_flow (#4691)

Use AsyncSession in crud log and find_flow
This commit is contained in:
Christophe Bornet 2024-12-04 16:24:23 +01:00 • committed by GitHub
commit 24f9cac9b5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 37 additions and 34 deletions

View file

@ -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.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.crud import log_vertex_build as crud_log_vertex_build
from langflow.services.database.models.vertex_builds.model import VertexBuildBase 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 from langflow.services.deps import get_db_service, get_settings_service
if TYPE_CHECKING: if TYPE_CHECKING:
@ -157,14 +157,14 @@ async def log_transaction(
error=error, error=error,
flow_id=flow_id if isinstance(flow_id, UUID) else UUID(flow_id), flow_id=flow_id if isinstance(flow_id, UUID) else UUID(flow_id),
) )
with session_getter(get_db_service()) as session: async with async_session_getter(get_db_service()) as session:
inserted = crud_log_transaction(session, transaction) inserted = await crud_log_transaction(session, transaction)
logger.debug(f"Logged transaction: {inserted.id}") logger.debug(f"Logged transaction: {inserted.id}")
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
logger.exception("Error logging transaction") logger.exception("Error logging transaction")
def log_vertex_build( async def log_vertex_build(
*, *,
flow_id: str, flow_id: str,
vertex_id: str, vertex_id: str,
@ -186,8 +186,8 @@ def log_vertex_build(
# ugly hack to get the model dump with weird datatypes # ugly hack to get the model dump with weird datatypes
artifacts=json.loads(json.dumps(artifacts, default=str)), artifacts=json.loads(json.dumps(artifacts, default=str)),
) )
with session_getter(get_db_service()) as session: async with async_session_getter(get_db_service()) as session:
inserted = crud_log_vertex_build(session, vertex_build) inserted = await crud_log_vertex_build(session, vertex_build)
logger.debug(f"Logged vertex build: {inserted.build_id}") logger.debug(f"Logged vertex build: {inserted.build_id}")
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
logger.exception("Error logging vertex build") logger.exception("Error logging vertex build")

View file

@ -435,7 +435,7 @@ class InterfaceVertex(ComponentVertex):
and hasattr(self.custom_component, "store_message") and hasattr(self.custom_component, "store_message")
): ):
self.custom_component.store_message(message) self.custom_component.store_message(message)
log_vertex_build( await log_vertex_build(
flow_id=self.graph.flow_id, flow_id=self.graph.flow_id,
vertex_id=self.id, vertex_id=self.id,
valid=True, valid=True,

View file

@ -10,7 +10,7 @@ from sqlmodel import select
from langflow.schema.schema import INPUT_FIELD_NAME from langflow.schema.schema import INPUT_FIELD_NAME
from langflow.services.database.models.flow import Flow from langflow.services.database.models.flow import Flow
from langflow.services.database.models.flow.model import FlowRead 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: if TYPE_CHECKING:
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
@ -53,13 +53,13 @@ async def load_flow(
msg = "Flow ID or Flow Name is required" msg = "Flow ID or Flow Name is required"
raise ValueError(msg) raise ValueError(msg)
if not flow_id and flow_name: 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: if not flow_id:
msg = f"Flow {flow_name} not found" msg = f"Flow {flow_name} not found"
raise ValueError(msg) raise ValueError(msg)
with session_scope() as session: async with async_session_scope() as session:
graph_data = flow.data if (flow := session.get(Flow, flow_id)) else None graph_data = flow.data if (flow := await session.get(Flow, flow_id)) else None
if not graph_data: if not graph_data:
msg = f"Flow {flow_id} not found" msg = f"Flow {flow_id} not found"
raise ValueError(msg) 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) 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: async def find_flow(flow_name: str, user_id: str) -> str | None:
with session_scope() as session: async with async_session_scope() as session:
flow = session.exec(select(Flow).where(Flow.name == flow_name).where(Flow.user_id == user_id)).first() 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 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: async 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 with async_session_scope() as session:
endpoint_name = None endpoint_name = None
try: try:
flow_id = UUID(flow_id_or_name) flow_id = UUID(flow_id_or_name)
flow = session.get(Flow, flow_id) flow = await session.get(Flow, flow_id)
except ValueError: except ValueError:
endpoint_name = flow_id_or_name endpoint_name = flow_id_or_name
stmt = select(Flow).where(Flow.endpoint_name == endpoint_name) stmt = select(Flow).where(Flow.endpoint_name == endpoint_name)
if user_id: if user_id:
stmt = stmt.where(Flow.user_id == 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: if flow is None:
raise HTTPException(status_code=404, detail=f"Flow identifier {flow_id_or_name} not found") raise HTTPException(status_code=404, detail=f"Flow identifier {flow_id_or_name} not found")
return FlowRead.model_validate(flow, from_attributes=True) return FlowRead.model_validate(flow, from_attributes=True)

View file

@ -8,19 +8,19 @@ from langflow.services.database.models.user.model import User, UserRead
from langflow.services.deps import get_db_service from langflow.services.deps import get_db_service
def get_user_by_flow_id_or_endpoint_name(flow_id_or_name: str) -> UserRead | None: async 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 with get_db_service().with_async_session() as session:
try: try:
flow_id = UUID(flow_id_or_name) flow_id = UUID(flow_id_or_name)
flow = session.get(Flow, flow_id) flow = await session.get(Flow, flow_id)
except ValueError: except ValueError:
stmt = select(Flow).where(Flow.endpoint_name == flow_id_or_name) 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: if flow is None:
raise HTTPException(status_code=404, detail=f"Flow identifier {flow_id_or_name} not found") 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: if user is None:
raise HTTPException(status_code=404, detail=f"User for flow {flow_id_or_name} not found") raise HTTPException(status_code=404, detail=f"User for flow {flow_id_or_name} not found")

View file

@ -1,7 +1,7 @@
from uuid import UUID from uuid import UUID
from sqlalchemy.exc import IntegrityError from sqlalchemy.exc import IntegrityError
from sqlmodel import Session, col, select from sqlmodel import col, select
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
from langflow.services.database.models.transactions.model import TransactionBase, TransactionTable from langflow.services.database.models.transactions.model import TransactionBase, TransactionTable
@ -21,12 +21,13 @@ async def get_transactions_by_flow_id(
return list(transactions) 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()) table = TransactionTable(**transaction.model_dump())
db.add(table) db.add(table)
try: try:
db.commit() await db.commit()
await db.refresh(table)
except IntegrityError: except IntegrityError:
db.rollback() await db.rollback()
raise raise
return table return table

View file

@ -1,7 +1,7 @@
from uuid import UUID from uuid import UUID
from sqlalchemy.exc import IntegrityError 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 sqlmodel.ext.asyncio.session import AsyncSession
from langflow.services.database.models.vertex_builds.model import VertexBuildBase, VertexBuildTable 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) 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()) table = VertexBuildTable(**vertex_build.model_dump())
db.add(table) db.add(table)
try: try:
db.commit() await db.commit()
await db.refresh(table)
except IntegrityError: except IntegrityError:
db.rollback() await db.rollback()
raise raise
return table return table

View file

@ -88,7 +88,7 @@ async def build_vertex(
result_dict = ResultDataResponse(results={}) result_dict = ResultDataResponse(results={})
artifacts = {} artifacts = {}
await set_cache(flow_id, graph) await set_cache(flow_id, graph)
log_vertex_build( await log_vertex_build(
flow_id=flow_id, flow_id=flow_id,
vertex_id=vertex_id, vertex_id=vertex_id,
valid=valid, valid=valid,

View file

@ -307,7 +307,7 @@ async def test_delete_flows_with_transaction_and_build(client: AsyncClient, logg
"vertex_id": "vid", "vertex_id": "vid",
"flow_id": flow_id, "flow_id": flow_id,
} }
log_vertex_build( await log_vertex_build(
flow_id=build["flow_id"], flow_id=build["flow_id"],
vertex_id=build["vertex_id"], vertex_id=build["vertex_id"],
valid=build["valid"], valid=build["valid"],
@ -376,7 +376,7 @@ async def test_delete_folder_with_flows_with_transaction_and_build(client: Async
"vertex_id": "vid", "vertex_id": "vid",
"flow_id": flow_id, "flow_id": flow_id,
} }
log_vertex_build( await log_vertex_build(
flow_id=build["flow_id"], flow_id=build["flow_id"],
vertex_id=build["vertex_id"], vertex_id=build["vertex_id"],
valid=build["valid"], valid=build["valid"],