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.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")

View file

@ -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,

View file

@ -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)

View file

@ -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")

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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"],