Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -25,3 +25,5 @@ config.*.yaml
venv
.venv
dist

.claude
12 changes: 9 additions & 3 deletions tests/api/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,9 +31,15 @@ async def app(db_backend_config: DatabaseConfig) -> AsyncGenerator[FastAPI, None
# create tables in the test database
async with app.state.db_engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
yield app
async with app.state.db_engine.begin() as conn:
await conn.run_sync(Base.metadata.drop_all)
try:
yield app

async with app.state.db_engine.begin() as conn:
await conn.run_sync(Base.metadata.drop_all)
finally:
# See tests/scheduler/conftest.py: the database is shared, so the
# engine has to be disposed even when the teardown above fails.
await app.state.db_engine.dispose()


@pytest_asyncio.fixture
Expand Down
23 changes: 14 additions & 9 deletions tests/scheduler/conftest.py
Original file line number Diff line number Diff line change
@@ -1,24 +1,29 @@
"""Pytest fixture and configurations"""

import pytest_asyncio
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
from sqlalchemy.ext.asyncio import async_sessionmaker

from warden.lib.db.database import Base, build_db_url
from warden.lib.db.database import Base, build_engine


@pytest_asyncio.fixture(scope="function")
async def db_engine(config_db):
engine = create_async_engine(build_db_url(config_db.database))
engine = build_engine(config_db.database)

async with engine.begin() as conn:
# Create all tables once
await conn.run_sync(Base.metadata.create_all)
yield engine

async with engine.begin() as conn:
# Delete tables
await conn.run_sync(Base.metadata.drop_all)
await engine.dispose()
try:
yield engine

async with engine.begin() as conn:
# Delete tables
await conn.run_sync(Base.metadata.drop_all)
finally:
# Always dispose: every backend shares one database across the test
# session, so an engine left open on a failing teardown leaks its
# connections - and their transactions - into every later test.
await engine.dispose()


@pytest_asyncio.fixture(scope="function")
Expand Down
6 changes: 3 additions & 3 deletions warden/api/routes/dependencies/db.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,15 @@
from typing import Annotated, AsyncGenerator

from fastapi import Depends, FastAPI, Request
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker

from warden.lib.config import DatabaseConfig
from warden.lib.db.database import build_db_url
from warden.lib.db.database import build_engine


def init_db(app: FastAPI, db_config: DatabaseConfig):
"""Initialize the async engine and session factory with the given DB URL."""
engine = create_async_engine(build_db_url(db_config), echo=db_config.echo)
engine = build_engine(db_config)

# TODO: ensure isolation between concurrent requests
session_factory = async_sessionmaker(bind=engine, expire_on_commit=False)
Expand Down
9 changes: 9 additions & 0 deletions warden/lib/db/database.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
"""Warden db utils"""

from sqlalchemy.engine.url import URL
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlalchemy.orm import declarative_base

from warden.lib.config import DatabaseConfig
Expand Down Expand Up @@ -35,3 +36,11 @@ def build_db_url(cfg: DatabaseConfig) -> str:
).render_as_string(hide_password=False)

raise ValueError(f"Unsupported backend: {cfg.backend}")


def build_engine(cfg: DatabaseConfig) -> AsyncEngine:
"""Build the async engine for `cfg`."""

engine = create_async_engine(build_db_url(cfg), echo=cfg.echo)

return engine
5 changes: 2 additions & 3 deletions warden/scheduler/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,11 +9,10 @@
from sqlalchemy.ext.asyncio import (
AsyncEngine,
async_sessionmaker,
create_async_engine,
)

from warden.lib.config import Config
from warden.lib.db.database import build_db_url
from warden.lib.db.database import build_engine
from warden.lib.models import Job
from warden.scheduler.cancellation_worker import cancellation_worker
from warden.scheduler.db import job_update_commiter
Expand Down Expand Up @@ -144,7 +143,7 @@ async def main_async(conf: Config | None = None):
conf = Config()

logging.config.dictConfig(config=conf.logging)
engine = create_async_engine(build_db_url(conf.database), echo=conf.database.echo)
engine = build_engine(conf.database)
loop = asyncio.get_running_loop()
stop_event = asyncio.Event()

Expand Down
Loading