736 lines
32 KiB
Python
736 lines
32 KiB
Python
"""Async SQLAlchemy engine lifecycle management.
|
||
|
||
Initializes at Gateway startup, provides session factory for
|
||
repositories, disposes at shutdown.
|
||
|
||
When database.backend="memory", init_engine is a no-op and
|
||
get_session_factory() returns None. Repositories must check for
|
||
None and fall back to in-memory implementations.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import asyncio
|
||
import json
|
||
import logging
|
||
import os
|
||
import sys
|
||
from importlib.util import find_spec
|
||
from pathlib import Path
|
||
|
||
from sqlalchemy import text
|
||
from sqlalchemy.engine import make_url
|
||
from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine
|
||
|
||
|
||
def _json_serializer(obj: object) -> str:
|
||
"""JSON serializer with ensure_ascii=False for Chinese character support."""
|
||
return json.dumps(obj, ensure_ascii=False)
|
||
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
_engine: AsyncEngine | None = None
|
||
_session_factory: async_sessionmaker[AsyncSession] | None = None
|
||
# The event loop on which the async engine was created. asyncmy/asyncpg connection
|
||
# pools bind their connections to the loop that opened them; awaiting a write on a
|
||
# *different* loop reuses those connections cross-loop and raises
|
||
# "got Future attached to a different loop". We capture the engine's home loop here so
|
||
# best-effort writes coming from worker threads / other loops can be routed back to it
|
||
# (see ``fire_and_forget_db_write``).
|
||
_engine_loop: asyncio.AbstractEventLoop | None = None
|
||
|
||
# These are the complete workflow-studio persistence tables owned by DeerFlow.
|
||
# Keep this explicit allow-list instead of treating ``workflow_%`` as a SQL
|
||
# wildcard: another application might share a schema and have its own table
|
||
# with that prefix.
|
||
_WORKFLOW_SCHEMA_TABLE_NAMES = (
|
||
"workflow_definitions",
|
||
"workflow_versions",
|
||
"workflow_runs",
|
||
"workflow_node_runs",
|
||
"workflow_run_events",
|
||
"workflow_run_artifacts",
|
||
"workflow_data_sources",
|
||
)
|
||
|
||
|
||
def _add_runtime_vendor_path() -> None:
|
||
"""Allow offline deployments to place optional DB drivers in DEER_FLOW_HOME.
|
||
|
||
The offline Docker image keeps the Python environment inside the image, while
|
||
``/app/backend/.deer-flow`` is a persistent bind mount. The MySQL switch
|
||
script installs ``asyncmy`` into ``{DEER_FLOW_HOME}/python-vendor`` so the
|
||
unchanged startup script can still import it after container recreation.
|
||
"""
|
||
try:
|
||
from deerflow.config.runtime_paths import runtime_home
|
||
|
||
vendor_dir = runtime_home() / "python-vendor"
|
||
except Exception:
|
||
vendor_dir = Path(".deer-flow/python-vendor").resolve()
|
||
|
||
vendor = str(vendor_dir)
|
||
if vendor_dir.exists() and vendor not in sys.path:
|
||
sys.path.insert(0, vendor)
|
||
|
||
|
||
def _mysql_driver_name(url: str) -> str:
|
||
if url.startswith("mysql+aiomysql://"):
|
||
return "aiomysql"
|
||
return "asyncmy"
|
||
|
||
|
||
async def _run_alembic_upgrade(url: str) -> None:
|
||
"""Apply pending Alembic migrations against ``url`` (``alembic upgrade head``).
|
||
|
||
Runs after ``Base.metadata.create_all`` so brand-new deployments still get
|
||
a complete schema even without Alembic, while existing deployments pick up
|
||
in-place schema fixes (column widenings, new indexes, etc.) at the next
|
||
restart — no manual ``alembic upgrade`` step required.
|
||
|
||
Notes:
|
||
|
||
* Each historical migration in ``migrations/versions/`` already guards
|
||
itself with ``_has_table`` / ``_has_column``, so re-running them against
|
||
a database that was originally built via ``create_all`` is a sequence of
|
||
no-ops up to the first migration that actually has work to do.
|
||
* The Alembic ``env.py`` calls ``asyncio.run`` internally, which cannot run
|
||
inside an already-running loop — so we hop into a worker thread.
|
||
* Failures here are logged but do not abort startup. The app remains
|
||
operable on the ``create_all`` schema; the operator just won't get
|
||
in-place fixes until they investigate.
|
||
"""
|
||
if os.getenv("DEER_FLOW_SKIP_ALEMBIC", "").lower() in ("1", "true", "yes"):
|
||
logger.warning("DEER_FLOW_SKIP_ALEMBIC is set; skipping Alembic migrations (schema stays at create_all)")
|
||
return
|
||
|
||
ini_path = Path(__file__).parent / "migrations" / "alembic.ini"
|
||
if not ini_path.exists():
|
||
logger.warning("Alembic config not found at %s; skipping auto-migrate", ini_path)
|
||
return
|
||
|
||
# Alembic's ``Config`` is backed by ConfigParser, which performs ``%``
|
||
# interpolation on every option value. URL-encoded credentials such as
|
||
# ``kland%4088124932`` (the ``@`` in ``kland@88124932``) trip that up with
|
||
# ``ValueError: invalid interpolation syntax ... at position N``. Doubling
|
||
# the percent signs neutralises the interpolation; ConfigParser un-doubles
|
||
# them back to the original value when read.
|
||
escaped_url = url.replace("%", "%%")
|
||
|
||
def _do_upgrade() -> None:
|
||
# Import inside the worker so a missing Alembic install can't break
|
||
# the engine import chain.
|
||
from alembic import command
|
||
from alembic.config import Config
|
||
|
||
cfg = Config(str(ini_path))
|
||
cfg.set_main_option("sqlalchemy.url", escaped_url)
|
||
# The ``script_location`` in alembic.ini is ``%(here)s`` which points
|
||
# at the directory of alembic.ini itself — already correct for us.
|
||
command.upgrade(cfg, "head")
|
||
|
||
# Bound the upgrade so a stuck DDL (e.g. MySQL metadata-lock wait) cannot
|
||
# hang startup forever. Tunable via DEER_FLOW_ALEMBIC_TIMEOUT (seconds).
|
||
# NOTE: asyncio.wait_for cancels the *await* on timeout; the worker thread
|
||
# itself keeps running (Python can't kill threads). That's acceptable —
|
||
# the orphan thread will exit when its MySQL connection eventually errors
|
||
# or when the process restarts. The point here is to unblock startup.
|
||
try:
|
||
timeout_s = float(os.getenv("DEER_FLOW_ALEMBIC_TIMEOUT", "60"))
|
||
except ValueError:
|
||
timeout_s = 60.0
|
||
|
||
try:
|
||
await asyncio.wait_for(asyncio.to_thread(_do_upgrade), timeout=timeout_s)
|
||
logger.info("Alembic migrations applied (or already up-to-date)")
|
||
except TimeoutError:
|
||
logger.error(
|
||
"Alembic upgrade did not finish within %.0fs; continuing startup on create_all schema. "
|
||
"Likely cause: MySQL metadata-lock wait — check SHOW FULL PROCESSLIST for a blocking transaction.",
|
||
timeout_s,
|
||
)
|
||
except Exception:
|
||
logger.exception("Auto-applying Alembic migrations failed; continuing with create_all schema")
|
||
|
||
|
||
async def _auto_create_postgres_db(url: str) -> None:
|
||
"""Connect to the ``postgres`` maintenance DB and CREATE DATABASE.
|
||
|
||
The target database name is extracted from *url*. The connection is
|
||
made to the default ``postgres`` database on the same server using
|
||
``AUTOCOMMIT`` isolation (CREATE DATABASE cannot run inside a
|
||
transaction).
|
||
"""
|
||
from sqlalchemy import text
|
||
from sqlalchemy.engine.url import make_url
|
||
|
||
parsed = make_url(url)
|
||
db_name = parsed.database
|
||
if not db_name:
|
||
raise ValueError("Cannot auto-create database: no database name in URL")
|
||
|
||
# Connect to the default 'postgres' database to issue CREATE DATABASE
|
||
maint_url = parsed.set(database="postgres")
|
||
maint_engine = create_async_engine(maint_url, isolation_level="AUTOCOMMIT")
|
||
try:
|
||
async with maint_engine.connect() as conn:
|
||
await conn.execute(text(f'CREATE DATABASE "{db_name}"'))
|
||
logger.info("Auto-created PostgreSQL-compatible database: %s", db_name)
|
||
finally:
|
||
await maint_engine.dispose()
|
||
|
||
|
||
async def _auto_create_mysql_db(url: str) -> None:
|
||
"""Create the target MySQL database when it is missing.
|
||
|
||
``create_all`` can create tables, but MySQL requires the database/schema to
|
||
exist before the application URL can connect. We connect to the built-in
|
||
``mysql`` maintenance database on the same server and issue a conservative
|
||
``CREATE DATABASE IF NOT EXISTS``.
|
||
"""
|
||
|
||
parsed = make_url(url)
|
||
db_name = parsed.database
|
||
if not db_name:
|
||
raise ValueError("Cannot auto-create MySQL database: no database name in URL")
|
||
|
||
maint_url = parsed.set(database="mysql")
|
||
maint_engine = create_async_engine(
|
||
maint_url.render_as_string(hide_password=False),
|
||
isolation_level="AUTOCOMMIT",
|
||
pool_pre_ping=True,
|
||
)
|
||
quoted = "`" + db_name.replace("`", "``") + "`"
|
||
try:
|
||
async with maint_engine.connect() as conn:
|
||
await conn.execute(
|
||
text(
|
||
f"CREATE DATABASE IF NOT EXISTS {quoted} "
|
||
"CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci"
|
||
)
|
||
)
|
||
logger.info("Auto-created MySQL database if missing: %s", db_name)
|
||
finally:
|
||
await maint_engine.dispose()
|
||
|
||
|
||
def _ensure_orm_columns_sync(sync_conn) -> None:
|
||
"""Add ORM-declared columns that are missing from already-existing tables.
|
||
|
||
``Base.metadata.create_all`` creates brand-new tables but never patches
|
||
existing ones, and Alembic auto-upgrade is disabled (see the commented-out
|
||
``_run_alembic_upgrade`` call) — so on an existing database a new
|
||
``mapped_column`` on an ORM model would otherwise error at first query
|
||
(e.g. ``no such column: roundtable_drafts.task_id`` on SQLite).
|
||
|
||
This sweep compares each existing table against the ORM metadata and issues
|
||
``ALTER TABLE … ADD COLUMN`` for every missing column that can be added
|
||
safely (nullable, or carrying a server_default). Indexes covering a newly
|
||
added column are created as well. Table-level ``UniqueConstraint`` objects
|
||
that are declared on the model but missing from the table are also created
|
||
(as a ``CREATE UNIQUE INDEX`` of the same name, which is what the migrations
|
||
emit too, and works on SQLite/MySQL/PostgreSQL/openGauss alike) — without this a newly
|
||
declared uniqueness rule (e.g. ``active_dedupe_key``) would silently be absent
|
||
on an existing database, weakening dedupe.
|
||
Works on sqlite/mysql/postgres/gauss — the column DDL is compiled through the
|
||
active dialect, so portable types (``PortableJSON``/``PortableLongText``/
|
||
``BeijingDateTime``) resolve to the right backend type. Everything is
|
||
per-column best-effort: a failure is logged and startup continues.
|
||
"""
|
||
from sqlalchemy import UniqueConstraint
|
||
from sqlalchemy import inspect as sa_inspect
|
||
from sqlalchemy.schema import CreateColumn
|
||
|
||
from deerflow.persistence.base import Base
|
||
|
||
try:
|
||
import deerflow.persistence.models # noqa: F401 — register all ORM models on Base.metadata
|
||
except ImportError:
|
||
pass
|
||
|
||
inspector = sa_inspect(sync_conn)
|
||
existing_tables = set(inspector.get_table_names())
|
||
preparer = sync_conn.dialect.identifier_preparer
|
||
|
||
for table in Base.metadata.sorted_tables:
|
||
if table.name not in existing_tables:
|
||
continue
|
||
existing_cols = {c["name"] for c in inspector.get_columns(table.name)}
|
||
missing = [c for c in table.columns if c.name not in existing_cols]
|
||
if not missing:
|
||
continue
|
||
|
||
added: set[str] = set()
|
||
for col in missing:
|
||
if not col.nullable and col.server_default is None:
|
||
logger.warning(
|
||
"Cannot auto-add NOT NULL column without server_default: %s.%s — add it manually",
|
||
table.name,
|
||
col.name,
|
||
)
|
||
continue
|
||
try:
|
||
spec = CreateColumn(col).compile(dialect=sync_conn.dialect).string
|
||
sync_conn.exec_driver_sql(f"ALTER TABLE {preparer.quote(table.name)} ADD COLUMN {spec}")
|
||
added.add(col.name)
|
||
logger.info("Auto-added missing column %s.%s", table.name, col.name)
|
||
except Exception:
|
||
logger.exception("Failed to auto-add column %s.%s", table.name, col.name)
|
||
|
||
if not added:
|
||
continue
|
||
try:
|
||
existing_idx = {ix["name"] for ix in inspector.get_indexes(table.name)}
|
||
except Exception:
|
||
existing_idx = set()
|
||
for idx in table.indexes:
|
||
if idx.name in existing_idx or not any(c.name in added for c in idx.columns):
|
||
continue
|
||
try:
|
||
idx.create(sync_conn)
|
||
logger.info("Auto-created index %s on %s", idx.name, table.name)
|
||
except Exception:
|
||
logger.exception("Failed to auto-create index %s on %s", idx.name, table.name)
|
||
|
||
# ── 表级 UNIQUE 约束补齐 ──────────────────────────────────────────────────
|
||
# create_all 只为全新表建 UniqueConstraint;既有表上新增声明的唯一约束(如
|
||
# roundtable_jobs.active_dedupe_key)会静默缺失,去重/幂等的正确性就靠它,故
|
||
# 这里显式补建。迁移可能用 CREATE UNIQUE INDEX(同名)实现唯一性,故同时查
|
||
# get_unique_constraints 与 get_indexes 的名字集合,避免重复建。
|
||
for table in Base.metadata.sorted_tables:
|
||
if table.name not in existing_tables:
|
||
continue
|
||
named_uqs = [
|
||
c for c in table.constraints
|
||
if isinstance(c, UniqueConstraint) and c.name is not None
|
||
]
|
||
if not named_uqs:
|
||
continue
|
||
try:
|
||
existing_uq_names = {
|
||
uq["name"] for uq in inspector.get_unique_constraints(table.name) if uq.get("name")
|
||
}
|
||
except Exception:
|
||
existing_uq_names = set()
|
||
try:
|
||
existing_idx_names = {
|
||
ix["name"] for ix in inspector.get_indexes(table.name) if ix.get("name")
|
||
}
|
||
except Exception:
|
||
existing_idx_names = set()
|
||
for uq in named_uqs:
|
||
if uq.name in existing_uq_names or uq.name in existing_idx_names:
|
||
continue
|
||
# 用 CREATE UNIQUE INDEX(而非 ALTER TABLE ADD CONSTRAINT):迁移本身也用
|
||
# CREATE UNIQUE INDEX,且它在 SQLite/MySQL/PostgreSQL 都受支持(SQLite 不
|
||
# 支持 ALTER TABLE ADD CONSTRAINT),功能等价地强制唯一性。
|
||
cols = ", ".join(preparer.quote(c.name) for c in uq.columns)
|
||
try:
|
||
sync_conn.exec_driver_sql(
|
||
f"CREATE UNIQUE INDEX {preparer.quote(uq.name)} "
|
||
f"ON {preparer.quote(table.name)} ({cols})"
|
||
)
|
||
logger.info("Auto-added unique index %s on %s", uq.name, table.name)
|
||
except Exception:
|
||
logger.exception("Failed to add unique index %s on %s", uq.name, table.name)
|
||
|
||
|
||
def _is_mysql_unknown_database_error(exc: Exception) -> bool:
|
||
message = str(exc).lower()
|
||
return "unknown database" in message or "1049" in message
|
||
|
||
|
||
async def init_engine(
|
||
backend: str,
|
||
*,
|
||
url: str = "",
|
||
echo: bool = False,
|
||
pool_size: int = 5,
|
||
max_overflow: int = 10,
|
||
pool_timeout: int = 30,
|
||
pool_recycle: int = 1800,
|
||
sqlite_dir: str = "",
|
||
) -> None:
|
||
"""Create the async engine and session factory, then auto-create tables.
|
||
|
||
Args:
|
||
backend: "memory", "sqlite", "postgres", "gauss", or "mysql".
|
||
url: SQLAlchemy async URL (for sqlite/postgres/gauss/mysql).
|
||
echo: Echo SQL to log.
|
||
pool_size: Base connection pool size (postgres/gauss/mysql).
|
||
max_overflow: Extra connections allowed beyond pool_size under load (postgres/gauss/mysql).
|
||
pool_timeout: Seconds to wait for a free pooled connection before erroring (postgres/gauss/mysql).
|
||
pool_recycle: Recycle connections after N seconds (postgres/gauss/mysql). Prevents stale
|
||
connections dropped by firewalls / MySQL wait_timeout / pgbouncer.
|
||
sqlite_dir: Directory to create for SQLite (ensured to exist).
|
||
"""
|
||
global _engine, _session_factory, _engine_loop
|
||
|
||
if backend == "memory":
|
||
logger.info("Persistence backend=memory -- ORM engine not initialized")
|
||
return
|
||
|
||
# Remember the loop that owns the engine's connection pool, so off-loop
|
||
# best-effort writes (metrics from worker threads, etc.) can be routed here
|
||
# instead of reusing pooled connections cross-loop.
|
||
try:
|
||
_engine_loop = asyncio.get_running_loop()
|
||
except RuntimeError:
|
||
_engine_loop = None
|
||
|
||
if backend in {"postgres", "gauss"}:
|
||
try:
|
||
import asyncpg # noqa: F401
|
||
except ImportError:
|
||
extra = "gauss" if backend == "gauss" else "postgres"
|
||
raise ImportError(
|
||
f"database.backend is set to {backend!r} but asyncpg is not installed.\n"
|
||
f"Install it with:\n uv sync --extra {extra}\n"
|
||
"Or switch to backend: sqlite in config.yaml for single-node deployment."
|
||
) from None
|
||
if backend == "gauss" and find_spec("opengauss_sqlalchemy") is None:
|
||
raise ImportError(
|
||
"database.backend is set to 'gauss' but the openGauss SQLAlchemy dialect "
|
||
"is not installed.\nInstall it with:\n uv sync --extra gauss"
|
||
)
|
||
if backend == "mysql":
|
||
_add_runtime_vendor_path()
|
||
driver = _mysql_driver_name(url)
|
||
if find_spec(driver) is None:
|
||
raise ImportError(
|
||
"database.backend is set to 'mysql' but the MySQL async driver is not installed.\n"
|
||
"Install asyncmy, or run switch_database_to_mysql.sh so it can install asyncmy into "
|
||
f"{Path('.deer-flow/python-vendor')} from the offline wheelhouse."
|
||
)
|
||
|
||
if backend == "sqlite":
|
||
import os
|
||
|
||
from sqlalchemy import event
|
||
|
||
os.makedirs(sqlite_dir or ".", exist_ok=True)
|
||
_engine = create_async_engine(url, echo=echo, json_serializer=_json_serializer)
|
||
|
||
# Enable WAL on every new connection. SQLite PRAGMA settings are
|
||
# per-connection, so we wire the listener instead of running PRAGMA
|
||
# once at startup. WAL gives concurrent reads + writers without
|
||
# blocking and is the standard recommendation for any production
|
||
# SQLite deployment (TC-UPG-06 in AUTH_TEST_PLAN.md). The companion
|
||
# ``synchronous=NORMAL`` is the safe-and-fast pairing — fsync only
|
||
# at WAL checkpoint boundaries instead of every commit.
|
||
# Note: we do not set PRAGMA busy_timeout here — Python's sqlite3
|
||
# driver already defaults to a 5-second busy timeout (see the
|
||
# ``timeout`` kwarg of ``sqlite3.connect``), and aiosqlite /
|
||
# SQLAlchemy's aiosqlite dialect inherit that default. Setting
|
||
# it again would be a no-op.
|
||
@event.listens_for(_engine.sync_engine, "connect")
|
||
def _enable_sqlite_wal(dbapi_conn, _record): # noqa: ARG001 — SQLAlchemy contract
|
||
cursor = dbapi_conn.cursor()
|
||
try:
|
||
cursor.execute("PRAGMA journal_mode=WAL;")
|
||
cursor.execute("PRAGMA synchronous=NORMAL;")
|
||
cursor.execute("PRAGMA foreign_keys=ON;")
|
||
finally:
|
||
cursor.close()
|
||
elif backend in {"postgres", "gauss"}:
|
||
_engine = create_async_engine(
|
||
url,
|
||
echo=echo,
|
||
pool_size=pool_size,
|
||
max_overflow=max_overflow,
|
||
pool_timeout=pool_timeout,
|
||
pool_pre_ping=True,
|
||
pool_recycle=pool_recycle,
|
||
json_serializer=_json_serializer,
|
||
)
|
||
elif backend == "mysql":
|
||
_engine = create_async_engine(
|
||
url,
|
||
echo=echo,
|
||
pool_size=pool_size,
|
||
max_overflow=max_overflow,
|
||
pool_timeout=pool_timeout,
|
||
pool_pre_ping=True,
|
||
pool_recycle=pool_recycle,
|
||
json_serializer=_json_serializer,
|
||
)
|
||
else:
|
||
raise ValueError(f"Unknown persistence backend: {backend!r}")
|
||
|
||
_session_factory = async_sessionmaker(_engine, expire_on_commit=False)
|
||
|
||
# Auto-create tables for fresh deployments, then let Alembic apply any
|
||
# in-place schema fixes on top (see ``_run_alembic_upgrade`` below). Both
|
||
# are idempotent so repeating them across restarts is harmless.
|
||
from deerflow.persistence.base import Base
|
||
|
||
# Import all models so Base.metadata discovers them.
|
||
# When no models exist yet (scaffolding phase), this is a no-op.
|
||
try:
|
||
import deerflow.persistence.models # noqa: F401
|
||
except ImportError:
|
||
# Models package not yet available — tables won't be auto-created.
|
||
# This is expected during initial scaffolding or minimal installs.
|
||
logger.debug("deerflow.persistence.models not found; skipping auto-create tables")
|
||
|
||
try:
|
||
async with _engine.begin() as conn:
|
||
await conn.run_sync(Base.metadata.create_all)
|
||
except Exception as exc:
|
||
if backend in {"postgres", "gauss"} and "does not exist" in str(exc).lower():
|
||
# Database not yet created — attempt to auto-create it, then retry.
|
||
await _auto_create_postgres_db(url)
|
||
# Rebuild engine against the now-existing database
|
||
await _engine.dispose()
|
||
_engine = create_async_engine(url, echo=echo, pool_size=pool_size, max_overflow=max_overflow, pool_timeout=pool_timeout, pool_pre_ping=True, pool_recycle=pool_recycle, json_serializer=_json_serializer)
|
||
_session_factory = async_sessionmaker(_engine, expire_on_commit=False)
|
||
async with _engine.begin() as conn:
|
||
await conn.run_sync(Base.metadata.create_all)
|
||
elif backend == "mysql" and _is_mysql_unknown_database_error(exc):
|
||
# Database not yet created. Create it, rebuild the app engine
|
||
# against the target DB, then let SQLAlchemy create all tables.
|
||
await _auto_create_mysql_db(url)
|
||
await _engine.dispose()
|
||
_engine = create_async_engine(
|
||
url,
|
||
echo=echo,
|
||
pool_size=pool_size,
|
||
max_overflow=max_overflow,
|
||
pool_timeout=pool_timeout,
|
||
pool_pre_ping=True,
|
||
pool_recycle=pool_recycle,
|
||
json_serializer=_json_serializer,
|
||
)
|
||
_session_factory = async_sessionmaker(_engine, expire_on_commit=False)
|
||
async with _engine.begin() as conn:
|
||
await conn.run_sync(Base.metadata.create_all)
|
||
else:
|
||
raise
|
||
|
||
# Patch existing tables with ORM-declared columns that ``create_all``
|
||
# cannot add (it only creates brand-new tables). Best-effort: a failure is
|
||
# logged and startup continues on the current schema.
|
||
try:
|
||
async with _engine.begin() as conn:
|
||
await conn.run_sync(_ensure_orm_columns_sync)
|
||
except Exception:
|
||
logger.exception("Auto-adding missing ORM columns failed; continuing with current schema")
|
||
|
||
# Apply in-place schema fixes (column widenings, etc.) that ``create_all``
|
||
# cannot do on its own. Failures don't abort startup — see the helper.
|
||
# await _run_alembic_upgrade(url)
|
||
|
||
logger.info("Persistence engine initialized: backend=%s dialect=%s", backend, _engine.dialect.name)
|
||
|
||
|
||
async def init_engine_from_config(config) -> None:
|
||
"""Convenience: init engine from a DatabaseConfig object."""
|
||
if config.backend == "memory":
|
||
await init_engine("memory")
|
||
return
|
||
await init_engine(
|
||
backend=config.backend,
|
||
url=config.app_sqlalchemy_url,
|
||
echo=config.echo_sql,
|
||
pool_size=config.pool_size,
|
||
max_overflow=getattr(config, "max_overflow", 10),
|
||
pool_timeout=getattr(config, "pool_timeout", 30),
|
||
pool_recycle=getattr(config, "pool_recycle", 1800),
|
||
sqlite_dir=config.sqlite_dir if config.backend == "sqlite" else "",
|
||
)
|
||
|
||
|
||
def get_session_factory() -> async_sessionmaker[AsyncSession] | None:
|
||
"""Return the async session factory, or None if backend=memory."""
|
||
return _session_factory
|
||
|
||
|
||
def get_engine() -> AsyncEngine | None:
|
||
"""Return the async engine, or None if not initialized."""
|
||
return _engine
|
||
|
||
|
||
def get_engine_loop() -> asyncio.AbstractEventLoop | None:
|
||
"""Return the event loop that owns the engine's connection pool (or None)."""
|
||
return _engine_loop
|
||
|
||
|
||
async def rebuild_workflow_schema_for_mysql() -> list[str]:
|
||
"""Drop and recreate DeerFlow's workflow tables in the current MySQL schema.
|
||
|
||
This is a deliberately narrow recovery operation for an unused Workflow
|
||
Studio installation. It never uses a prefix wildcard, never changes the
|
||
selected database, and refuses to run unless the engine is connected to
|
||
MySQL and ``SELECT DATABASE()`` confirms the URL's exact target schema.
|
||
|
||
Returns the allow-listed workflow tables that existed before the rebuild.
|
||
"""
|
||
engine = get_engine()
|
||
if engine is None:
|
||
raise RuntimeError("Cannot rebuild workflow schema: persistence engine is not initialized")
|
||
if engine.dialect.name != "mysql":
|
||
raise RuntimeError(
|
||
"DEER_FLOW_REBUILD_WORKFLOW_SCHEMA is supported only for a MySQL persistence backend"
|
||
)
|
||
|
||
expected_database = engine.url.database
|
||
if not expected_database:
|
||
raise RuntimeError("Cannot rebuild workflow schema: MySQL URL has no database name")
|
||
|
||
# Make the operation safe when this helper is invoked independently from
|
||
# normal app startup as well.
|
||
import deerflow.persistence.models # noqa: F401 — register ORM models
|
||
from deerflow.persistence.base import Base
|
||
|
||
missing_models = [name for name in _WORKFLOW_SCHEMA_TABLE_NAMES if name not in Base.metadata.tables]
|
||
if missing_models:
|
||
raise RuntimeError(
|
||
"Cannot rebuild workflow schema: ORM metadata is missing " + ", ".join(missing_models)
|
||
)
|
||
workflow_tables = [Base.metadata.tables[name] for name in _WORKFLOW_SCHEMA_TABLE_NAMES]
|
||
|
||
def _rebuild(sync_conn) -> list[str]:
|
||
from sqlalchemy import inspect as sa_inspect
|
||
|
||
current_database = sync_conn.exec_driver_sql("SELECT DATABASE()").scalar_one_or_none()
|
||
if current_database != expected_database:
|
||
raise RuntimeError(
|
||
"Refusing workflow schema rebuild: connected database "
|
||
f"{current_database!r} does not match configured database {expected_database!r}"
|
||
)
|
||
|
||
existing = set(sa_inspect(sync_conn).get_table_names())
|
||
present = [name for name in _WORKFLOW_SCHEMA_TABLE_NAMES if name in existing]
|
||
# SQLAlchemy orders parent/child tables around the foreign key from
|
||
# workflow_versions to workflow_definitions. Only tables in the exact
|
||
# allow-list above are ever passed to drop_all/create_all.
|
||
Base.metadata.drop_all(sync_conn, tables=workflow_tables, checkfirst=True)
|
||
Base.metadata.create_all(sync_conn, tables=workflow_tables, checkfirst=True)
|
||
return present
|
||
|
||
async with engine.begin() as conn:
|
||
rebuilt = await conn.run_sync(_rebuild)
|
||
|
||
logger.warning(
|
||
"Rebuilt workflow schema in current MySQL database %s; tables reset: %s",
|
||
expected_database,
|
||
", ".join(rebuilt) if rebuilt else "none (created fresh)",
|
||
)
|
||
return rebuilt
|
||
|
||
|
||
def fire_and_forget_db_write(coro) -> None:
|
||
"""Schedule a best-effort DB write on the engine's home loop, from any loop/thread.
|
||
|
||
The async engine's connection pool is bound to the loop that created it
|
||
(``init_engine``, normally the Gateway's main loop). Running a write on a
|
||
*different* loop — e.g. a model/tool call that LangGraph executed in a worker
|
||
thread, which then does ``asyncio.run(store.record(...))`` — reuses those pooled
|
||
connections cross-loop and raises ``got Future attached to a different loop``.
|
||
Under concurrency this happens constantly (many calls land in worker threads at
|
||
once), spamming the log and churning the pool; single-threaded it almost never does.
|
||
|
||
This routes every such write back onto the engine's home loop:
|
||
- already on that loop -> ``create_task`` (in-loop fire-and-forget);
|
||
- on another loop/thread -> ``run_coroutine_threadsafe`` to the home loop.
|
||
|
||
The write itself stays best-effort: if no engine loop is captured (memory backend
|
||
/ not initialized) the coroutine is closed (avoids a "never awaited" warning) and
|
||
the call is a no-op. Never raises.
|
||
"""
|
||
loop = _engine_loop
|
||
if loop is None or loop.is_closed():
|
||
close = getattr(coro, "close", None)
|
||
if callable(close):
|
||
close()
|
||
return
|
||
try:
|
||
running = asyncio.get_running_loop()
|
||
except RuntimeError:
|
||
running = None
|
||
try:
|
||
if running is loop:
|
||
loop.create_task(coro)
|
||
else:
|
||
asyncio.run_coroutine_threadsafe(coro, loop)
|
||
except Exception: # noqa: BLE001 — best-effort; scheduling must never break the caller
|
||
close = getattr(coro, "close", None)
|
||
if callable(close):
|
||
close()
|
||
logger.debug("fire_and_forget_db_write: failed to schedule DB write", exc_info=True)
|
||
|
||
|
||
def _run_coro_on_throwaway_loop(coro):
|
||
"""Last-resort: run *coro* to completion on a fresh event loop.
|
||
|
||
Only reached when there is no usable engine home loop, or when the caller is
|
||
*synchronously* inside the home loop (routing back would deadlock). This
|
||
touches the shared connection pool from a throwaway loop, so it can still leak
|
||
a dead-loop-bound connection — but the common worker-thread path no longer
|
||
gets here (see ``run_db_blocking``).
|
||
"""
|
||
try:
|
||
asyncio.get_running_loop()
|
||
in_loop = True
|
||
except RuntimeError:
|
||
in_loop = False
|
||
|
||
if in_loop:
|
||
import concurrent.futures
|
||
|
||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||
return pool.submit(asyncio.run, coro).result()
|
||
return asyncio.run(coro)
|
||
|
||
|
||
def run_db_blocking(coro, *, timeout: float | None = 30.0):
|
||
"""Run an awaited DB coroutine from sync code, safe w.r.t. the engine's home loop.
|
||
|
||
The blocking-read sibling of :func:`fire_and_forget_db_write`: where that one
|
||
drops the result, this returns it — used by the sync→async bridges that must
|
||
read a row before proceeding (thread-path owner resolution, always-on skills,
|
||
skill visibility).
|
||
|
||
The async engine's connection pool is bound to the loop that created it
|
||
(``init_engine``). Running a DB coroutine on a *throwaway* loop
|
||
(``asyncio.run``) checks out pooled connections cross-loop and **leaks** them:
|
||
the connection ends up bound to a loop that is then closed, and when the pool
|
||
later recycles/disposes it, asyncmy's ``_terminate_graceful_close`` fires on the
|
||
wrong loop and raises ``got Future <...> attached to a different loop``. Under
|
||
concurrency (scheduler / subagent / sandbox-tool worker threads each spinning a
|
||
fresh loop) this floods the logs.
|
||
|
||
So when we are **off** the home loop — the common case, a worker thread with no
|
||
running loop — route the coroutine onto the home loop and block the *caller*
|
||
thread for its result. The pool is only ever touched on its own loop, so nothing
|
||
leaks. Only when there is no home loop, or we are synchronously inside it (where
|
||
blocking would deadlock), do we fall back to a throwaway loop.
|
||
|
||
Never swallows the coroutine's own exception — the sync bridges wrap this call in
|
||
their own try/except and degrade gracefully.
|
||
"""
|
||
loop = _engine_loop
|
||
try:
|
||
running = asyncio.get_running_loop()
|
||
except RuntimeError:
|
||
running = None
|
||
if loop is not None and not loop.is_closed() and running is not loop:
|
||
# Off the home loop (worker thread, or a different loop): execute the coro
|
||
# on the engine's loop and block this thread for the result. asyncmy's
|
||
# pooled connections stay on their own loop, so none leak cross-loop.
|
||
return asyncio.run_coroutine_threadsafe(coro, loop).result(timeout)
|
||
# No home loop, or we're already synchronously inside it: last resort.
|
||
return _run_coro_on_throwaway_loop(coro)
|
||
|
||
|
||
async def close_engine() -> None:
|
||
"""Dispose the engine, release all connections."""
|
||
global _engine, _session_factory, _engine_loop
|
||
if _engine is not None:
|
||
await _engine.dispose()
|
||
logger.info("Persistence engine closed")
|
||
_engine = None
|
||
_session_factory = None
|
||
_engine_loop = None
|