aps-agent/server/db/database.py

160 lines
5.3 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# ============================================================
# 数据库引擎与会话(moduleId: db-database, 可重生 ✅)
# SQLite 单文件:server/data/master.db(APS_DB_PATH 可覆盖)
# ============================================================
from __future__ import annotations
import os
import threading
from sqlalchemy import create_engine
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, sessionmaker
_PATH_ENV = "APS_DB_PATH"
_URL_ENV = "APS_DATABASE_URL"
_APPROVAL_SQLITE_TEST_ENV = "APS_APPROVAL_DATABASE_ALLOW_SQLITE"
_MIGRATION_MANAGED_TABLES = frozenset({
"aps_approval_requests",
"aps_approval_grants",
"aps_approval_events",
})
def _default_db_path() -> str:
try:
from server.aps_home import default_db_path
return default_db_path()
except Exception:
return "server/data/master.db"
_DEFAULT_PATH = "server/data/master.db"
_lock = threading.Lock()
_engine: Engine | None = None
_engine_path: str | None = None
_SessionLocal: sessionmaker | None = None
def _db_path() -> str:
return os.environ.get(_PATH_ENV) or _default_db_path()
def _database_url() -> str:
configured = (os.environ.get(_URL_ENV) or "").strip()
if configured:
return configured
return f"sqlite:///{_db_path()}"
def get_engine() -> Engine:
"""懒加载引擎;APS_DB_PATH 变化时自动重建(测试隔离用)。"""
global _engine, _engine_path, _SessionLocal
url = _database_url()
with _lock:
if _engine is None or _engine_path != url:
kwargs = {"pool_pre_ping": True}
if url.startswith("sqlite:///"):
path = url.removeprefix("sqlite:///")
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
kwargs["connect_args"] = {"check_same_thread": False}
_engine = create_engine(url, **kwargs)
_engine_path = url
_SessionLocal = sessionmaker(bind=_engine, expire_on_commit=False)
init_db(_engine)
return _engine
def reset_engine() -> None:
"""丢弃当前引擎(测试切换 APS_DB_PATH 后调用)。"""
global _engine, _engine_path, _SessionLocal
with _lock:
if _engine is not None:
_engine.dispose()
_engine = None
_engine_path = None
_SessionLocal = None
def get_session() -> Session:
get_engine()
assert _SessionLocal is not None
return _SessionLocal()
def init_db(engine: Engine | None = None) -> None:
"""建表(幂等)。"""
from server.db.models import Base
target = engine or get_engine()
if target.dialect.name == "sqlite":
_migrate_legacy_sqlite(target)
allow_sqlite_approval_ddl = (
target.dialect.name == "sqlite"
and os.environ.get(_APPROVAL_SQLITE_TEST_ENV) == "1"
)
runtime_tables = None
if not allow_sqlite_approval_ddl:
runtime_tables = [
table
for table in Base.metadata.sorted_tables
if table.name not in _MIGRATION_MANAGED_TABLES
]
Base.metadata.create_all(target, tables=runtime_tables)
def _migrate_legacy_sqlite(engine: Engine) -> None:
"""Small compatibility migration for databases created before tenant columns existed."""
from sqlalchemy import inspect, text
inspector = inspect(engine)
additions = {
"projects": {
"tenant_uuid": "VARCHAR(32) NOT NULL DEFAULT 'platform'",
"creator_id": "BIGINT",
"updater_id": "BIGINT",
"updated_at": "VARCHAR(32) NOT NULL DEFAULT ''",
"deleted": "BIGINT NOT NULL DEFAULT 0",
"owner_user_id": "BIGINT",
"data_version": "INTEGER NOT NULL DEFAULT 0",
},
"master_records": {
"tenant_uuid": "VARCHAR(32) NOT NULL DEFAULT 'platform'",
"creator_id": "BIGINT",
"updater_id": "BIGINT",
"updated_at": "VARCHAR(32) NOT NULL DEFAULT ''",
"deleted": "BIGINT NOT NULL DEFAULT 0",
},
"project_settings": {
"tenant_uuid": "VARCHAR(32) NOT NULL DEFAULT 'platform'",
"creator_id": "BIGINT",
"updater_id": "BIGINT",
"updated_at": "VARCHAR(32) NOT NULL DEFAULT ''",
"deleted": "BIGINT NOT NULL DEFAULT 0",
},
"routing_templates": {
"tenant_uuid": "VARCHAR(32) NOT NULL DEFAULT 'platform'",
"creator_id": "BIGINT",
"updater_id": "BIGINT",
"updated_at": "VARCHAR(32) NOT NULL DEFAULT ''",
"deleted": "BIGINT NOT NULL DEFAULT 0",
},
"routing_template_steps": {
"tenant_uuid": "VARCHAR(32) NOT NULL DEFAULT 'platform'",
"creator_id": "BIGINT",
"updater_id": "BIGINT",
"updated_at": "VARCHAR(32) NOT NULL DEFAULT ''",
"deleted": "BIGINT NOT NULL DEFAULT 0",
},
}
tables = set(inspector.get_table_names())
with engine.begin() as connection:
for table, columns in additions.items():
if table not in tables:
continue
existing = {col["name"] for col in inspector.get_columns(table)}
for name, ddl in columns.items():
if name not in existing:
connection.execute(text(f"ALTER TABLE {table} ADD COLUMN {name} {ddl}"))