verify_seed.py 1.4 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556
  1. from sqlalchemy import func, select
  2. from app.core.config import Settings
  3. from app.infrastructure.mysql.core_models import (
  4. AdminUserRecord,
  5. PlanRecord,
  6. ProductRecord,
  7. ProductVersionRecord,
  8. )
  9. from app.infrastructure.mysql.sessions import (
  10. DatabaseName,
  11. create_mysql_engine,
  12. create_session_factory,
  13. )
  14. EXPECTED_COUNTS = {
  15. "admin_users": 9,
  16. "products": 4,
  17. "product_versions": 7,
  18. "plans": 10,
  19. }
  20. def read_seed_counts(settings: Settings | None = None) -> dict[str, int]:
  21. resolved_settings = settings or Settings()
  22. engine = create_mysql_engine(resolved_settings, DatabaseName.CORE)
  23. factory = create_session_factory(engine)
  24. models = {
  25. "admin_users": AdminUserRecord,
  26. "products": ProductRecord,
  27. "product_versions": ProductVersionRecord,
  28. "plans": PlanRecord,
  29. }
  30. with factory() as session:
  31. counts = {
  32. name: session.scalar(select(func.count()).select_from(model)) or 0
  33. for name, model in models.items()
  34. }
  35. engine.dispose()
  36. return counts
  37. def main() -> None:
  38. counts = read_seed_counts()
  39. failures = {
  40. name: (counts[name], expected)
  41. for name, expected in EXPECTED_COUNTS.items()
  42. if counts[name] != expected
  43. }
  44. if failures:
  45. raise SystemExit(f"种子数据数量不符合固定契约:{failures}")
  46. print(f"种子数据验证通过:{counts}")
  47. if __name__ == "__main__":
  48. main()