import ast
import unittest
from pathlib import Path

ROOT = Path(__file__).resolve().parents[2]


def parse_module(path: str) -> ast.Module:
    return ast.parse((ROOT / path).read_text(encoding="utf-8"))


def has_function(module: ast.Module, function_name: str) -> bool:
    return any(
        isinstance(node, ast.FunctionDef) and node.name == function_name
        for node in module.body
    )


def require_permission_values(module: ast.Module) -> set[str]:
    values: set[str] = set()
    for node in ast.walk(module):
        if not isinstance(node, ast.Call):
            continue
        if not isinstance(node.func, ast.Name) or node.func.id != "require_permission":
            continue
        if not node.args or not isinstance(node.args[0], ast.Constant):
            continue
        if isinstance(node.args[0].value, str):
            values.add(node.args[0].value)
    return values


def route_handler_permission_values(module: ast.Module) -> set[str]:
    values: set[str] = set()

    def collect_from_depends_call(call: ast.Call | None) -> None:
        if not isinstance(call, ast.Call):
            return
        if not isinstance(call.func, ast.Name) or call.func.id != "Depends":
            return
        if not call.args or not isinstance(call.args[0], ast.Call):
            return
        guard_call = call.args[0]
        if not isinstance(guard_call.func, ast.Name) or guard_call.func.id != "require_permission":
            return
        if guard_call.args and isinstance(guard_call.args[0], ast.Constant) and isinstance(guard_call.args[0].value, str):
            values.add(guard_call.args[0].value)

    for node in module.body:
        if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
            continue
        if not any(
            isinstance(decorator, ast.Call)
            and isinstance(decorator.func, ast.Attribute)
            and isinstance(decorator.func.value, ast.Name)
            and decorator.func.value.id == "router"
            for decorator in node.decorator_list
        ):
            continue
        for argument in node.args.args:
            annotation = argument.annotation
            if (
                isinstance(annotation, ast.Subscript)
                and isinstance(annotation.value, ast.Name)
                and annotation.value.id == "Annotated"
            ):
                annotation_items = annotation.slice.elts if isinstance(annotation.slice, ast.Tuple) else [annotation.slice]
                for item in annotation_items[1:]:
                    collect_from_depends_call(item)

            default = None
            positional_defaults = list(node.args.defaults)
            if positional_defaults:
                default_offset = len(node.args.args) - len(positional_defaults)
                if node.args.args.index(argument) >= default_offset:
                    default = positional_defaults[node.args.args.index(argument) - default_offset]
            collect_from_depends_call(default)
    return values


def function_permission_values(module: ast.Module, function_name: str) -> set[str]:
    function = next(
        node
        for node in module.body
        if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name == function_name
    )
    return {
        call.args[0].value
        for call in ast.walk(function)
        if isinstance(call, ast.Call)
        and isinstance(call.func, ast.Name)
        and call.func.id == "require_permission"
        and call.args
        and isinstance(call.args[0], ast.Constant)
        and isinstance(call.args[0].value, str)
    }


class PermissionGuardsContractTests(unittest.TestCase):
    def test_business_routers_use_permission_guards(self):
        security = parse_module("backend/app/core/security.py")
        customers = parse_module("backend/app/api/customers.py")
        invoices = parse_module("backend/app/api/invoices.py")
        expenses = parse_module("backend/app/api/expenses.py")
        finance = parse_module("backend/app/api/finance.py")
        reminders = parse_module("backend/app/api/reminders.py")
        email_templates = parse_module("backend/app/api/email_templates.py")
        config = parse_module("backend/app/api/config.py")

        self.assertTrue(has_function(security, "require_permission"))

        self.assertTrue({"customers.read", "customers.write"}.issubset(route_handler_permission_values(customers)))
        self.assertTrue({"invoices.read", "invoices.write", "invoices.send"}.issubset(route_handler_permission_values(invoices)))
        self.assertTrue({"expenses.read", "expenses.write", "expenses.confirm"}.issubset(route_handler_permission_values(expenses)))
        self.assertTrue({"finance.read", "gst.write"}.issubset(route_handler_permission_values(finance)))
        self.assertTrue({"reminders.read", "reminders.write"}.issubset(route_handler_permission_values(reminders)))
        self.assertTrue({"email_templates.read", "email_templates.write"}.issubset(route_handler_permission_values(email_templates)))
        self.assertTrue({"settings.read", "settings.write"}.issubset(route_handler_permission_values(config)))

    def test_remaining_read_routes_enforce_their_declared_permissions(self):
        ai = parse_module("backend/app/api/ai.py")
        stats = parse_module("backend/app/api/stats.py")
        email_logs = parse_module("backend/app/api/email_logs.py")
        payroll = parse_module("backend/app/api/payroll.py")

        self.assertIn("ai_chat.read", function_permission_values(ai, "chat"))
        for function_name in (
            "get_overview",
            "get_monthly_revenue",
            "get_invoice_status_distribution",
            "get_income_by_type",
        ):
            self.assertIn("dashboard.read", function_permission_values(stats, function_name))
        self.assertIn("settings.read", function_permission_values(email_logs, "list_email_logs"))
        self.assertIn("payroll.read", function_permission_values(payroll, "get_payroll_record"))


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