generated from john/python-template
Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ec6617a1c4 | ||
|
|
9eb0f40c08 | ||
|
|
f769d29da1 | ||
|
|
8afc462a6d | ||
|
|
3d6daec561 | ||
|
|
1cc2f319d5 |
@@ -0,0 +1,20 @@
|
|||||||
|
import uvicorn
|
||||||
|
|
||||||
|
from .config import LOGGING_CONFIG
|
||||||
|
from .config import get_settings
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
settings = get_settings()
|
||||||
|
uvicorn.run(
|
||||||
|
"transcription.app:create_app",
|
||||||
|
factory=True,
|
||||||
|
host=settings.host,
|
||||||
|
port=settings.port,
|
||||||
|
log_level=LOGGING_CONFIG.get("root", {}).get("level", "info").lower(),
|
||||||
|
reload=settings.reload,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -16,11 +16,14 @@ from fastapi.staticfiles import StaticFiles
|
|||||||
|
|
||||||
from .api.errors import register_error_handlers
|
from .api.errors import register_error_handlers
|
||||||
from .api.health import router as health_router
|
from .api.health import router as health_router
|
||||||
|
from .config import Settings
|
||||||
from .config import configure_logging
|
from .config import configure_logging
|
||||||
from .config import get_settings
|
from .config import get_settings
|
||||||
from .db import create_all
|
from .db import create_all
|
||||||
from .db import dispose_database_runtime
|
|
||||||
from .db import initialize_database_runtime
|
from .db import initialize_database_runtime
|
||||||
|
from .db.engine import get_database_url
|
||||||
|
from .db.engine import resolve_engine
|
||||||
|
from .db.session import dispose_session_factory
|
||||||
from .services import ServiceBundle
|
from .services import ServiceBundle
|
||||||
from .services.jobs import JobService
|
from .services.jobs import JobService
|
||||||
from .ui import register_pages
|
from .ui import register_pages
|
||||||
@@ -39,7 +42,7 @@ async def _lifespan(app: FastAPI):
|
|||||||
app.state.runtime = initialize_database_runtime(settings=settings)
|
app.state.runtime = initialize_database_runtime(settings=settings)
|
||||||
|
|
||||||
if settings.should_bootstrap_schema:
|
if settings.should_bootstrap_schema:
|
||||||
await create_all(engine=app.state.runtime.engine)
|
await create_all(engine=resolve_engine(settings=settings))
|
||||||
|
|
||||||
settings.upload_dir.mkdir(parents=True, exist_ok=True)
|
settings.upload_dir.mkdir(parents=True, exist_ok=True)
|
||||||
settings.prompt_dir.mkdir(parents=True, exist_ok=True)
|
settings.prompt_dir.mkdir(parents=True, exist_ok=True)
|
||||||
@@ -47,7 +50,10 @@ async def _lifespan(app: FastAPI):
|
|||||||
await _recover_stale_processing_jobs(app)
|
await _recover_stale_processing_jobs(app)
|
||||||
|
|
||||||
async with AsyncExitStack() as stack:
|
async with AsyncExitStack() as stack:
|
||||||
stack.push_async_callback(dispose_database_runtime)
|
stack.push_async_callback(
|
||||||
|
dispose_session_factory,
|
||||||
|
database_url=get_database_url(settings),
|
||||||
|
)
|
||||||
stop_event, worker_notifier = await stack.enter_async_context(
|
stop_event, worker_notifier = await stack.enter_async_context(
|
||||||
worker_consumer_lifespan(
|
worker_consumer_lifespan(
|
||||||
session_factory=app.state.runtime.session_factory,
|
session_factory=app.state.runtime.session_factory,
|
||||||
@@ -73,14 +79,14 @@ async def _recover_stale_processing_jobs(app: FastAPI) -> None:
|
|||||||
logger.warning("Recovered %s stale processing job(s) at startup", recovered)
|
logger.warning("Recovered %s stale processing job(s) at startup", recovered)
|
||||||
|
|
||||||
|
|
||||||
def create_app() -> FastAPI:
|
def create_app(settings: Settings | None = None) -> FastAPI:
|
||||||
"""Create and configure the FastAPI application."""
|
"""Create and configure the FastAPI application."""
|
||||||
app = FastAPI(title="Transcription", lifespan=_lifespan)
|
app = FastAPI(title="Transcription", lifespan=_lifespan)
|
||||||
settings = get_settings()
|
active_settings = settings or get_settings()
|
||||||
app.state.settings = settings
|
app.state.settings = active_settings
|
||||||
app.mount(
|
app.mount(
|
||||||
"/uploads",
|
"/uploads",
|
||||||
StaticFiles(directory=settings.upload_dir, check_dir=False),
|
StaticFiles(directory=active_settings.upload_dir, check_dir=False),
|
||||||
name="uploads",
|
name="uploads",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -92,6 +98,10 @@ def create_app() -> FastAPI:
|
|||||||
async def ui_redirect() -> RedirectResponse:
|
async def ui_redirect() -> RedirectResponse:
|
||||||
return RedirectResponse(url="/ui/upload", status_code=status.HTTP_307_TEMPORARY_REDIRECT)
|
return RedirectResponse(url="/ui/upload", status_code=status.HTTP_307_TEMPORARY_REDIRECT)
|
||||||
|
|
||||||
|
@app.get("/healthz")
|
||||||
|
def health() -> dict[str, str]:
|
||||||
|
return {"status": "ok"}
|
||||||
|
|
||||||
register_error_handlers(app)
|
register_error_handlers(app)
|
||||||
register_pages(app)
|
register_pages(app)
|
||||||
app.include_router(health_router)
|
app.include_router(health_router)
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from sqlalchemy.ext.asyncio import async_sessionmaker
|
|||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
from transcription.db.runtime import DatabaseRuntime
|
from transcription.db.runtime import DatabaseRuntime
|
||||||
from transcription.db.runtime import get_session_factory
|
from transcription.db.session import get_session_factory
|
||||||
from transcription.worker import WorkerNotifier
|
from transcription.worker import WorkerNotifier
|
||||||
from transcription.worker import resolve_worker_notifier
|
from transcription.worker import resolve_worker_notifier
|
||||||
|
|
||||||
|
|||||||
+44
-13
@@ -6,12 +6,16 @@ are resolved by the provider adapters, not here.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import logging.config
|
import logging.config
|
||||||
from contextvars import ContextVar
|
|
||||||
from enum import StrEnum
|
from enum import StrEnum
|
||||||
|
from functools import cache
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Annotated
|
||||||
|
from typing import Any
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
|
|
||||||
|
from pydantic import BaseModel
|
||||||
from pydantic import Field
|
from pydantic import Field
|
||||||
|
from pydantic import SecretStr
|
||||||
from pydantic_settings import BaseSettings
|
from pydantic_settings import BaseSettings
|
||||||
from pydantic_settings import SettingsConfigDict
|
from pydantic_settings import SettingsConfigDict
|
||||||
|
|
||||||
@@ -22,13 +26,42 @@ class Provider(StrEnum):
|
|||||||
OPENROUTER = "openrouter"
|
OPENROUTER = "openrouter"
|
||||||
|
|
||||||
|
|
||||||
|
class SqliteSettings(BaseModel):
|
||||||
|
driver: Literal["sqlite"] = "sqlite"
|
||||||
|
path: str = "app.db"
|
||||||
|
|
||||||
|
|
||||||
|
class PostgresSettings(BaseModel):
|
||||||
|
driver: Literal["postgres"] = "postgres"
|
||||||
|
host: str
|
||||||
|
port: int = 5432
|
||||||
|
database: str
|
||||||
|
user: str
|
||||||
|
password: SecretStr
|
||||||
|
|
||||||
|
|
||||||
|
DatabaseSettings = Annotated[
|
||||||
|
SqliteSettings | PostgresSettings,
|
||||||
|
Field(discriminator="driver"),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
class Settings(BaseSettings):
|
class Settings(BaseSettings):
|
||||||
model_config = SettingsConfigDict(
|
model_config = SettingsConfigDict(
|
||||||
env_file=".env",
|
env_file=".env",
|
||||||
env_file_encoding="utf-8",
|
env_file_encoding="utf-8",
|
||||||
extra="ignore",
|
extra="ignore",
|
||||||
|
cli_parse_args=True,
|
||||||
|
cli_implicit_flags=True,
|
||||||
|
cli_kebab_case=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# --- NiceGUI Server ---
|
||||||
|
host: str = "0.0.0.0"
|
||||||
|
port: int = 8000
|
||||||
|
log_level: Literal["critical", "error", "warning", "info", "debug", "trace"] = "info"
|
||||||
|
reload: bool = False
|
||||||
|
|
||||||
# --- AI provider ---
|
# --- AI provider ---
|
||||||
provider: Provider = Provider.OPENROUTER
|
provider: Provider = Provider.OPENROUTER
|
||||||
openrouter_api_key: str
|
openrouter_api_key: str
|
||||||
@@ -40,8 +73,9 @@ class Settings(BaseSettings):
|
|||||||
environment: Literal["development", "test", "production"] = "development"
|
environment: Literal["development", "test", "production"] = "development"
|
||||||
|
|
||||||
# --- persistence ---
|
# --- persistence ---
|
||||||
|
database: DatabaseSettings = Field(default_factory=SqliteSettings)
|
||||||
database_url: str = "sqlite:///./transcription.db"
|
database_url: str = "sqlite:///./transcription.db"
|
||||||
bootstrap_schema_on_startup: bool | None = None
|
bootstrap_schema_on_startup: bool = False
|
||||||
sqlite_check_same_thread: bool = False
|
sqlite_check_same_thread: bool = False
|
||||||
|
|
||||||
# --- filesystem paths ---
|
# --- filesystem paths ---
|
||||||
@@ -64,18 +98,12 @@ class Settings(BaseSettings):
|
|||||||
return self.environment in {"development", "test"}
|
return self.environment in {"development", "test"}
|
||||||
|
|
||||||
|
|
||||||
_settings: ContextVar[Settings | None] = ContextVar("settings", default=None)
|
@cache
|
||||||
|
|
||||||
|
|
||||||
def get_settings(**kwargs) -> Settings:
|
def get_settings(**kwargs) -> Settings:
|
||||||
settings = _settings.get()
|
return Settings(**kwargs)
|
||||||
if settings is None:
|
|
||||||
settings = Settings(**kwargs) # pyright: ignore[reportCallIssue]
|
|
||||||
_settings.set(settings)
|
|
||||||
return settings
|
|
||||||
|
|
||||||
|
|
||||||
LOGGING_CONFIG: dict[str, object] = {
|
LOGGING_CONFIG: dict[str, Any] = {
|
||||||
"version": 1,
|
"version": 1,
|
||||||
"disable_existing_loggers": False,
|
"disable_existing_loggers": False,
|
||||||
"formatters": {
|
"formatters": {
|
||||||
@@ -105,7 +133,10 @@ LOGGING_CONFIG: dict[str, object] = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def configure_logging() -> None:
|
def configure_logging(settings: Settings | None = None) -> None:
|
||||||
"""Configure root logging once at startup."""
|
"""Configure root logging once at startup."""
|
||||||
logging.config.dictConfig(LOGGING_CONFIG)
|
cfg = LOGGING_CONFIG.copy()
|
||||||
|
active_settings = settings or get_settings()
|
||||||
|
cfg["loggers"]["transcription"]["level"] = active_settings.log_level.upper()
|
||||||
|
logging.config.dictConfig(cfg)
|
||||||
logger.debug("Logging configured")
|
logger.debug("Logging configured")
|
||||||
|
|||||||
@@ -1,6 +1,13 @@
|
|||||||
from .operations import create_all
|
from .operations import create_all
|
||||||
from .runtime import dispose_database_runtime
|
from .runtime import dispose_database_runtime
|
||||||
from .runtime import get_session
|
|
||||||
from .runtime import initialize_database_runtime
|
from .runtime import initialize_database_runtime
|
||||||
|
from .session import session_scope
|
||||||
|
from .session import transaction_scope
|
||||||
|
|
||||||
__all__ = ["create_all", "dispose_database_runtime", "get_session", "initialize_database_runtime"]
|
__all__ = [
|
||||||
|
"create_all",
|
||||||
|
"dispose_database_runtime",
|
||||||
|
"initialize_database_runtime",
|
||||||
|
"session_scope",
|
||||||
|
"transaction_scope",
|
||||||
|
]
|
||||||
|
|||||||
@@ -0,0 +1,60 @@
|
|||||||
|
from functools import cache
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy import URL
|
||||||
|
from sqlalchemy import StaticPool
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||||
|
from sqlalchemy.ext.asyncio import create_async_engine
|
||||||
|
|
||||||
|
from ..config import PostgresSettings
|
||||||
|
from ..config import Settings
|
||||||
|
from ..config import SqliteSettings
|
||||||
|
from ..config import get_settings
|
||||||
|
|
||||||
|
|
||||||
|
def get_database_url(settings: Settings) -> str:
|
||||||
|
match settings.database:
|
||||||
|
case SqliteSettings(path=path):
|
||||||
|
url = URL.create(
|
||||||
|
drivername="sqlite+aiosqlite",
|
||||||
|
database=path,
|
||||||
|
)
|
||||||
|
case PostgresSettings() as database:
|
||||||
|
url = URL.create(
|
||||||
|
drivername="postgresql+asyncpg",
|
||||||
|
host=database.host,
|
||||||
|
port=database.port,
|
||||||
|
database=database.database,
|
||||||
|
username=database.user,
|
||||||
|
password=database.password.get_secret_value(),
|
||||||
|
)
|
||||||
|
return url.render_as_string(hide_password=False)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_engine(settings: Settings | None = None) -> AsyncEngine:
|
||||||
|
active_settings = settings or get_settings()
|
||||||
|
return get_engine(get_database_url(active_settings))
|
||||||
|
|
||||||
|
|
||||||
|
@cache
|
||||||
|
def get_engine(database_url: str) -> AsyncEngine:
|
||||||
|
kwargs: dict[str, Any] = {"echo": False, "pool_pre_ping": True}
|
||||||
|
if database_url.startswith("sqlite"):
|
||||||
|
kwargs["connect_args"] = {"check_same_thread": False}
|
||||||
|
if ":memory:" in database_url:
|
||||||
|
kwargs["poolclass"] = StaticPool
|
||||||
|
|
||||||
|
return create_async_engine(database_url, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
async def dispose_engine(database_url: str) -> None:
|
||||||
|
engine = get_engine(database_url)
|
||||||
|
try:
|
||||||
|
await engine.dispose()
|
||||||
|
finally:
|
||||||
|
get_engine.cache_clear()
|
||||||
|
|
||||||
|
|
||||||
|
async def refresh_engine(database_url: str) -> AsyncEngine:
|
||||||
|
await dispose_engine(database_url)
|
||||||
|
return get_engine(database_url)
|
||||||
@@ -10,13 +10,25 @@ from sqlmodel import SQLModel
|
|||||||
from sqlmodel import select
|
from sqlmodel import select
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
from ..models import Job
|
from .engine import resolve_engine
|
||||||
from ..models import JobStatus
|
from .models import Job
|
||||||
from .runtime import get_engine
|
from .models import JobStatus
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
async def create_all(*, engine: AsyncEngine | None = None) -> None:
|
||||||
|
"""Create all tables on the selected engine."""
|
||||||
|
# Import models so SQLModel metadata is fully registered before bootstrap.
|
||||||
|
from transcription.db import models as _models # noqa: F401
|
||||||
|
|
||||||
|
active_engine = engine or resolve_engine()
|
||||||
|
async with active_engine.begin() as connection:
|
||||||
|
await connection.run_sync(SQLModel.metadata.create_all)
|
||||||
|
await connection.run_sync(_ensure_sqlite_compat_columns)
|
||||||
|
logger.debug("Database schema bootstrap complete for database_url=%s", active_engine.url)
|
||||||
|
|
||||||
|
|
||||||
async def get_next_queued_job(*, session: AsyncSession) -> Job | None:
|
async def get_next_queued_job(*, session: AsyncSession) -> Job | None:
|
||||||
"""Get the next queued job, if any."""
|
"""Get the next queued job, if any."""
|
||||||
result = await session.exec(
|
result = await session.exec(
|
||||||
@@ -28,18 +40,6 @@ async def get_next_queued_job(*, session: AsyncSession) -> Job | None:
|
|||||||
return result.first()
|
return result.first()
|
||||||
|
|
||||||
|
|
||||||
async def create_all(*, engine: AsyncEngine | None = None) -> None:
|
|
||||||
"""Create all tables on the selected engine."""
|
|
||||||
# Import models so SQLModel metadata is fully registered before bootstrap.
|
|
||||||
from transcription import models as _models # noqa: F401
|
|
||||||
|
|
||||||
active_engine = engine or get_engine()
|
|
||||||
async with active_engine.begin() as connection:
|
|
||||||
await connection.run_sync(SQLModel.metadata.create_all)
|
|
||||||
await connection.run_sync(_ensure_sqlite_compat_columns)
|
|
||||||
logger.debug("Database schema bootstrap complete for database_url=%s", active_engine.url)
|
|
||||||
|
|
||||||
|
|
||||||
def _ensure_sqlite_compat_columns(connection: Connection) -> None:
|
def _ensure_sqlite_compat_columns(connection: Connection) -> None:
|
||||||
"""Apply lightweight dev/test SQLite compatibility column patches.
|
"""Apply lightweight dev/test SQLite compatibility column patches.
|
||||||
|
|
||||||
@@ -68,12 +68,8 @@ def _ensure_sqlite_compat_columns(connection: Connection) -> None:
|
|||||||
break
|
break
|
||||||
if not has_unique_source:
|
if not has_unique_source:
|
||||||
connection.execute(
|
connection.execute(
|
||||||
text(
|
text("CREATE UNIQUE INDEX IF NOT EXISTS ux_revision_source_id ON revision(source_id)")
|
||||||
"CREATE UNIQUE INDEX IF NOT EXISTS "
|
|
||||||
"ux_revision_source_id ON revision(source_id)"
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Applied SQLite compatibility schema patch "
|
"Applied SQLite compatibility schema patch table=revision unique_index=ux_revision_source_id"
|
||||||
"table=revision unique_index=ux_revision_source_id"
|
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,18 +1,16 @@
|
|||||||
import logging
|
import logging
|
||||||
from collections.abc import AsyncGenerator
|
|
||||||
from contextlib import asynccontextmanager
|
|
||||||
from contextvars import ContextVar
|
from contextvars import ContextVar
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from functools import partial
|
|
||||||
|
|
||||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||||
from sqlalchemy.ext.asyncio import async_sessionmaker
|
from sqlalchemy.ext.asyncio import async_sessionmaker
|
||||||
from sqlalchemy.ext.asyncio import create_async_engine
|
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
from sqlmodel.pool import StaticPool
|
|
||||||
|
|
||||||
from ..config import Settings
|
from ..config import Settings
|
||||||
from ..config import get_settings
|
from ..config import get_settings
|
||||||
|
from .engine import get_database_url
|
||||||
|
from .engine import get_engine
|
||||||
|
from .session import get_session_factory
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -37,33 +35,6 @@ async def dispose_database_runtime() -> None:
|
|||||||
_runtime.set(None)
|
_runtime.set(None)
|
||||||
|
|
||||||
|
|
||||||
def _to_async_database_url(database_url: str) -> str:
|
|
||||||
"""Normalize configured database URL to an async SQLAlchemy driver URL."""
|
|
||||||
if database_url.startswith("sqlite://") and not database_url.startswith("sqlite+aiosqlite://"):
|
|
||||||
return database_url.replace("sqlite://", "sqlite+aiosqlite://", 1)
|
|
||||||
if database_url.startswith("postgresql://") and not database_url.startswith("postgresql+asyncpg://"):
|
|
||||||
return database_url.replace("postgresql://", "postgresql+asyncpg://", 1)
|
|
||||||
return database_url
|
|
||||||
|
|
||||||
|
|
||||||
def _build_engine(settings: Settings) -> AsyncEngine:
|
|
||||||
database_url = _to_async_database_url(settings.database_url)
|
|
||||||
engine_factory = partial(
|
|
||||||
create_async_engine,
|
|
||||||
url=database_url,
|
|
||||||
echo=False,
|
|
||||||
pool_pre_ping=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
if database_url.startswith("sqlite"):
|
|
||||||
sqlite_connect_settings = {"check_same_thread": settings.sqlite_check_same_thread}
|
|
||||||
engine_factory = partial(engine_factory, connect_args=sqlite_connect_settings)
|
|
||||||
if ":memory:" in database_url:
|
|
||||||
engine_factory = partial(engine_factory, poolclass=StaticPool)
|
|
||||||
|
|
||||||
return engine_factory()
|
|
||||||
|
|
||||||
|
|
||||||
def initialize_database_runtime(*, settings: Settings | None = None) -> DatabaseRuntime:
|
def initialize_database_runtime(*, settings: Settings | None = None) -> DatabaseRuntime:
|
||||||
"""Initialize lifespan-owned async DB resources once per process."""
|
"""Initialize lifespan-owned async DB resources once per process."""
|
||||||
runtime = _runtime.get()
|
runtime = _runtime.get()
|
||||||
@@ -71,33 +42,10 @@ def initialize_database_runtime(*, settings: Settings | None = None) -> Database
|
|||||||
return runtime
|
return runtime
|
||||||
|
|
||||||
active_settings = settings or get_settings()
|
active_settings = settings or get_settings()
|
||||||
engine = _build_engine(active_settings)
|
database_url = get_database_url(active_settings)
|
||||||
session_factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)
|
engine = get_engine(database_url)
|
||||||
|
session_factory = get_session_factory(database_url)
|
||||||
runtime = DatabaseRuntime(engine=engine, session_factory=session_factory)
|
runtime = DatabaseRuntime(engine=engine, session_factory=session_factory)
|
||||||
_runtime.set(runtime)
|
_runtime.set(runtime)
|
||||||
logger.debug("Initialized async database runtime for database_url=%s", engine.url)
|
logger.debug("Initialized async database runtime for database_url=%s", engine.url)
|
||||||
return runtime
|
return runtime
|
||||||
|
|
||||||
|
|
||||||
def get_engine(settings: Settings | None = None) -> AsyncEngine:
|
|
||||||
"""Return the current async SQLAlchemy engine."""
|
|
||||||
runtime = _runtime.get() or initialize_database_runtime(settings=settings)
|
|
||||||
return runtime.engine
|
|
||||||
|
|
||||||
|
|
||||||
def get_session_factory(settings: Settings | None = None) -> async_sessionmaker[AsyncSession]:
|
|
||||||
"""Return the shared async session factory."""
|
|
||||||
runtime = _runtime.get() or initialize_database_runtime(settings=settings)
|
|
||||||
return runtime.session_factory
|
|
||||||
|
|
||||||
|
|
||||||
@asynccontextmanager
|
|
||||||
async def get_session(
|
|
||||||
*,
|
|
||||||
settings: Settings | None = None,
|
|
||||||
session_factory: async_sessionmaker[AsyncSession] | None = None,
|
|
||||||
) -> AsyncGenerator[AsyncSession]:
|
|
||||||
"""Yield a database session and ensure cleanup."""
|
|
||||||
active_session_factory = session_factory or get_session_factory(settings)
|
|
||||||
async with active_session_factory() as session:
|
|
||||||
yield session
|
|
||||||
|
|||||||
@@ -0,0 +1,79 @@
|
|||||||
|
from collections.abc import AsyncGenerator
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
|
from functools import cache
|
||||||
|
from typing import Annotated
|
||||||
|
|
||||||
|
from fastapi import Depends
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSessionTransaction
|
||||||
|
from sqlalchemy.ext.asyncio import async_sessionmaker
|
||||||
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
from ..config import get_settings
|
||||||
|
from .engine import dispose_engine
|
||||||
|
from .engine import get_database_url
|
||||||
|
from .engine import get_engine
|
||||||
|
|
||||||
|
type SessionFactory = async_sessionmaker[AsyncSession]
|
||||||
|
|
||||||
|
|
||||||
|
@cache
|
||||||
|
def get_session_factory(database_url: str) -> SessionFactory:
|
||||||
|
return async_sessionmaker(
|
||||||
|
bind=get_engine(database_url),
|
||||||
|
class_=AsyncSession,
|
||||||
|
expire_on_commit=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_session_factory(database_url: str | None = None) -> SessionFactory:
|
||||||
|
return get_session_factory(database_url or get_database_url(get_settings()))
|
||||||
|
|
||||||
|
|
||||||
|
type SessionFactoryDep = Annotated[SessionFactory, Depends(resolve_session_factory)]
|
||||||
|
|
||||||
|
|
||||||
|
async def dispose_session_factory(database_url: str) -> None:
|
||||||
|
get_session_factory.cache_clear()
|
||||||
|
await dispose_engine(database_url)
|
||||||
|
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def session_scope(
|
||||||
|
*,
|
||||||
|
database_url: str | None = None,
|
||||||
|
session: AsyncSession | None = None,
|
||||||
|
) -> AsyncGenerator[AsyncSession]:
|
||||||
|
if session is not None:
|
||||||
|
yield session
|
||||||
|
return
|
||||||
|
|
||||||
|
session_factory = resolve_session_factory(database_url)
|
||||||
|
async with session_factory() as owned_session:
|
||||||
|
yield owned_session
|
||||||
|
|
||||||
|
|
||||||
|
type SessionScopeDep = Annotated[AsyncSession, Depends(session_scope)]
|
||||||
|
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def transaction_scope(
|
||||||
|
*,
|
||||||
|
database_url: str | None = None,
|
||||||
|
session: AsyncSessionTransaction | None = None,
|
||||||
|
) -> AsyncGenerator[AsyncSessionTransaction]:
|
||||||
|
match session:
|
||||||
|
case AsyncSession() as async_session:
|
||||||
|
if not async_session.in_transaction():
|
||||||
|
raise RuntimeError("A supplied session must have an active transaction")
|
||||||
|
yield async_session
|
||||||
|
return
|
||||||
|
case AsyncSessionTransaction() as async_transaction:
|
||||||
|
yield async_transaction
|
||||||
|
return
|
||||||
|
|
||||||
|
session_factory = resolve_session_factory(database_url)
|
||||||
|
async with session_factory().begin() as owned_session:
|
||||||
|
yield owned_session
|
||||||
|
|
||||||
|
|
||||||
|
type TransactionScopeDep = Annotated[AsyncSessionTransaction, Depends(transaction_scope)]
|
||||||
@@ -8,7 +8,8 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
|||||||
|
|
||||||
from ..config import Settings
|
from ..config import Settings
|
||||||
from ..config import get_settings
|
from ..config import get_settings
|
||||||
from ..db.runtime import get_session_factory
|
from ..db.session import resolve_session_factory
|
||||||
|
from ..db.session import session_scope
|
||||||
|
|
||||||
|
|
||||||
class ServiceBase(ABC):
|
class ServiceBase(ABC):
|
||||||
@@ -24,19 +25,14 @@ class ServiceBase(ABC):
|
|||||||
queue: asyncio.Queue | None = None,
|
queue: asyncio.Queue | None = None,
|
||||||
):
|
):
|
||||||
self.settings = get_settings()
|
self.settings = get_settings()
|
||||||
self.session_factory = session_factory or get_session_factory()
|
self.session_factory = session_factory or resolve_session_factory()
|
||||||
self.queue = queue or asyncio.Queue()
|
self.queue = queue or asyncio.Queue()
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def _session_scope(self, session: AsyncSession | None = None):
|
async def _session_scope(self, session: AsyncSession | None = None):
|
||||||
"""Provide a transactional scope around a series of operations."""
|
"""Provide a transactional scope around a series of operations."""
|
||||||
if session is not None:
|
async with session_scope(session=session) as active_session:
|
||||||
# Reuse the provided session if one is passed in
|
yield active_session
|
||||||
yield session
|
|
||||||
else:
|
|
||||||
# Otherwise, create a new session for this scope
|
|
||||||
async with self.session_factory() as new_session:
|
|
||||||
yield new_session
|
|
||||||
|
|
||||||
async def _finalize(
|
async def _finalize(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -9,9 +9,9 @@ from sqlalchemy.orm import selectinload
|
|||||||
from sqlmodel import select
|
from sqlmodel import select
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
from ..db.models import Document
|
||||||
from ..errors import AppError
|
from ..errors import AppError
|
||||||
from ..errors import ErrorCategory
|
from ..errors import ErrorCategory
|
||||||
from ..models import Document
|
|
||||||
from .base import ServiceBase
|
from .base import ServiceBase
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|||||||
@@ -7,9 +7,9 @@ from sqlalchemy.orm import selectinload
|
|||||||
from sqlmodel import select
|
from sqlmodel import select
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
from ..models import Job
|
from ..db.models import Job
|
||||||
from ..models import JobStatus
|
from ..db.models import JobStatus
|
||||||
from ..models import Source
|
from ..db.models import Source
|
||||||
from .base import ServiceBase
|
from .base import ServiceBase
|
||||||
|
|
||||||
|
|
||||||
@@ -170,11 +170,7 @@ class JobService(ServiceBase):
|
|||||||
``stale_before`` are considered stale and re-queued.
|
``stale_before`` are considered stale and re-queued.
|
||||||
"""
|
"""
|
||||||
async with self._session_scope(session) as _session:
|
async with self._session_scope(session) as _session:
|
||||||
query = (
|
query = select(Job).where(Job.status == JobStatus.PROCESSING).where(Job.date_updated < stale_before)
|
||||||
select(Job)
|
|
||||||
.where(Job.status == JobStatus.PROCESSING)
|
|
||||||
.where(Job.date_updated < stale_before)
|
|
||||||
)
|
|
||||||
stale_jobs = (await _session.exec(query)).all()
|
stale_jobs = (await _session.exec(query)).all()
|
||||||
if not stale_jobs:
|
if not stale_jobs:
|
||||||
return 0
|
return 0
|
||||||
|
|||||||
@@ -11,9 +11,9 @@ from transcription.config import get_settings
|
|||||||
from transcription.errors import AppError
|
from transcription.errors import AppError
|
||||||
from transcription.errors import ErrorCategory
|
from transcription.errors import ErrorCategory
|
||||||
|
|
||||||
from ..models import Document
|
from ..db.models import Document
|
||||||
from ..models import Job
|
from ..db.models import Job
|
||||||
from ..models import Source
|
from ..db.models import Source
|
||||||
from .documents import UploadJobResult
|
from .documents import UploadJobResult
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|||||||
@@ -18,11 +18,11 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
|||||||
|
|
||||||
from transcription.config import Settings
|
from transcription.config import Settings
|
||||||
from transcription.config import get_settings
|
from transcription.config import get_settings
|
||||||
|
from transcription.db.models import Job
|
||||||
|
from transcription.db.models import Revision
|
||||||
|
from transcription.db.models import Source
|
||||||
from transcription.errors import AppError
|
from transcription.errors import AppError
|
||||||
from transcription.errors import ErrorCategory
|
from transcription.errors import ErrorCategory
|
||||||
from transcription.models import Job
|
|
||||||
from transcription.models import Revision
|
|
||||||
from transcription.models import Source
|
|
||||||
from transcription.providers import ProviderAuthError
|
from transcription.providers import ProviderAuthError
|
||||||
from transcription.providers import ProviderError
|
from transcription.providers import ProviderError
|
||||||
from transcription.providers import ProviderResponseError
|
from transcription.providers import ProviderResponseError
|
||||||
|
|||||||
@@ -5,13 +5,13 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
|||||||
|
|
||||||
from ..config import Settings
|
from ..config import Settings
|
||||||
from ..config import get_settings
|
from ..config import get_settings
|
||||||
|
from ..db.models import Job
|
||||||
|
from ..db.models import JobStatus
|
||||||
|
from ..db.models import Source
|
||||||
from ..errors import AppError
|
from ..errors import AppError
|
||||||
from ..errors import ErrorCategory
|
from ..errors import ErrorCategory
|
||||||
from ..errors import classify_unexpected_error
|
from ..errors import classify_unexpected_error
|
||||||
from ..errors import format_error_detail
|
from ..errors import format_error_detail
|
||||||
from ..models import Job
|
|
||||||
from ..models import JobStatus
|
|
||||||
from ..models import Source
|
|
||||||
from ..providers import TranscriptionResult
|
from ..providers import TranscriptionResult
|
||||||
from . import ServiceBundle
|
from . import ServiceBundle
|
||||||
from .transcription import DEFAULT_PROMPT_FILE
|
from .transcription import DEFAULT_PROMPT_FILE
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ from uuid import uuid4
|
|||||||
from nicegui import ui
|
from nicegui import ui
|
||||||
|
|
||||||
from transcription.config import get_settings
|
from transcription.config import get_settings
|
||||||
from transcription.models import Source
|
from transcription.db.models import Source
|
||||||
|
|
||||||
PANGOZOOM_CDN_URL = "https://unpkg.com/@panzoom/[email protected]/dist/panzoom.min.js"
|
PANGOZOOM_CDN_URL = "https://unpkg.com/@panzoom/[email protected]/dist/panzoom.min.js"
|
||||||
UPLOADS_URL_PREFIX = "/uploads"
|
UPLOADS_URL_PREFIX = "/uploads"
|
||||||
|
|||||||
@@ -6,9 +6,9 @@ import logging
|
|||||||
|
|
||||||
from nicegui import ui
|
from nicegui import ui
|
||||||
|
|
||||||
from transcription.models import Job
|
from transcription.db.models import Job
|
||||||
from transcription.models import Revision
|
from transcription.db.models import Revision
|
||||||
from transcription.models import Source
|
from transcription.db.models import Source
|
||||||
from transcription.ui.components.document_panzoom import render_document_panzoom
|
from transcription.ui.components.document_panzoom import render_document_panzoom
|
||||||
from transcription.ui.components.transcript import render_original_transcription_card
|
from transcription.ui.components.transcript import render_original_transcription_card
|
||||||
from transcription.ui.components.transcript import render_revision_row
|
from transcription.ui.components.transcript import render_revision_row
|
||||||
|
|||||||
@@ -9,8 +9,8 @@ from typing import Any
|
|||||||
|
|
||||||
from nicegui import ui
|
from nicegui import ui
|
||||||
|
|
||||||
from transcription.models import Job
|
from transcription.db.models import Job
|
||||||
from transcription.models import Revision
|
from transcription.db.models import Revision
|
||||||
|
|
||||||
type RevisionAction = Callable[[Revision], Awaitable[None] | None]
|
type RevisionAction = Callable[[Revision], Awaitable[None] | None]
|
||||||
|
|
||||||
|
|||||||
@@ -4,19 +4,18 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from fastapi import Request
|
|
||||||
from nicegui import ui
|
from nicegui import ui
|
||||||
|
|
||||||
from transcription.app_state import resolve_session_factory
|
from transcription.db.models import Job
|
||||||
from transcription.models import Job
|
from transcription.db.models import JobStatus
|
||||||
from transcription.models import JobStatus
|
from transcription.db.models import Source
|
||||||
from transcription.models import Source
|
|
||||||
from transcription.services.jobs import JobService
|
from transcription.services.jobs import JobService
|
||||||
from transcription.services.transcription import TranscriptionService
|
from transcription.services.transcription import TranscriptionService
|
||||||
from transcription.ui.components.app_shell import render_navigation_header
|
from transcription.ui.components.app_shell import render_navigation_header
|
||||||
from transcription.ui.components.error_presenter import show_error
|
from transcription.ui.components.error_presenter import show_error
|
||||||
from transcription.ui.components.table.jobs import render_jobs_table
|
from transcription.ui.components.table.jobs import render_jobs_table
|
||||||
|
|
||||||
|
from ...db.session import SessionFactoryDep
|
||||||
from ..components.document_panzoom import render_document_panzoom
|
from ..components.document_panzoom import render_document_panzoom
|
||||||
from ..components.table.jobs import JobTableRow
|
from ..components.table.jobs import JobTableRow
|
||||||
from ..components.transcript import render_original_transcription_card
|
from ..components.transcript import render_original_transcription_card
|
||||||
@@ -27,8 +26,7 @@ def register_page() -> None: # noqa: PLR0915
|
|||||||
"""Register jobs list and detail routes."""
|
"""Register jobs list and detail routes."""
|
||||||
|
|
||||||
@ui.page("/jobs")
|
@ui.page("/jobs")
|
||||||
async def jobs_page(request: Request) -> None:
|
async def jobs_page(session_factory: SessionFactoryDep) -> None:
|
||||||
session_factory = resolve_session_factory(request.app.state)
|
|
||||||
jobs_service = JobService(session_factory=session_factory)
|
jobs_service = JobService(session_factory=session_factory)
|
||||||
render_navigation_header(current_path="/jobs")
|
render_navigation_header(current_path="/jobs")
|
||||||
|
|
||||||
@@ -51,8 +49,7 @@ def register_page() -> None: # noqa: PLR0915
|
|||||||
await render_table()
|
await render_table()
|
||||||
|
|
||||||
@ui.page("/jobs/{job_id}")
|
@ui.page("/jobs/{job_id}")
|
||||||
async def job_detail_page(job_id: str, request: Request) -> None: # noqa: PLR0915
|
async def job_detail_page(job_id: str, session_factory: SessionFactoryDep) -> None: # noqa: PLR0915
|
||||||
session_factory = resolve_session_factory(request.app.state)
|
|
||||||
jobs_service = JobService(session_factory=session_factory)
|
jobs_service = JobService(session_factory=session_factory)
|
||||||
transcription_service = TranscriptionService(session_factory=session_factory)
|
transcription_service = TranscriptionService(session_factory=session_factory)
|
||||||
render_navigation_header(current_path="/jobs")
|
render_navigation_header(current_path="/jobs")
|
||||||
|
|||||||
@@ -5,8 +5,7 @@ from __future__ import annotations
|
|||||||
from fastapi import Request
|
from fastapi import Request
|
||||||
from nicegui import ui
|
from nicegui import ui
|
||||||
|
|
||||||
from transcription.app_state import resolve_session_factory
|
from transcription.db import session_scope
|
||||||
from transcription.db import get_session
|
|
||||||
from transcription.services.store import create_upload_job
|
from transcription.services.store import create_upload_job
|
||||||
from transcription.ui.components.app_shell import render_navigation_header
|
from transcription.ui.components.app_shell import render_navigation_header
|
||||||
from transcription.ui.components.upload import render_upload_widget
|
from transcription.ui.components.upload import render_upload_widget
|
||||||
@@ -19,10 +18,9 @@ def register_page() -> None:
|
|||||||
@ui.page("/upload", title="Upload Document")
|
@ui.page("/upload", title="Upload Document")
|
||||||
def upload_page(request: Request) -> None:
|
def upload_page(request: Request) -> None:
|
||||||
render_navigation_header(current_path="/upload")
|
render_navigation_header(current_path="/upload")
|
||||||
session_factory = resolve_session_factory(request.app.state)
|
|
||||||
|
|
||||||
async def submit_upload(filename: str, file_bytes: bytes):
|
async def submit_upload(filename: str, file_bytes: bytes):
|
||||||
async with get_session(session_factory=session_factory) as session:
|
async with session_scope() as session:
|
||||||
return await create_upload_job(
|
return await create_upload_job(
|
||||||
filename=filename,
|
filename=filename,
|
||||||
file_bytes=file_bytes,
|
file_bytes=file_bytes,
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ from uuid import UUID
|
|||||||
from sqlalchemy.ext.asyncio import async_sessionmaker
|
from sqlalchemy.ext.asyncio import async_sessionmaker
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
from transcription.db import get_session
|
from transcription.db import session_scope
|
||||||
from transcription.errors import AppError
|
from transcription.errors import AppError
|
||||||
from transcription.errors import classify_unexpected_error
|
from transcription.errors import classify_unexpected_error
|
||||||
|
|
||||||
@@ -178,7 +178,7 @@ async def process_next_queued_job(
|
|||||||
)
|
)
|
||||||
|
|
||||||
if session is None:
|
if session is None:
|
||||||
async with get_session(session_factory=session_factory) as local_session:
|
async with session_scope(session_factory=session_factory) as local_session:
|
||||||
return await process_next_queued_job_workflow(services=services, session=local_session)
|
return await process_next_queued_job_workflow(services=services, session=local_session)
|
||||||
|
|
||||||
return await process_next_queued_job_workflow(services=services, session=session)
|
return await process_next_queued_job_workflow(services=services, session=session)
|
||||||
|
|||||||
+12
-8
@@ -13,11 +13,12 @@ from sqlmodel.pool import StaticPool
|
|||||||
|
|
||||||
from transcription.config import Settings
|
from transcription.config import Settings
|
||||||
from transcription.config import get_settings
|
from transcription.config import get_settings
|
||||||
|
from transcription.db.engine import get_database_url
|
||||||
|
from transcription.db.engine import get_engine
|
||||||
from transcription.db.operations import create_all
|
from transcription.db.operations import create_all
|
||||||
from transcription.db.runtime import dispose_database_runtime
|
from transcription.db.session import dispose_session_factory
|
||||||
from transcription.db.runtime import get_engine
|
from transcription.db.session import get_session_factory
|
||||||
from transcription.db.runtime import get_session
|
from transcription.db.session import session_scope
|
||||||
from transcription.db.runtime import get_session_factory
|
|
||||||
from transcription.services.documents import DocumentService
|
from transcription.services.documents import DocumentService
|
||||||
from transcription.services.jobs import JobService
|
from transcription.services.jobs import JobService
|
||||||
|
|
||||||
@@ -39,23 +40,26 @@ def session():
|
|||||||
async def default_settings():
|
async def default_settings():
|
||||||
"""Provide default settings for tests."""
|
"""Provide default settings for tests."""
|
||||||
settings = get_settings(database_url="sqlite:///:memory:")
|
settings = get_settings(database_url="sqlite:///:memory:")
|
||||||
await create_all(engine=get_engine(settings=settings))
|
db_url = get_database_url(settings)
|
||||||
|
await create_all(engine=get_engine(database_url=db_url))
|
||||||
return settings
|
return settings
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture
|
@pytest_asyncio.fixture
|
||||||
async def async_session(default_settings: Settings):
|
async def async_session(default_settings: Settings):
|
||||||
"""Provide a clean asynchronous database session for async tests."""
|
"""Provide a clean asynchronous database session for async tests."""
|
||||||
async with get_session(settings=default_settings) as async_session:
|
db_url = get_database_url(default_settings)
|
||||||
|
async with session_scope(database_url=db_url) as async_session:
|
||||||
yield async_session
|
yield async_session
|
||||||
|
|
||||||
await dispose_database_runtime()
|
await dispose_session_factory(db_url)
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def default_session_factory(default_settings: Settings):
|
def default_session_factory(default_settings: Settings):
|
||||||
"""Provide a base fixture for tests that require database access."""
|
"""Provide a base fixture for tests that require database access."""
|
||||||
session_factory = get_session_factory(settings=default_settings)
|
db_url = get_database_url(default_settings)
|
||||||
|
session_factory = get_session_factory(database_url=db_url)
|
||||||
return session_factory
|
return session_factory
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -5,8 +5,8 @@ from pathlib import Path
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from transcription.config import Settings
|
from transcription.config import Settings
|
||||||
from transcription.models import Job
|
from transcription.db.models import Job
|
||||||
from transcription.models import JobStatus
|
from transcription.db.models import JobStatus
|
||||||
from transcription.providers.base import TranscriptionResult
|
from transcription.providers.base import TranscriptionResult
|
||||||
from transcription.services.store import create_upload_job
|
from transcription.services.store import create_upload_job
|
||||||
from transcription.worker import process_next_queued_job
|
from transcription.worker import process_next_queued_job
|
||||||
|
|||||||
@@ -2,10 +2,10 @@ from uuid import uuid4
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from transcription.models import Document
|
from transcription.db.models import Document
|
||||||
from transcription.models import Job
|
from transcription.db.models import Job
|
||||||
from transcription.models import JobStatus
|
from transcription.db.models import JobStatus
|
||||||
from transcription.models import Source
|
from transcription.db.models import Source
|
||||||
from transcription.services.documents import DocumentService
|
from transcription.services.documents import DocumentService
|
||||||
from transcription.services.jobs import JobService
|
from transcription.services.jobs import JobService
|
||||||
|
|
||||||
|
|||||||
@@ -4,10 +4,10 @@ from uuid import uuid4
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from transcription.models import Document
|
from transcription.db.models import Document
|
||||||
from transcription.models import Job
|
from transcription.db.models import Job
|
||||||
from transcription.models import JobStatus
|
from transcription.db.models import JobStatus
|
||||||
from transcription.models import Source
|
from transcription.db.models import Source
|
||||||
from transcription.services.documents import DocumentService
|
from transcription.services.documents import DocumentService
|
||||||
from transcription.services.jobs import JobService
|
from transcription.services.jobs import JobService
|
||||||
from transcription.services.transcription import TranscriptionService
|
from transcription.services.transcription import TranscriptionService
|
||||||
|
|||||||
@@ -6,10 +6,10 @@ from uuid import uuid4
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from transcription.config import Settings
|
from transcription.config import Settings
|
||||||
from transcription.models import Document
|
from transcription.db.models import Document
|
||||||
from transcription.models import Job
|
from transcription.db.models import Job
|
||||||
from transcription.models import JobStatus
|
from transcription.db.models import JobStatus
|
||||||
from transcription.models import Source
|
from transcription.db.models import Source
|
||||||
from transcription.services import ServiceBundle
|
from transcription.services import ServiceBundle
|
||||||
from transcription.services.workflows import process_queued_job
|
from transcription.services.workflows import process_queued_job
|
||||||
|
|
||||||
|
|||||||
+5
-4
@@ -4,17 +4,18 @@ import pytest
|
|||||||
from sqlalchemy import inspect
|
from sqlalchemy import inspect
|
||||||
|
|
||||||
from transcription.config import Settings
|
from transcription.config import Settings
|
||||||
|
from transcription.config import SqliteSettings
|
||||||
from transcription.db import create_all
|
from transcription.db import create_all
|
||||||
from transcription.db import dispose_database_runtime
|
from transcription.db import dispose_database_runtime
|
||||||
from transcription.db import get_session
|
|
||||||
from transcription.db import initialize_database_runtime
|
from transcription.db import initialize_database_runtime
|
||||||
|
from transcription.db import session_scope
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_create_all_creates_expected_tables(tmp_path):
|
async def test_create_all_creates_expected_tables(tmp_path):
|
||||||
settings = Settings(
|
settings = Settings(
|
||||||
openrouter_api_key="test-key",
|
openrouter_api_key="test-key",
|
||||||
database_url=f"sqlite:///{tmp_path / 'schema.db'}",
|
database=SqliteSettings(path=str(tmp_path / "schema.db")),
|
||||||
environment="test",
|
environment="test",
|
||||||
)
|
)
|
||||||
runtime = initialize_database_runtime(settings=settings)
|
runtime = initialize_database_runtime(settings=settings)
|
||||||
@@ -36,13 +37,13 @@ async def test_create_all_creates_expected_tables(tmp_path):
|
|||||||
async def test_get_session_yields_async_session(tmp_path):
|
async def test_get_session_yields_async_session(tmp_path):
|
||||||
settings = Settings(
|
settings = Settings(
|
||||||
openrouter_api_key="test-key",
|
openrouter_api_key="test-key",
|
||||||
database_url=f"sqlite:///{tmp_path / 'session.db'}",
|
database=SqliteSettings(path=str(tmp_path / "session.db")),
|
||||||
environment="test",
|
environment="test",
|
||||||
)
|
)
|
||||||
initialize_database_runtime(settings=settings)
|
initialize_database_runtime(settings=settings)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
async with get_session(settings=settings) as session:
|
async with session_scope(settings=settings) as session:
|
||||||
assert session is not None
|
assert session is not None
|
||||||
finally:
|
finally:
|
||||||
await dispose_database_runtime()
|
await dispose_database_runtime()
|
||||||
|
|||||||
@@ -5,11 +5,11 @@ from uuid import UUID
|
|||||||
import pytest
|
import pytest
|
||||||
from sqlalchemy.exc import IntegrityError
|
from sqlalchemy.exc import IntegrityError
|
||||||
|
|
||||||
from transcription.models import Document
|
from transcription.db.models import Document
|
||||||
from transcription.models import Job
|
from transcription.db.models import Job
|
||||||
from transcription.models import JobStatus
|
from transcription.db.models import JobStatus
|
||||||
from transcription.models import Revision
|
from transcription.db.models import Revision
|
||||||
from transcription.models import Source
|
from transcription.db.models import Source
|
||||||
|
|
||||||
|
|
||||||
def _make_document(**overrides) -> Document:
|
def _make_document(**overrides) -> Document:
|
||||||
|
|||||||
+12
-12
@@ -4,6 +4,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
|
from collections.abc import Generator
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
@@ -14,32 +15,31 @@ from sqlmodel import delete
|
|||||||
|
|
||||||
from transcription.app import create_app
|
from transcription.app import create_app
|
||||||
from transcription.config import Settings
|
from transcription.config import Settings
|
||||||
from transcription.config import _settings
|
from transcription.config import SqliteSettings
|
||||||
from transcription.db import create_all
|
from transcription.db import create_all
|
||||||
from transcription.db import get_session
|
|
||||||
from transcription.db import initialize_database_runtime
|
from transcription.db import initialize_database_runtime
|
||||||
from transcription.models import Document
|
from transcription.db import session_scope
|
||||||
from transcription.models import Job
|
from transcription.db.models import Document
|
||||||
from transcription.models import JobStatus
|
from transcription.db.models import Job
|
||||||
from transcription.models import Revision
|
from transcription.db.models import JobStatus
|
||||||
from transcription.models import Source
|
from transcription.db.models import Revision
|
||||||
|
from transcription.db.models import Source
|
||||||
|
|
||||||
RevisionSeed = str
|
RevisionSeed = str
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="session")
|
@pytest.fixture(scope="session")
|
||||||
def app_client(tmp_path_factory: pytest.TempPathFactory) -> tuple[FastAPI, TestClient]:
|
def app_client(tmp_path_factory: pytest.TempPathFactory) -> Generator[tuple[FastAPI, TestClient]]:
|
||||||
"""Provide a real application and test client backed by in-memory SQLite."""
|
"""Provide a real application and test client backed by in-memory SQLite."""
|
||||||
tmp_path = tmp_path_factory.mktemp("ui")
|
tmp_path = tmp_path_factory.mktemp("ui")
|
||||||
settings = Settings(
|
settings = Settings(
|
||||||
openrouter_api_key="test-key",
|
openrouter_api_key="test-key",
|
||||||
database_url="sqlite:///:memory:",
|
database=SqliteSettings(path=":memory:"),
|
||||||
environment="test",
|
environment="test",
|
||||||
bootstrap_schema_on_startup=True,
|
bootstrap_schema_on_startup=True,
|
||||||
upload_dir=tmp_path / "uploads",
|
upload_dir=tmp_path / "uploads",
|
||||||
prompt_dir=tmp_path / "prompts",
|
prompt_dir=tmp_path / "prompts",
|
||||||
)
|
)
|
||||||
_settings.set(settings)
|
|
||||||
|
|
||||||
app = create_app()
|
app = create_app()
|
||||||
app.state.runtime = initialize_database_runtime(settings=settings)
|
app.state.runtime = initialize_database_runtime(settings=settings)
|
||||||
@@ -54,7 +54,7 @@ def clear_ui_database(app_client: tuple[FastAPI, TestClient]) -> None:
|
|||||||
app, _ = app_client
|
app, _ = app_client
|
||||||
|
|
||||||
async def _clear() -> None:
|
async def _clear() -> None:
|
||||||
async with get_session(session_factory=app.state.runtime.session_factory) as session:
|
async with session_scope() as session:
|
||||||
await session.exec(delete(Revision))
|
await session.exec(delete(Revision))
|
||||||
await session.exec(delete(Source))
|
await session.exec(delete(Source))
|
||||||
await session.exec(delete(Job))
|
await session.exec(delete(Job))
|
||||||
@@ -80,7 +80,7 @@ def seed_job(app_client: tuple[FastAPI, TestClient]) -> Callable[..., UUID]:
|
|||||||
source_file: Path | None = None,
|
source_file: Path | None = None,
|
||||||
) -> UUID:
|
) -> UUID:
|
||||||
async def _insert() -> UUID:
|
async def _insert() -> UUID:
|
||||||
async with get_session(session_factory=app.state.runtime.session_factory) as session:
|
async with session_scope() as session:
|
||||||
stored_path = app.state.settings.upload_dir / filename
|
stored_path = app.state.settings.upload_dir / filename
|
||||||
stored_path.parent.mkdir(parents=True, exist_ok=True)
|
stored_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
source_path = source_file or fixtures_dir / "small_png.png"
|
source_path = source_file or fixtures_dir / "small_png.png"
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from uuid import uuid4
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from transcription.models import JobStatus
|
from transcription.db.models import JobStatus
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
|
|||||||
Reference in New Issue
Block a user