fix(migrations): enforce schema ownership safety
This commit is contained in:
129
server/src/app/db/migration_preflight.py
Normal file
129
server/src/app/db/migration_preflight.py
Normal file
@@ -0,0 +1,129 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
|
||||
from sqlalchemy import create_engine, inspect, text
|
||||
from sqlalchemy.engine import Connection, Engine
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.db.schema_ownership import MIGRATION_OWNED_TABLES
|
||||
|
||||
# 迁移脚本新增或删除自有表时,必须同步更新该映射并补充对应测试。
|
||||
MIGRATION_OWNED_TABLES_BY_REVISION: dict[str, frozenset[str]] = {
|
||||
"20260713_0001": frozenset(
|
||||
{
|
||||
"expense_cases",
|
||||
"expense_case_links",
|
||||
"business_events",
|
||||
}
|
||||
),
|
||||
"20260713_0002": frozenset(
|
||||
{
|
||||
"expense_cases",
|
||||
"expense_case_links",
|
||||
"business_events",
|
||||
"auth_sessions",
|
||||
}
|
||||
),
|
||||
}
|
||||
if MIGRATION_OWNED_TABLES_BY_REVISION["20260713_0002"] != MIGRATION_OWNED_TABLES:
|
||||
raise RuntimeError("latest Alembic revision must own the centralized migration table set")
|
||||
|
||||
|
||||
class MigrationPreflightError(RuntimeError):
|
||||
"""Raised when the database schema cannot be safely advanced by Alembic."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MigrationPreflightState:
|
||||
revision: str | None
|
||||
owned_tables: frozenset[str]
|
||||
|
||||
|
||||
def _format_tables(table_names: frozenset[str]) -> str:
|
||||
return ", ".join(sorted(table_names)) or "none"
|
||||
|
||||
|
||||
def _validate_connection(connection: Connection) -> MigrationPreflightState:
|
||||
table_names = frozenset(inspect(connection).get_table_names())
|
||||
owned_tables = table_names & MIGRATION_OWNED_TABLES
|
||||
|
||||
if "alembic_version" not in table_names:
|
||||
if owned_tables:
|
||||
raise MigrationPreflightError(
|
||||
"unversioned database contains migration-owned tables "
|
||||
f"({_format_tables(owned_tables)}); refusing to guess, stamp, or repair"
|
||||
)
|
||||
return MigrationPreflightState(revision=None, owned_tables=owned_tables)
|
||||
|
||||
revisions = tuple(
|
||||
str(revision)
|
||||
for revision in connection.execute(
|
||||
text("SELECT version_num FROM alembic_version")
|
||||
).scalars()
|
||||
)
|
||||
if not revisions:
|
||||
if owned_tables:
|
||||
raise MigrationPreflightError(
|
||||
"alembic_version has no recorded revision but migration-owned tables exist "
|
||||
f"({_format_tables(owned_tables)}); refusing to guess, stamp, or repair"
|
||||
)
|
||||
return MigrationPreflightState(revision=None, owned_tables=owned_tables)
|
||||
|
||||
if len(revisions) > 1:
|
||||
raise MigrationPreflightError(
|
||||
"multiple Alembic revisions are recorded "
|
||||
f"({', '.join(sorted(revisions))}); branch state is unsupported"
|
||||
)
|
||||
|
||||
revision = revisions[0]
|
||||
expected_tables = MIGRATION_OWNED_TABLES_BY_REVISION.get(revision)
|
||||
if expected_tables is None:
|
||||
raise MigrationPreflightError(
|
||||
f"unknown Alembic revision {revision!r}; refusing to run migrations"
|
||||
)
|
||||
|
||||
if owned_tables != expected_tables:
|
||||
missing_tables = expected_tables - owned_tables
|
||||
unexpected_tables = owned_tables - expected_tables
|
||||
raise MigrationPreflightError(
|
||||
f"migration-owned table set does not match revision {revision}: "
|
||||
f"missing={_format_tables(missing_tables)}; "
|
||||
f"unexpected={_format_tables(unexpected_tables)}"
|
||||
)
|
||||
|
||||
return MigrationPreflightState(revision=revision, owned_tables=owned_tables)
|
||||
|
||||
|
||||
def validate_migration_state(bind: Engine | Connection) -> MigrationPreflightState:
|
||||
"""Inspect the database without changing its schema or Alembic revision state."""
|
||||
|
||||
if isinstance(bind, Engine):
|
||||
with bind.connect() as connection:
|
||||
return _validate_connection(connection)
|
||||
return _validate_connection(bind)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
settings = get_settings()
|
||||
engine = create_engine(settings.resolved_database_url, pool_pre_ping=True)
|
||||
try:
|
||||
state = validate_migration_state(engine)
|
||||
except (MigrationPreflightError, SQLAlchemyError) as exc:
|
||||
print(f"Database migration preflight failed: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
revision = state.revision or "unversioned/base"
|
||||
print(
|
||||
"Database migration preflight passed: "
|
||||
f"revision={revision}; owned_tables={_format_tables(state.owned_tables)}"
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
22
server/src/app/db/schema_ownership.py
Normal file
22
server/src/app/db/schema_ownership.py
Normal file
@@ -0,0 +1,22 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.engine import Connection, Engine
|
||||
|
||||
from app.db.base import Base
|
||||
|
||||
MIGRATION_OWNED_TABLES: frozenset[str] = frozenset(
|
||||
{
|
||||
"auth_sessions",
|
||||
"expense_cases",
|
||||
"expense_case_links",
|
||||
"business_events",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def create_legacy_schema(bind: Engine | Connection) -> None:
|
||||
"""创建仍由旧 bootstrap 管理的表,不越过 Alembic 的表所有权边界。"""
|
||||
legacy_tables = [
|
||||
table for table in Base.metadata.sorted_tables if table.name not in MIGRATION_OWNED_TABLES
|
||||
]
|
||||
Base.metadata.create_all(bind=bind, tables=legacy_tables)
|
||||
@@ -2,31 +2,16 @@ from __future__ import annotations
|
||||
|
||||
import threading
|
||||
|
||||
from sqlalchemy import inspect, select, text
|
||||
from sqlalchemy import inspect, text
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.logging import get_logger
|
||||
from app.db.base import Base
|
||||
from app.db.schema_ownership import create_legacy_schema
|
||||
from app.db.session import get_session_factory
|
||||
from app.models.agent_asset import AgentAsset
|
||||
from app.services.agent_foundation_asset_helpers import AgentFoundationAssetHelperMixin
|
||||
from app.services.agent_foundation_asset_seed import AgentFoundationAssetSeedMixin
|
||||
from app.services.agent_foundation_asset_topup import AgentFoundationAssetTopUpMixin
|
||||
from app.services.agent_foundation_constants import (
|
||||
ATTACHMENT_RULE_ASSET_CODE,
|
||||
ATTACHMENT_RULE_RUNTIME_CONFIG,
|
||||
COMPANY_COMMUNICATION_RULE_SCENARIO_JSON,
|
||||
COMPANY_COMMUNICATION_RULE_VERSION,
|
||||
COMPANY_TRAVEL_RULE_SCENARIO_JSON,
|
||||
COMPANY_TRAVEL_RULE_VERSION,
|
||||
DEMO_EXPENSE_CLAIM_SIGNATURES,
|
||||
DEMO_PAYABLE_SIGNATURES,
|
||||
DEMO_RECEIVABLE_SIGNATURES,
|
||||
LEGACY_RULE_CODES,
|
||||
PLATFORM_DESTINATION_LOCATION_RULE_CODE,
|
||||
PLATFORM_DESTINATION_LOCATION_RULE_FILENAME,
|
||||
)
|
||||
from app.services.agent_foundation_digital_employee_tasks import (
|
||||
AgentFoundationDigitalEmployeeTaskMixin,
|
||||
)
|
||||
@@ -38,12 +23,6 @@ from app.services.agent_foundation_spreadsheets import AgentFoundationSpreadshee
|
||||
logger = get_logger("app.services.agent_foundation")
|
||||
_foundation_ready_lock = threading.RLock()
|
||||
_foundation_ready_keys: set[str] = set()
|
||||
MIGRATION_OWNED_TABLES = {
|
||||
"auth_sessions",
|
||||
"expense_cases",
|
||||
"expense_case_links",
|
||||
"business_events",
|
||||
}
|
||||
|
||||
|
||||
def prepare_agent_foundation() -> None:
|
||||
@@ -83,12 +62,7 @@ class AgentFoundationService(
|
||||
|
||||
def _prepare_foundation(self) -> None:
|
||||
try:
|
||||
legacy_bootstrap_tables = [
|
||||
table
|
||||
for table in Base.metadata.sorted_tables
|
||||
if table.name not in MIGRATION_OWNED_TABLES
|
||||
]
|
||||
Base.metadata.create_all(bind=self.db.get_bind(), tables=legacy_bootstrap_tables)
|
||||
create_legacy_schema(self.db.get_bind())
|
||||
self._ensure_agent_asset_schema()
|
||||
self._ensure_financial_record_schema()
|
||||
self._seed_agent_assets()
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import Any
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.db.base import Base
|
||||
from app.db.schema_ownership import create_legacy_schema
|
||||
from app.models.budget import BudgetAllocation, BudgetReservation
|
||||
from app.models.financial_record import ExpenseClaim
|
||||
from app.schemas.budget import (
|
||||
@@ -32,7 +32,7 @@ class BudgetService(BudgetPaginationMixin, BudgetSupportMixin):
|
||||
|
||||
def ensure_budget_ready(self) -> None:
|
||||
# 复用当前 Session 连接,避免在业务事务中通过 Engine 隐式提交已 flush 的单据。
|
||||
Base.metadata.create_all(bind=self.db.connection())
|
||||
create_legacy_schema(self.db.connection())
|
||||
exists = self.db.scalar(select(BudgetAllocation.id).limit(1))
|
||||
if exists:
|
||||
return
|
||||
|
||||
@@ -10,7 +10,7 @@ from sqlalchemy import or_, select
|
||||
from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from app.core.security import hash_password
|
||||
from app.db.base import Base
|
||||
from app.db.schema_ownership import create_legacy_schema
|
||||
from app.models.budget import BudgetAllocation, BudgetReservation, BudgetTransaction
|
||||
from app.models.employee import Employee
|
||||
from app.models.financial_record import ExpenseClaim, ExpenseClaimItem
|
||||
@@ -78,7 +78,7 @@ class HalfYearExpenseSimulationSeeder:
|
||||
return self._run(apply=True)
|
||||
|
||||
def _run(self, *, apply: bool) -> SimulationSummary:
|
||||
Base.metadata.create_all(bind=self.db.get_bind())
|
||||
create_legacy_schema(self.db.get_bind())
|
||||
departments = self._department_refs(apply=apply)
|
||||
current_employee_count = self._employee_count()
|
||||
planned_employees = self._build_new_employee_refs(departments, current_employee_count)
|
||||
|
||||
@@ -7,7 +7,7 @@ from sqlalchemy import or_, select
|
||||
from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from app.core.agent_enums import AgentName, AgentRunSource
|
||||
from app.db.base import Base
|
||||
from app.db.schema_ownership import create_legacy_schema
|
||||
from app.models.agent_run import AgentRun, AgentToolCall
|
||||
from app.schemas.digital_employee_dashboard import DigitalEmployeeDashboardRead
|
||||
|
||||
@@ -186,7 +186,7 @@ class DigitalEmployeeDashboardService:
|
||||
)
|
||||
|
||||
def _ensure_storage_ready(self) -> None:
|
||||
Base.metadata.create_all(bind=self.db.get_bind())
|
||||
create_legacy_schema(self.db.get_bind())
|
||||
|
||||
def _fetch_runs(self, *, start: datetime, limit: int) -> list[AgentRun]:
|
||||
stmt = (
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import Counter
|
||||
from datetime import UTC, date, datetime
|
||||
import threading
|
||||
from collections import Counter
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
@@ -11,7 +11,7 @@ from sqlalchemy.orm import Session
|
||||
from app.core.config import get_settings
|
||||
from app.core.logging import get_logger
|
||||
from app.core.security import hash_password
|
||||
from app.db.base import Base
|
||||
from app.db.schema_ownership import create_legacy_schema
|
||||
from app.db.session import get_session_factory
|
||||
from app.models.employee import Employee
|
||||
from app.models.employee_change_log import EmployeeChangeLog
|
||||
@@ -28,11 +28,9 @@ from app.schemas.employee import (
|
||||
EmployeeStatusSummaryRead,
|
||||
EmployeeUpdate,
|
||||
)
|
||||
from app.services.employee_import import EmployeeImportCoordinator
|
||||
from app.services.employee_bank_info import apply_default_bank_info
|
||||
from app.services.employee_import import EmployeeImportCoordinator
|
||||
from app.services.employee_schema import ensure_employee_schema
|
||||
from app.services.employee_serialization import serialize_employee
|
||||
from app.services.employee_spreadsheet import build_import_template_bytes
|
||||
from app.services.employee_seed import (
|
||||
CANONICAL_DEPARTMENT_CODES,
|
||||
EMPLOYEE_DEFINITIONS,
|
||||
@@ -44,6 +42,8 @@ from app.services.employee_seed import (
|
||||
ROLE_PERMISSION_MAP,
|
||||
normalize_organization_unit_code,
|
||||
)
|
||||
from app.services.employee_serialization import serialize_employee
|
||||
from app.services.employee_spreadsheet import build_import_template_bytes
|
||||
from app.services.employee_time import (
|
||||
format_date,
|
||||
format_datetime,
|
||||
@@ -108,7 +108,7 @@ class EmployeeService:
|
||||
|
||||
def _ensure_directory_ready_uncached(self) -> None:
|
||||
try:
|
||||
Base.metadata.create_all(bind=self.db.get_bind())
|
||||
create_legacy_schema(self.db.get_bind())
|
||||
ensure_employee_schema(self.db)
|
||||
self._prune_extra_seed_employees()
|
||||
self._seed_roles()
|
||||
|
||||
@@ -8,16 +8,20 @@ from datetime import datetime
|
||||
from sqlalchemy import inspect, text
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.admin_secret import legacy_admin_secret_to_password_hash, read_admin_secret, verify_admin_secret
|
||||
from app.core.admin_secret import (
|
||||
legacy_admin_secret_to_password_hash,
|
||||
read_admin_secret,
|
||||
verify_admin_secret,
|
||||
)
|
||||
from app.core.config import get_settings
|
||||
from app.core.secret_box import decrypt_secret, encrypt_secret
|
||||
from app.core.security import hash_password, verify_password
|
||||
from app.db.base import Base
|
||||
from app.db.schema_ownership import create_legacy_schema
|
||||
from app.db.session import get_session_factory
|
||||
from app.models.hermes_config import HermesTaskConfig
|
||||
from app.models.system_model_setting import SystemModelSetting
|
||||
from app.models.system_setting import SystemSetting
|
||||
from app.models.system_setting_secret import SystemSettingSecret
|
||||
from app.models.hermes_config import HermesTaskConfig
|
||||
from app.repositories.settings import SETTINGS_ROW_ID, SettingsRepository
|
||||
from app.schemas.settings import SettingsRead, SettingsWrite
|
||||
from app.services.hermes_sync import (
|
||||
@@ -161,7 +165,7 @@ class SettingsService:
|
||||
if cache_key not in self._schema_ready_keys:
|
||||
with self._schema_ready_lock:
|
||||
if cache_key not in self._schema_ready_keys:
|
||||
Base.metadata.create_all(bind=self.db.get_bind())
|
||||
create_legacy_schema(self.db.get_bind())
|
||||
self._ensure_settings_schema()
|
||||
self._schema_ready_keys.add(cache_key)
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Any
|
||||
from sqlalchemy import or_, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.db.base import Base
|
||||
from app.db.schema_ownership import create_legacy_schema
|
||||
from app.models.agent_feedback import AgentOperationFeedback
|
||||
from app.models.agent_run import AgentRun, AgentToolCall
|
||||
from app.models.user_session_metric import UserSessionMetric
|
||||
@@ -143,7 +143,7 @@ class SystemDashboardService:
|
||||
)
|
||||
|
||||
def _ensure_storage_ready(self) -> None:
|
||||
Base.metadata.create_all(bind=self.db.get_bind())
|
||||
create_legacy_schema(self.db.get_bind())
|
||||
|
||||
def _fetch_runs(self, start: datetime, *, before: datetime | None = None) -> list[_DashboardRun]:
|
||||
stmt = (
|
||||
|
||||
Reference in New Issue
Block a user