import asyncio

import aiohttp

from defence360agent.rpc_tools import ValidationError, lookup
from defence360agent.utils import Scope
from im360.subsys.panels import hosting_panel

ADMIN_USERNAME = "admin"

CALLER_HEADER = "X-Imunify-Caller"
DOMAIN_OWNED_HEADER = "X-Imunify-Domain-Owned"

DOMAIN_OWNED_YES = "yes"
DOMAIN_OWNED_NO = "no"

_BASE_URL = "http://uam"
_TIMEOUT = aiohttp.ClientTimeout(total=10)
_SOCKET_PATH = "/var/run/defence360agent/uam.sock"


def _query(**kwargs):
    return {
        key: value for key, value in kwargs.items() if value not in (None, "")
    }


def _wafd_error(status, body):
    # wafd errors use a {code, message} envelope; fall back to a generic
    # message/code when the body is missing or not that shape so a malformed
    # response surfaces as a clean RPC error rather than a KeyError/decode error.
    code = "internal"
    message = f"wafd request failed with status {status}"
    if isinstance(body, dict):
        code = body.get("code") or code
        message = body.get("message") or message
    return ValidationError(message, extra_data={"code": code})


class WafdUAMClient:
    def __init__(self, socket_path=None):
        self._socket_path = socket_path or _SOCKET_PATH

    async def _request(
        self,
        method,
        path,
        caller,
        domain_owned,
        *,
        params=None,
        json_body=None,
    ):
        headers = {
            CALLER_HEADER: caller,
            DOMAIN_OWNED_HEADER: domain_owned,
        }
        try:
            async with aiohttp.ClientSession(
                connector=aiohttp.UnixConnector(path=self._socket_path),
                timeout=_TIMEOUT,
            ) as session:
                async with session.request(
                    method,
                    _BASE_URL + path,
                    params=params,
                    json=json_body,
                    headers=headers,
                ) as response:
                    if response.status == 204:
                        return None
                    try:
                        body = await response.json(content_type=None)
                    except ValueError:
                        body = None
                        parsed = False
                    else:
                        parsed = True
                    if 200 <= response.status < 300:
                        if not parsed:
                            raise ValidationError(
                                "wafd returned a non-JSON success response",
                                extra_data={"code": "internal"},
                            )
                        return body
                    raise _wafd_error(response.status, body)
        except (aiohttp.ClientError, asyncio.TimeoutError) as e:
            raise ValidationError(str(e), extra_data={"code": "internal"})

    async def create_rule(self, caller, domain_owned, body):
        return await self._request(
            "POST", "/uam/v1/rules", caller, domain_owned, json_body=body
        )

    async def update_rule(self, caller, domain_owned, rule_id, body):
        return await self._request(
            "PATCH",
            f"/uam/v1/rules/{rule_id}",
            caller,
            domain_owned,
            json_body=body,
        )

    async def delete_rule(self, caller, domain_owned, rule_id):
        return await self._request(
            "DELETE",
            f"/uam/v1/rules/{rule_id}",
            caller,
            domain_owned,
        )

    async def list_rules(self, caller, domain_owned, owner="", domain=""):
        return await self._request(
            "GET",
            "/uam/v1/rules",
            caller,
            domain_owned,
            params=_query(owner=owner, domain=domain),
        )

    async def test_url(self, caller, url):
        return await self._request(
            "GET",
            "/uam/v1/test",
            caller,
            DOMAIN_OWNED_YES,
            params=_query(url=url),
        )

    async def get_counters(
        self,
        caller,
        domain_owned,
        *,
        owner="",
        domain="",
        rule_id="",
        since="",
    ):
        return await self._request(
            "GET",
            "/uam/v1/counters",
            caller,
            domain_owned,
            params=_query(
                owner=owner, domain=domain, rule_id=rule_id, since=since
            ),
        )

    async def get_settings(self, caller):
        return await self._request(
            "GET", "/uam/v1/service/settings", caller, DOMAIN_OWNED_YES
        )

    async def set_settings(self, caller, enabled):
        return await self._request(
            "PUT",
            "/uam/v1/service/settings",
            caller,
            DOMAIN_OWNED_YES,
            json_body={"enabled": enabled},
        )

    async def get_visibility(self, caller):
        return await self._request(
            "GET", "/uam/v1/service/visibility", caller, DOMAIN_OWNED_YES
        )

    async def set_visibility(self, caller, allowed_for_users):
        return await self._request(
            "PUT",
            "/uam/v1/service/visibility",
            caller,
            DOMAIN_OWNED_YES,
            json_body={"allowed_for_users": allowed_for_users},
        )


def _caller(user):
    # The RPC server injects params["user"] only for non-root callers
    # (SO_PEERCRED username or JWT claim); the server admin connects over the
    # root socket where it stays unset, so an absent/empty user is the admin.
    return user or ADMIN_USERNAME


def _rule_id(value):
    # rule_id is interpolated into the wafd path; require a positive integer so a
    # missing/None value can't build /uam/v1/rules/None and a crafted string
    # can't inject query/path segments (?, #, /, ..) into the request target.
    try:
        rule_id = int(value)
    except (TypeError, ValueError):
        rule_id = 0
    if rule_id < 1:
        raise ValidationError(
            "rule_id must be a positive integer",
            extra_data={"code": "validation_failed"},
        )
    return rule_id


def _scope_owner(caller, owner):
    # A non-admin caller may only read its own rules; never let a user-supplied
    # owner widen a list/counters query beyond the authenticated identity.
    return owner if caller == ADMIN_USERNAME else caller


def _base_domain(domain):
    # The domain a wildcard rule scopes to: strip the leading `*.` or `.`
    # wildcard prefix (the two forms the matcher understands). A plain domain is
    # returned unchanged. Ownership is decided on this base domain.
    for prefix in ("*.", "."):
        if domain.startswith(prefix):
            return domain[len(prefix) :]
    return domain


class UAMEndpoints(lookup.CommonEndpoints):
    SCOPE = Scope.IM360

    def __init__(self, sink):
        super().__init__(sink)
        self.hp = hosting_panel.HostingPanel()
        self._wafd = WafdUAMClient()

    async def _domain_owned(self, caller, domain):
        if not domain:
            return DOMAIN_OWNED_YES
        owned = (await self.hp.get_domains_per_user()).get(caller, [])
        # A wildcard rule (`*.example.com` / `.example.com`) is owned exactly
        # when its base domain is, so an unprivileged user may put a wildcard
        # of a domain they own under attack -- the same base the Web UI offered
        # the wildcard for. An unowned base still resolves to "no".
        return (
            DOMAIN_OWNED_YES
            if _base_domain(domain) in owned
            else DOMAIN_OWNED_NO
        )

    @lookup.bind("uam", "add")
    async def add(self, user=None, domain=None, **kwargs):
        caller = _caller(user)
        domain_owned = await self._domain_owned(caller, domain)
        return await self._wafd.create_rule(
            caller, domain_owned, {"domain": domain, **kwargs}
        )

    @lookup.bind("uam", "delete")
    async def delete(self, user=None, rule_id=None, domain=""):
        caller = _caller(user)
        rule_id = _rule_id(rule_id)
        domain_owned = await self._domain_owned(caller, domain)
        return await self._wafd.delete_rule(caller, domain_owned, rule_id)

    @lookup.bind("uam", "edit")
    async def edit(self, user=None, rule_id=None, domain="", **kwargs):
        caller = _caller(user)
        rule_id = _rule_id(rule_id)
        domain_owned = await self._domain_owned(caller, domain)
        return await self._wafd.update_rule(
            caller, domain_owned, rule_id, dict(kwargs)
        )

    @lookup.bind("uam", "list")
    async def list_rules(self, user=None, domain="", owner=""):
        caller = _caller(user)
        owner = _scope_owner(caller, owner)
        domain_owned = await self._domain_owned(caller, domain)
        return await self._wafd.list_rules(
            caller, domain_owned, owner=owner, domain=domain
        )

    @lookup.bind("uam", "test")
    async def test_url(self, user=None, url=None):
        return await self._wafd.test_url(_caller(user), url)

    @lookup.bind("uam", "counters")
    async def counters(
        self, user=None, domain="", owner="", rule_id="", since=""
    ):
        caller = _caller(user)
        owner = _scope_owner(caller, owner)
        domain_owned = await self._domain_owned(caller, domain)
        return await self._wafd.get_counters(
            caller,
            domain_owned,
            owner=owner,
            domain=domain,
            rule_id=rule_id,
            since=since,
        )


class UAMServiceStatusEndpoints(lookup.CommonEndpoints):
    SCOPE = Scope.IM360

    def __init__(self, sink):
        super().__init__(sink)
        self.hp = hosting_panel.HostingPanel()
        self._wafd = WafdUAMClient()

    @lookup.bind("uam", "service", "settings", "get")
    async def settings_get(self, user=None):
        return await self._wafd.get_settings(_caller(user))

    @lookup.bind("uam", "service", "visibility", "get")
    async def visibility_get(self, user=None):
        return await self._wafd.get_visibility(_caller(user))

    @lookup.bind("uam", "domains")
    async def domains(self, user=None):
        caller = _caller(user)
        per_user = await self.hp.get_domains_per_user()
        if caller == ADMIN_USERNAME:
            found = {d for domains in per_user.values() for d in domains}
        else:
            found = set(per_user.get(caller, []))
        return {"items": sorted(found)}


class UAMServiceEndpoints(lookup.RootEndpoints):
    SCOPE = Scope.IM360

    def __init__(self, sink):
        super().__init__(sink)
        self._wafd = WafdUAMClient()

    @lookup.bind("uam", "service", "settings", "set")
    async def settings_set(self, enabled):
        return await self._wafd.set_settings(ADMIN_USERNAME, enabled)

    @lookup.bind("uam", "service", "visibility", "set")
    async def visibility_set(self, allowed_for_users):
        return await self._wafd.set_visibility(
            ADMIN_USERNAME, allowed_for_users
        )
