fix: Use AsyncSession in crud log and find_flow (#4691)
Use AsyncSession in crud log and find_flow
This commit is contained in:
parent
ba9dea5547
commit
24f9cac9b5
8 changed files with 37 additions and 34 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue