feat: Use Alembic with async driver (#6258)
Use Alembic with async driver
This commit is contained in:
parent
93bc185ca9
commit
ffbc97bfc9
6 changed files with 95 additions and 74 deletions
|
|
@ -1,9 +1,11 @@
|
|||
# noqa: INP001
|
||||
import asyncio
|
||||
from logging.config import fileConfig
|
||||
|
||||
from alembic import context
|
||||
from sqlalchemy import engine_from_config, pool, text
|
||||
from sqlalchemy import pool, text
|
||||
from sqlalchemy.event import listen
|
||||
from sqlalchemy.ext.asyncio import async_engine_from_config
|
||||
|
||||
from langflow.services.database.service import SQLModel
|
||||
|
||||
|
|
@ -68,14 +70,18 @@ def _sqlite_do_begin(conn):
|
|||
conn.exec_driver_sql("BEGIN EXCLUSIVE")
|
||||
|
||||
|
||||
def run_migrations_online() -> None:
|
||||
"""Run migrations in 'online' mode.
|
||||
def _do_run_migrations(connection):
|
||||
context.configure(connection=connection, target_metadata=target_metadata, render_as_batch=True)
|
||||
|
||||
In this scenario we need to create an Engine
|
||||
and associate a connection with the context.
|
||||
with context.begin_transaction():
|
||||
if connection.dialect.name == "postgresql":
|
||||
connection.execute(text("SET LOCAL lock_timeout = '60s';"))
|
||||
connection.execute(text("SELECT pg_advisory_xact_lock(112233);"))
|
||||
context.run_migrations()
|
||||
|
||||
"""
|
||||
connectable = engine_from_config(
|
||||
|
||||
async def _run_async_migrations() -> None:
|
||||
connectable = async_engine_from_config(
|
||||
config.get_section(config.config_ini_section, {}),
|
||||
prefix="sqlalchemy.",
|
||||
poolclass=pool.NullPool,
|
||||
|
|
@ -83,17 +89,23 @@ def run_migrations_online() -> None:
|
|||
|
||||
if connectable.dialect.name == "sqlite":
|
||||
# See https://docs.sqlalchemy.org/en/20/dialects/sqlite.html#serializable-isolation-savepoints-transactional-ddl
|
||||
listen(connectable, "connect", _sqlite_do_connect)
|
||||
listen(connectable, "begin", _sqlite_do_begin)
|
||||
listen(connectable.sync_engine, "connect", _sqlite_do_connect)
|
||||
listen(connectable.sync_engine, "begin", _sqlite_do_begin)
|
||||
|
||||
with connectable.connect() as connection:
|
||||
context.configure(connection=connection, target_metadata=target_metadata, render_as_batch=True)
|
||||
async with connectable.connect() as connection:
|
||||
await connection.run_sync(_do_run_migrations)
|
||||
|
||||
with context.begin_transaction():
|
||||
if connection.dialect.name == "postgresql":
|
||||
connection.execute(text("SET LOCAL lock_timeout = '60s';"))
|
||||
connection.execute(text("SELECT pg_advisory_xact_lock(112233);"))
|
||||
context.run_migrations()
|
||||
await connectable.dispose()
|
||||
|
||||
|
||||
def run_migrations_online() -> None:
|
||||
"""Run migrations in 'online' mode.
|
||||
|
||||
In this scenario we need to create an Engine
|
||||
and associate a connection with the context.
|
||||
|
||||
"""
|
||||
asyncio.run(_run_async_migrations())
|
||||
|
||||
|
||||
if context.is_offline_mode():
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ import sqlalchemy as sa
|
|||
from alembic import command, util
|
||||
from alembic.config import Config
|
||||
from loguru import logger
|
||||
from sqlalchemy import AsyncAdaptedQueuePool, event, exc, inspect
|
||||
from sqlalchemy import event, exc, inspect
|
||||
from sqlalchemy.dialects import sqlite as dialect_sqlite
|
||||
from sqlalchemy.engine import Engine
|
||||
from sqlalchemy.exc import OperationalError
|
||||
|
|
@ -81,12 +81,23 @@ class DatabaseService(Service):
|
|||
self.engine = self._create_engine()
|
||||
|
||||
def _sanitize_database_url(self):
|
||||
if self.database_url.startswith("postgres://"):
|
||||
self.database_url = self.database_url.replace("postgres://", "postgresql://")
|
||||
logger.warning(
|
||||
"Fixed postgres dialect in database URL. Replacing postgres:// with postgresql://. "
|
||||
"To avoid this warning, update the database URL."
|
||||
)
|
||||
"""Create the engine for the database."""
|
||||
url_components = self.database_url.split("://", maxsplit=1)
|
||||
|
||||
driver = url_components[0]
|
||||
|
||||
if driver == "sqlite":
|
||||
driver = "sqlite+aiosqlite"
|
||||
elif driver in {"postgresql", "postgres"}:
|
||||
if driver == "postgres":
|
||||
logger.warning(
|
||||
"The postgres dialect in the database URL is deprecated. "
|
||||
"Use postgresql instead. "
|
||||
"To avoid this warning, update the database URL."
|
||||
)
|
||||
driver = "postgresql+psycopg"
|
||||
|
||||
self.database_url = f"{driver}://{url_components[1]}"
|
||||
|
||||
def _build_connection_kwargs(self):
|
||||
"""Build connection kwargs by merging deprecated settings with db_connection_settings.
|
||||
|
|
@ -109,28 +120,13 @@ class DatabaseService(Service):
|
|||
return connection_kwargs
|
||||
|
||||
def _create_engine(self) -> AsyncEngine:
|
||||
"""Create the engine for the database."""
|
||||
url_components = self.database_url.split("://", maxsplit=1)
|
||||
|
||||
# Get connection settings from config, with defaults if not specified
|
||||
# if the user specifies an empty dict, we allow it.
|
||||
kwargs = self._build_connection_kwargs()
|
||||
|
||||
if url_components[0].startswith("sqlite"):
|
||||
scheme = "sqlite+aiosqlite"
|
||||
# Even though the docs say this is the default, it raises an error
|
||||
# if we don't specify it.
|
||||
# https://docs.sqlalchemy.org/en/20/errors.html#pool-class-cannot-be-used-with-asyncio-engine-or-vice-versa
|
||||
pool = AsyncAdaptedQueuePool
|
||||
else:
|
||||
scheme = "postgresql+psycopg" if url_components[0].startswith("postgresql") else url_components[0]
|
||||
pool = None
|
||||
|
||||
database_url = f"{scheme}://{url_components[1]}"
|
||||
return create_async_engine(
|
||||
database_url,
|
||||
self.database_url,
|
||||
connect_args=self._get_connect_args(),
|
||||
poolclass=pool,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -70,7 +70,10 @@ class Settings(BaseSettings):
|
|||
dev: bool = False
|
||||
"""If True, Langflow will run in development mode."""
|
||||
database_url: str | None = None
|
||||
"""Database URL for Langflow. If not provided, Langflow will use a SQLite database."""
|
||||
"""Database URL for Langflow. If not provided, Langflow will use a SQLite database.
|
||||
The driver shall be an async one like `sqlite+aiosqlite` (`sqlite` and `postgresql`
|
||||
will be automatically converted to the async drivers `sqlite+aiosqlite` and
|
||||
`postgresql+psycopg` respectively)."""
|
||||
database_connection_retry: bool = False
|
||||
"""If True, Langflow will retry to connect to the database if it fails."""
|
||||
pool_size: int = 10
|
||||
|
|
|
|||
|
|
@ -75,6 +75,12 @@ def blockbuster(request):
|
|||
.can_block_in("langchain_core/_api/internal.py", "is_caller_internal")
|
||||
)
|
||||
|
||||
for func in ["os.stat", "os.path.abspath", "os.scandir"]:
|
||||
bb.functions[func].can_block_in("alembic/util/pyfiles.py", "load_python_file")
|
||||
|
||||
for func in ["os.path.abspath", "os.scandir"]:
|
||||
bb.functions[func].can_block_in("alembic/script/base.py", "_load_revisions")
|
||||
|
||||
(
|
||||
bb.functions["os.path.abspath"]
|
||||
.can_block_in("loguru/_better_exceptions.py", {"_get_lib_dirs", "_format_exception"})
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue