fix(migrations): enforce schema ownership safety

This commit is contained in:
caoxiaozhu
2026-07-14 09:23:34 +08:00
parent 1347366b95
commit 11275e4ba6
18 changed files with 755 additions and 57 deletions

View 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())

View 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)

View File

@@ -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()

View File

@@ -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

View File

@@ -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)

View File

@@ -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 = (

View File

@@ -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()

View File

@@ -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)

View File

@@ -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 = (