from __future__ import annotations import json import tempfile import unittest from pathlib import Path from unittest.mock import MagicMock, patch from sqlalchemy import create_engine, inspect, select, text from sqlalchemy.orm import Session, sessionmaker import app.models # noqa: F401 from app.models.base import Base from app.models.org import User from app.services.auth import verify_password def _options(*, confirm_reset: bool = False, mode: str = "reset"): from app.services.system_initializer import SystemInitializeOptions return SystemInitializeOptions( mode=mode, company_name="百华", admin_name="超级管理员", admin_phone="13800000000", admin_password="secret123", confirm_reset=confirm_reset, ) class SystemInitializerResetTest(unittest.TestCase): def setUp(self) -> None: self.engine = create_engine("sqlite+pysqlite:///:memory:", future=True) Base.metadata.create_all(self.engine) self.SessionLocal = sessionmaker(bind=self.engine, autoflush=False, autocommit=False, future=True) self.db: Session = self.SessionLocal() def tearDown(self) -> None: self.db.close() def _seed_dirty_data(self) -> None: self.db.execute( text( """ INSERT INTO sys_department (id, dept_code, dept_name, org_node_type, dept_type, status, sort_no, created_at, updated_at) VALUES (100, 'DIRTY_ROOT', '旧公司', 'COMPANY', 'ADMIN', 'ACTIVE', 0, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) """ ) ) self.db.execute( text( """ INSERT INTO hr_employee (id, employee_code, employee_name, dept_id, mobile, is_operator, is_workshop_staff, status, created_at, updated_at) VALUES (100, 'EMP_OLD', '旧管理员', 100, '13900000000', 0, 0, 'ACTIVE', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) """ ) ) self.db.execute( text( """ INSERT INTO sys_user (id, username, password_hash, employee_id, dept_id, nickname, is_super_admin, status, created_at, updated_at) VALUES (100, '13900000000', 'old', 100, 100, '旧管理员', 0, 'ACTIVE', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) """ ) ) self.db.execute( text( """ INSERT INTO personnel (phone, name, is_temporary, created_at, updated_at) VALUES ('13900000000', '旧小程序用户', 0, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) """ ) ) self.db.execute( text( """ INSERT INTO md_customer (id, customer_code, customer_name, status, created_at, updated_at) VALUES (100, 'C_OLD', '旧客户', 'ACTIVE', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) """ ) ) self.db.commit() def test_reset_requires_confirm_reset_and_does_not_clear_data(self) -> None: from app.services.system_initializer import reset_existing_database self._seed_dirty_data() with self.assertRaisesRegex(ValueError, "confirm-reset"): reset_existing_database(self.db, _options(confirm_reset=False)) self.db.rollback() self.assertEqual(self.db.execute(text("SELECT COUNT(*) FROM sys_user")).scalar_one(), 1) self.assertEqual(self.db.execute(text("SELECT COUNT(*) FROM md_customer")).scalar_one(), 1) def test_reset_clears_business_account_and_miniapp_data_then_recreates_admin(self) -> None: from app.services.system_initializer import reset_existing_database self._seed_dirty_data() result = reset_existing_database(self.db, _options(confirm_reset=True)) self.db.commit() self.assertEqual(result.mode, "reset") self.assertEqual(result.counts_before["sys_user"], 1) self.assertEqual(result.counts_before["md_customer"], 1) self.assertEqual(result.counts_before["personnel"], 1) self.assertEqual(result.counts_after["md_customer"], 0) self.assertEqual(result.counts_after["sys_user"], 1) self.assertEqual(result.counts_after["personnel"], 1) self.assertEqual(result.seeded["admin_username"], "13800000000") self.assertEqual(self.db.execute(text("SELECT COUNT(*) FROM md_customer")).scalar_one(), 0) self.assertIsNone(self.db.scalar(select(User).where(User.username == "13900000000"))) admin = self.db.scalar(select(User).where(User.username == "13800000000")) self.assertIsNotNone(admin) assert admin is not None self.assertEqual(admin.is_super_admin, 1) self.assertTrue(verify_password("secret123", admin.password_hash)) def test_reset_keeps_sqlite_foreign_keys_enabled(self) -> None: from app.services.system_initializer import reset_existing_database self.db.execute(text("PRAGMA foreign_keys = ON")) self._seed_dirty_data() reset_existing_database(self.db, _options(confirm_reset=True)) self.db.commit() self.assertEqual(self.db.execute(text("PRAGMA foreign_keys")).scalar_one(), 1) def test_write_summary_report_writes_expected_json(self) -> None: from app.services.system_initializer import SystemInitializeResult, write_summary_report result = SystemInitializeResult( mode="reset", database="erp_test", backup_path="/tmp/backup.sql", counts_before={"sys_user": 2}, counts_after={"sys_user": 1}, seeded={"admin_username": "13800000000"}, ) with tempfile.TemporaryDirectory() as tmpdir: summary_path = write_summary_report(Path(tmpdir), result) payload = json.loads(Path(summary_path).read_text(encoding="utf-8")) self.assertEqual(payload["mode"], "reset") self.assertEqual(payload["database"], "erp_test") self.assertEqual(payload["backup_path"], "/tmp/backup.sql") self.assertEqual(payload["summary_path"], summary_path) self.assertEqual(payload["counts_before"], {"sys_user": 2}) self.assertEqual(payload["counts_after"], {"sys_user": 1}) self.assertEqual(payload["seeded"], {"admin_username": "13800000000"}) self.assertRegex(payload["created_at"], r"^\d{4}-\d{2}-\d{2}T") def test_default_output_dir_uses_cleanup_backups_folder(self) -> None: from app.services.system_initializer import default_output_dir self.assertEqual(default_output_dir().name, "cleanup_backups") self.assertEqual(default_output_dir().parent.name, "outputs") def test_backup_mysql_database_uses_defaults_file_without_password_argument(self) -> None: from app.services.system_initializer import backup_mysql_database class Settings: mysql_host = "127.0.0.1" mysql_port = 3306 mysql_user = "root" mysql_password = "secret" mysql_database = "jiaheng_erp" commands: list[list[str]] = [] def fake_run(command: list[str], check: bool) -> None: commands.append(command) defaults_arg = next(part for part in command if part.startswith("--defaults-extra-file=")) defaults_path = Path(defaults_arg.split("=", 1)[1]) self.assertTrue(defaults_path.exists()) self.assertIn("password=secret", defaults_path.read_text(encoding="utf-8")) self.assertEqual(defaults_path.stat().st_mode & 0o777, 0o600) with tempfile.TemporaryDirectory() as tmpdir: with patch("app.services.system_initializer.subprocess.run", side_effect=fake_run): backup_path = backup_mysql_database(Settings(), tmpdir) self.assertTrue(backup_path.endswith(".sql")) self.assertEqual(len(commands), 1) self.assertIn("--no-tablespaces", commands[0]) self.assertFalse(any("secret" in part for part in commands[0])) defaults_arg = next(part for part in commands[0] if part.startswith("--defaults-extra-file=")) self.assertFalse(Path(defaults_arg.split("=", 1)[1]).exists()) class SystemInitializerFreshTest(unittest.TestCase): def test_fresh_create_all_builds_key_tables_and_seeds_skeleton(self) -> None: from app.services.system_initializer import fresh_initialize_database engine = create_engine("sqlite+pysqlite:///:memory:", future=True) result = fresh_initialize_database(engine, _options(mode="fresh")) created_tables = set(inspect(engine).get_table_names()) self.assertIn("sys_user", created_tables) self.assertIn("md_customer", created_tables) self.assertIn("attendance_points", created_tables) self.assertEqual(result.mode, "fresh") self.assertEqual(result.counts_after["sys_user"], 1) self.assertEqual(result.counts_after["wh_warehouse"], 6) self.assertEqual(result.seeded["admin_username"], "13800000000") with Session(engine, future=True) as db: self.assertIsNotNone(db.scalar(select(User).where(User.username == "13800000000"))) class SystemInitializerFileCleanupTest(unittest.TestCase): def test_delete_managed_files_removes_only_children_and_keeps_directories(self) -> None: from app.services.system_initializer import delete_managed_files with tempfile.TemporaryDirectory() as tmpdir: root = Path(tmpdir) archive_dir = root / "document_archives" photo_dir = root / "uploads" / "logistics" nested_dir = archive_dir / "batch" archive_dir.mkdir(parents=True) photo_dir.mkdir(parents=True) nested_dir.mkdir() (archive_dir / "a.pdf").write_text("pdf", encoding="utf-8") (nested_dir / "nested.pdf").write_text("nested", encoding="utf-8") (photo_dir / "b.png").write_text("png", encoding="utf-8") summary = delete_managed_files([archive_dir, photo_dir]) self.assertEqual(summary["deleted_file_count"], 3) self.assertEqual(summary["deleted_dir_count"], 1) self.assertTrue(archive_dir.exists()) self.assertTrue(photo_dir.exists()) self.assertEqual(list(archive_dir.iterdir()), []) self.assertEqual(list(photo_dir.iterdir()), []) def test_reset_with_backup_skips_file_cleanup_by_default(self) -> None: from app.services import system_initializer from app.services.system_initializer import SystemInitializeResult, reset_existing_database_with_backup class Settings: mysql_database = "erp_test" reset_result = SystemInitializeResult( mode="reset", counts_before={"sys_user": 2}, counts_after={"sys_user": 1}, seeded={"admin_username": "13800000000"}, ) with tempfile.TemporaryDirectory() as tmpdir: with ( patch.object(system_initializer, "backup_mysql_database", return_value="/tmp/backup.sql"), patch.object(system_initializer, "reset_existing_database", return_value=reset_result), patch.object(system_initializer, "delete_managed_files") as delete_mock, patch.object(system_initializer, "Session") as session_mock, ): session_mock.return_value.__enter__.return_value = MagicMock() result = reset_existing_database_with_backup( engine=object(), settings=Settings(), options=_options(confirm_reset=True), output_dir=tmpdir, ) delete_mock.assert_not_called() self.assertNotIn("deleted_file_count", result.seeded) def test_reset_with_backup_includes_file_cleanup_when_requested(self) -> None: from app.services import system_initializer from app.services.system_initializer import SystemInitializeResult, SystemInitializeOptions, reset_existing_database_with_backup class Settings: mysql_database = "erp_test" reset_result = SystemInitializeResult( mode="reset", counts_before={"sys_user": 2}, counts_after={"sys_user": 1}, seeded={"admin_username": "13800000000"}, ) options = SystemInitializeOptions( mode="reset", company_name="百华", admin_name="超级管理员", admin_phone="13800000000", admin_password="secret123", confirm_reset=True, delete_files=True, ) with tempfile.TemporaryDirectory() as tmpdir: with ( patch.object(system_initializer, "backup_mysql_database", return_value="/tmp/backup.sql"), patch.object(system_initializer, "reset_existing_database", return_value=reset_result), patch.object( system_initializer, "delete_managed_files", return_value={"deleted_file_count": 3, "deleted_dir_count": 2}, ) as delete_mock, patch.object(system_initializer, "Session") as session_mock, ): session_mock.return_value.__enter__.return_value = MagicMock() result = reset_existing_database_with_backup( engine=object(), settings=Settings(), options=options, output_dir=tmpdir, ) delete_mock.assert_called_once() self.assertEqual(result.seeded["deleted_file_count"], 3) self.assertEqual(result.seeded["deleted_dir_count"], 2) if __name__ == "__main__": unittest.main()