ref: Make initialize_database async (#5163)
Make initialize_database async
This commit is contained in:
parent
e545d12c40
commit
63bdcb9d03
10 changed files with 70 additions and 104 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue