ref: Make initialize_database async (#5163)

Make initialize_database async
This commit is contained in:
Christophe Bornet 2024-12-10 07:44:34 +01:00 • committed by GitHub
commit 63bdcb9d03
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 70 additions and 104 deletions

View file

@ -5,9 +5,8 @@ import shutil
# we need to import tmpdir
import tempfile
from collections.abc import AsyncGenerator
from contextlib import contextmanager, suppress
from contextlib import suppress
from pathlib import Path
from typing import TYPE_CHECKING
from uuid import UUID
import anyio
@ -27,7 +26,7 @@ from langflow.services.database.models.folder.model import Folder
from langflow.services.database.models.transactions.model import TransactionTable
from langflow.services.database.models.user.model import User, UserCreate, UserRead
from langflow.services.database.models.vertex_builds.crud import delete_vertex_builds_by_flow_id
from langflow.services.database.utils import async_session_getter
from langflow.services.database.utils import session_getter
from langflow.services.deps import get_db_service
from loguru import logger
from sqlalchemy.ext.asyncio import create_async_engine
@ -39,10 +38,6 @@ from typer.testing import CliRunner
from tests.api_keys import get_openai_api_key
if TYPE_CHECKING:
from langflow.services.database.service import DatabaseService
load_dotenv()
@ -369,17 +364,6 @@ async def client_fixture(
await anyio.Path(db_path).unlink()
# create a fixture for session_getter above
@pytest.fixture(name="session_getter")
def session_getter_fixture(client): # noqa: ARG001
@contextmanager
def blank_session_getter(db_service: "DatabaseService"):
with Session(db_service.engine) as session:
yield session
return blank_session_getter
@pytest.fixture
def runner():
return CliRunner()
@ -489,7 +473,7 @@ async def flow(
flow_data = FlowCreate(name="test_flow", data=loaded_json.get("data"), user_id=active_user.id)
flow = Flow.model_validate(flow_data)
async with async_session_getter(get_db_service()) as session:
async with session_getter(get_db_service()) as session:
session.add(flow)
await session.commit()
await session.refresh(flow)
@ -600,7 +584,7 @@ async def created_api_key(active_user):
hashed_api_key=hashed,
)
db_manager = get_db_service()
async with async_session_getter(db_manager) as session:
async with session_getter(db_manager) as session:
stmt = select(ApiKey).where(ApiKey.api_key == api_key.api_key)
if existing_api_key := (await session.exec(stmt)).first():
yield existing_api_key
@ -630,7 +614,7 @@ async def get_simple_api_test(client, logged_in_headers, json_simple_api_test):
@pytest.fixture(name="starter_project")
async def get_starter_project(active_user):
# once the client is created, we can get the starter project
async with async_session_getter(get_db_service()) as session:
async with session_getter(get_db_service()) as session:
stmt = (
select(Flow)
.where(Flow.folder.has(Folder.name == STARTER_FOLDER_NAME))