verify_full_demo_import.py 4.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106
  1. """Verify a DMS data export by importing it into an isolated scratch database."""
  2. from __future__ import annotations
  3. import argparse
  4. import os
  5. from pathlib import Path
  6. from alembic import command
  7. from alembic.config import Config
  8. from dotenv import dotenv_values
  9. from pymysql.constants import CLIENT
  10. from sqlalchemy import create_engine, text
  11. from sqlalchemy.engine import make_url
  12. TABLES = {
  13. "sys_organization": 3,
  14. "sys_user": 3,
  15. "doc_category": 21,
  16. "doc_document": 31,
  17. "doc_attachment_binding": 19,
  18. "doc_permission": 7,
  19. "sys_audit_log": 291,
  20. }
  21. def execute_script(database_url: str, path: Path) -> None:
  22. engine = create_engine(
  23. database_url,
  24. connect_args={"client_flag": CLIENT.MULTI_STATEMENTS},
  25. )
  26. connection = engine.raw_connection()
  27. try:
  28. driver_connection = connection.driver_connection
  29. driver_connection.autocommit(True)
  30. with driver_connection.cursor() as cursor:
  31. cursor.execute(path.read_text(encoding="utf-8"))
  32. while cursor.nextset():
  33. pass
  34. finally:
  35. connection.close()
  36. engine.dispose()
  37. def main() -> None:
  38. parser = argparse.ArgumentParser()
  39. parser.add_argument("--env-file", type=Path, required=True)
  40. parser.add_argument("--schema", type=Path, required=True)
  41. parser.add_argument("--data", type=Path, required=True)
  42. parser.add_argument("--backend", type=Path, required=True)
  43. parser.add_argument("--scratch-database", required=True)
  44. args = parser.parse_args()
  45. if not args.scratch_database.startswith("dms_full_demo_verify_"):
  46. raise RuntimeError("Scratch database name must start with dms_full_demo_verify_")
  47. source_url = make_url(dotenv_values(args.env_file).get("DMS_DATABASE_URL") or "")
  48. server_url = source_url.set(database=None)
  49. scratch_url = source_url.set(database=args.scratch_database)
  50. server_engine = create_engine(server_url)
  51. scratch_created = False
  52. try:
  53. with server_engine.connect().execution_options(isolation_level="AUTOCOMMIT") as connection:
  54. exists = connection.execute(
  55. text("SELECT SCHEMA_NAME FROM INFORMATION_SCHEMA.SCHEMATA WHERE SCHEMA_NAME=:name"),
  56. {"name": args.scratch_database},
  57. ).first()
  58. if exists:
  59. raise RuntimeError("Scratch database already exists; refusing to overwrite it")
  60. connection.execute(text(f"CREATE DATABASE `{args.scratch_database}` CHARACTER SET utf8mb4"))
  61. scratch_created = True
  62. scratch_url_text = scratch_url.render_as_string(hide_password=False)
  63. execute_script(scratch_url_text, args.schema.resolve())
  64. previous_directory = Path.cwd()
  65. try:
  66. os.chdir(args.backend.resolve())
  67. os.environ["DMS_DATABASE_URL"] = scratch_url_text
  68. os.environ["DMS_JWT_SECRET"] = "verification-only-secret-not-for-deployment"
  69. command.upgrade(Config("migrations/alembic.ini"), "head")
  70. finally:
  71. os.chdir(previous_directory)
  72. execute_script(scratch_url_text, args.data.resolve())
  73. scratch_engine = create_engine(scratch_url)
  74. try:
  75. with scratch_engine.connect() as connection:
  76. for table_name, expected in TABLES.items():
  77. actual = connection.execute(text(f"SELECT COUNT(*) FROM `{table_name}`")).scalar_one()
  78. if actual != expected:
  79. raise RuntimeError(f"{table_name}: expected {expected}, got {actual}")
  80. print(f"{table_name}={actual}")
  81. finally:
  82. scratch_engine.dispose()
  83. print("scratch import verification passed")
  84. finally:
  85. if scratch_created:
  86. with server_engine.connect().execution_options(isolation_level="AUTOCOMMIT") as connection:
  87. connection.execute(text(f"DROP DATABASE `{args.scratch_database}`"))
  88. print("scratch database removed")
  89. server_engine.dispose()
  90. if __name__ == "__main__":
  91. main()