| 1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556 |
- from sqlalchemy import func, select
- from app.core.config import Settings
- from app.infrastructure.mysql.core_models import (
- AdminUserRecord,
- PlanRecord,
- ProductRecord,
- ProductVersionRecord,
- )
- from app.infrastructure.mysql.sessions import (
- DatabaseName,
- create_mysql_engine,
- create_session_factory,
- )
- EXPECTED_COUNTS = {
- "admin_users": 9,
- "products": 4,
- "product_versions": 7,
- "plans": 10,
- }
- def read_seed_counts(settings: Settings | None = None) -> dict[str, int]:
- resolved_settings = settings or Settings()
- engine = create_mysql_engine(resolved_settings, DatabaseName.CORE)
- factory = create_session_factory(engine)
- models = {
- "admin_users": AdminUserRecord,
- "products": ProductRecord,
- "product_versions": ProductVersionRecord,
- "plans": PlanRecord,
- }
- with factory() as session:
- counts = {
- name: session.scalar(select(func.count()).select_from(model)) or 0
- for name, model in models.items()
- }
- engine.dispose()
- return counts
- def main() -> None:
- counts = read_seed_counts()
- failures = {
- name: (counts[name], expected)
- for name, expected in EXPECTED_COUNTS.items()
- if counts[name] != expected
- }
- if failures:
- raise SystemExit(f"种子数据数量不符合固定契约:{failures}")
- print(f"种子数据验证通过:{counts}")
- if __name__ == "__main__":
- main()
|