import importlib.util
import sqlite3
import unittest
from pathlib import Path

spec = importlib.util.spec_from_file_location("master_worker", Path(__file__).resolve().parents[2] / "scripts/master-api/worker.py")
worker = importlib.util.module_from_spec(spec)
spec.loader.exec_module(worker)


class Cursor:
    """Exercise parameterized write SQL against an isolated SQLite database."""
    def __init__(self, connection):
        self.cursor = connection.cursor()

    def execute(self, sql, args=()):
        return self.cursor.execute(sql.replace("%s", "?"), args)

    @property
    def lastrowid(self):
        return self.cursor.lastrowid


def record(i, project="project-a", **overrides):
    return dict(source_id=i, email=f"account{i}@example.test", project_name=project,
                client_id="client-" + project, client_secret="test-secret", redirect_uri="https://example.test/callback",
                access_token="new-access", refresh_token="new-refresh", **overrides)


class WorkerTest(unittest.TestCase):
    def setUp(self):
        self.db = sqlite3.connect(":memory:")
        self.db.row_factory = sqlite3.Row
        self.db.executescript("""
            PRAGMA foreign_keys = ON;
            CREATE TABLE servers (id INTEGER PRIMARY KEY);
            INSERT INTO servers VALUES (147), (188), (15);
            CREATE TABLE google_credentials (
                id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER, project_name TEXT UNIQUE,
                server_id INTEGER REFERENCES servers(id), client_id TEXT, client_secret TEXT,
                redirect_uri TEXT, refreshedd_at TEXT, created_at TEXT, updated_at TEXT);
            CREATE TABLE google_api_accounts (
                id INTEGER PRIMARY KEY AUTOINCREMENT, google_credential_id INTEGER REFERENCES google_credentials(id),
                email TEXT UNIQUE COLLATE NOCASE, name TEXT, access_token TEXT, token_type TEXT, refresh_token TEXT,
                token_expires_at TEXT, last_connection_at TEXT, created_at TEXT, updated_at TEXT);
        """)

    def tearDown(self):
        self.db.close()

    def snapshot(self):
        projects = [dict(row) for row in self.db.execute("SELECT id, project_name, server_id FROM google_credentials ORDER BY id")]
        accounts = [dict(row) for row in self.db.execute("SELECT id, google_credential_id, email FROM google_api_accounts ORDER BY id")]
        return projects, accounts

    def request(self, records=None, **options):
        return {"operation": "import", "mode": "append", "servers": [147, 188], "records": records or [record(1)], **options}

    def apply(self, request):
        plan = worker.make_plan(*self.snapshot(), request)
        worker.apply_plan(Cursor(self.db), plan, request)
        self.db.commit()
        return plan

    def test_insert_preserves_project_link_and_user_one(self):
        self.apply(self.request([record(1), record(2)]))
        project = self.db.execute("SELECT * FROM google_credentials").fetchone()
        accounts = self.db.execute("SELECT * FROM google_api_accounts").fetchall()
        self.assertEqual(project["user_id"], 1)
        self.assertEqual(project["client_secret"], "test-secret")
        self.assertEqual([row["google_credential_id"] for row in accounts], [project["id"]] * 2)
        self.assertTrue(all(len(row["token_type"]) == 6 for row in accounts))

    def test_duplicates_update_tokens_without_adding_rows_or_replacing_creation_metadata(self):
        self.apply(self.request())
        self.db.execute("UPDATE google_api_accounts SET name = 'Existing name', token_type = 'abcdef', created_at = 'old-date'")
        self.db.commit()
        row = record(1)
        row.update(email="ACCOUNT1@example.test", refresh_token="updated-refresh", access_token="updated-access")
        plan = self.apply(self.request([row]))
        self.assertEqual(plan["summary"]["updated"], 1)
        self.assertEqual(plan["summary"]["inserted"], 0)
        accounts = self.db.execute("SELECT * FROM google_api_accounts").fetchall()
        self.assertEqual(len(accounts), 1)
        self.assertEqual(accounts[0]["refresh_token"], "updated-refresh")
        self.assertEqual(accounts[0]["name"], "Existing name")
        self.assertEqual(accounts[0]["token_type"], "abcdef")
        self.assertEqual(accounts[0]["created_at"], "old-date")

    def test_replace_clears_both_complete_tables(self):
        self.apply(self.request([record(1), record(2, "project-b")]))
        plan = self.apply(self.request([record(3, "project-c")], mode="replace"))
        projects, accounts = self.snapshot()
        self.assertEqual([p["project_name"] for p in projects], ["project-c"])
        self.assertEqual([a["email"] for a in accounts], ["account3@example.test"])
        self.assertEqual(plan["summary"]["deleted_accounts"], 2)
        self.assertEqual(plan["summary"]["deleted_projects"], 2)

    def test_failure_after_delete_rolls_back_both_tables(self):
        self.apply(self.request())
        before = self.snapshot()
        request = self.request([record(2)], mode="replace", servers=[999])
        plan = worker.make_plan(*before, request)
        with self.assertRaises(sqlite3.IntegrityError):
            try:
                worker.apply_plan(Cursor(self.db), plan, request)
            except Exception:
                self.db.rollback()
                raise
        self.assertEqual(before, self.snapshot())

    def test_whole_projects_remain_together_and_larger_projects_are_balanced_first(self):
        records = [record(i, "project-a" if i <= 7 else "project-b" if i <= 12 else "project-c") for i in range(1, 16)]
        plan = self.apply(self.request(records))
        self.assertEqual([row["accounts"] for row in plan["summary"]["distribution"]], [7, 8])
        self.assertEqual(len(self.snapshot()[0]), 3)
        self.assertEqual(plan["summary"]["affected_accounts"], 15)

    def test_seventy_eight_single_account_projects_distribute_across_ten_servers(self):
        request = self.request([record(i, f"project-{i}") for i in range(1, 79)], servers=list(range(1, 11)))
        plan = worker.make_plan([], [], request)
        self.assertEqual(sorted(row["accounts"] for row in plan["summary"]["distribution"]), [7, 7] + [8] * 8)

    def test_rebalance_changes_only_assignment_and_update_time(self):
        self.apply(self.request([record(i, f"project-{i}") for i in range(1, 11)], servers=[147]))
        before = [tuple(row) for row in self.db.execute("SELECT * FROM google_api_accounts ORDER BY id")]
        plan = self.apply(self.request(operation="rebalance", servers=[188, 15]))
        after = [tuple(row) for row in self.db.execute("SELECT * FROM google_api_accounts ORDER BY id")]
        self.assertEqual(before, after)
        self.assertEqual([row["accounts"] for row in plan["summary"]["distribution"]], [5, 5])

    def test_append_accounts_for_existing_project_includes_all_its_accounts_in_distribution(self):
        self.apply(self.request([record(1), record(2)]))
        plan = self.apply(self.request([record(3)], servers=[188]))
        self.assertEqual(plan["summary"]["selected_accounts"], 1)
        self.assertEqual(plan["summary"]["affected_accounts"], 3)
        self.assertEqual(plan["summary"]["distribution"][0]["accounts"], 3)

    def test_preview_plan_does_not_write_and_does_not_include_credentials_in_summary(self):
        plan = worker.make_plan(*self.snapshot(), self.request())
        self.assertEqual(self.snapshot(), ([], []))
        summary = worker.json.dumps(plan["summary"])
        for secret in ["test-secret", "new-access", "new-refresh"]:
            self.assertNotIn(secret, summary)

    def test_rebalance_fingerprint_changes_when_assignments_change(self):
        self.apply(self.request())
        before = worker.fingerprint(*self.snapshot())
        self.apply(self.request(operation="rebalance", servers=[188]))
        self.assertNotEqual(before, worker.fingerprint(*self.snapshot()))


if __name__ == "__main__":
    unittest.main()
