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

View file

@ -12,7 +12,7 @@ from langflow.initial_setup.setup import load_starter_projects
from langflow.services.database.models.base import orjson_dumps
from langflow.services.database.models.flow import Flow, FlowCreate, FlowUpdate
from langflow.services.database.models.folder.model import FolderCreate
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
@ -530,7 +530,7 @@ async def test_download_file(
]
)
db_manager = get_db_service()
async with async_session_getter(db_manager) as _session:
async with session_getter(db_manager) as _session:
saved_flows = []
for flow in flow_list.flows:
flow.user_id = active_user.id

View file

@ -5,7 +5,7 @@ from httpx import AsyncClient
from langflow.services.auth.utils import create_super_user, get_password_hash
from langflow.services.database.models.user import UserUpdate
from langflow.services.database.models.user.model import User
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, get_settings_service
from sqlmodel import select
@ -14,7 +14,7 @@ from sqlmodel import select
async def super_user(client): # noqa: ARG001
settings_manager = get_settings_service()
auth_settings = settings_manager.auth_settings
async with async_session_getter(get_db_service()) as db:
async with session_getter(get_db_service()) as db:
return await create_super_user(
db=db,
username=auth_settings.SUPERUSER,
@ -42,7 +42,7 @@ async def super_user_headers(
@pytest.fixture
async def deactivated_user(client): # noqa: ARG001
async with async_session_getter(get_db_service()) as session:
async with session_getter(get_db_service()) as session:
user = User(
username="deactivateduser",
password=get_password_hash("testpassword"),
@ -61,7 +61,7 @@ async def test_user_waiting_for_approval(client):
password = "testpassword" # noqa: S105
# Debug: Check if the user already exists
async with async_session_getter(get_db_service()) as session:
async with session_getter(get_db_service()) as session:
stmt = select(User).where(User.username == username)
existing_user = (await session.exec(stmt)).first()
if existing_user:
@ -70,7 +70,7 @@ async def test_user_waiting_for_approval(client):
)
# Create a user that is not active and has never logged in
async with async_session_getter(get_db_service()) as session:
async with session_getter(get_db_service()) as session:
user = User(
username=username,
password=get_password_hash(password),
@ -86,7 +86,7 @@ async def test_user_waiting_for_approval(client):
assert response.json()["detail"] == "Waiting for approval"
# Debug: Check if the user still exists after the test
async with async_session_getter(get_db_service()) as session:
async with session_getter(get_db_service()) as session:
stmt = select(User).where(User.username == username)
existing_user = (await session.exec(stmt)).first()
if existing_user:
@ -140,7 +140,7 @@ async def test_data_consistency_after_delete(client: AsyncClient, test_user, sup
@pytest.mark.api_key_required
async def test_inactive_user(client: AsyncClient):
# Create a user that is not active and has a last_login_at value
async with async_session_getter(get_db_service()) as session:
async with session_getter(get_db_service()) as session:
user = User(
username="inactiveuser",
password=get_password_hash("testpassword"),