"""
This program is free software: you can redistribute it and/or modify it under
the terms of the GNU General Public License as published by
the Free Software Foundation, either version 3 of the License,
or (at your option) any later version.


This program is distributed in the hope that it will be useful,
but WITHOUT ANY WARRANTY; without even the implied warranty of
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. 
See the GNU General Public License for more details.


You should have received a copy of the GNU General Public License
 along with this program.  If not, see <https://www.gnu.org/licenses/>.

Copyright © 2019 Cloud Linux Software Inc.

This software is also available under ImunifyAV commercial license,
see <https://www.imunify360.com/legal/eula>
"""
import asyncio
import os
from collections import defaultdict
from logging import getLogger

from defence360agent.contracts.config import Malware as Config
from defence360agent.contracts.messages import MessageType
from defence360agent.model.simplification import run_in_executor
from defence360agent.contracts.plugins import (
    MessageSink,
    MessageSource,
    expect,
)
from imav.malwarelib.config import MalwareScanType
from imav.malwarelib.model import MalwareHit, MalwareIgnorePath
from imav.malwarelib.tenant_path import (
    TenantPath,
    split_prefixed,
    strip_prefix,
    to_prefixed,
)
from imav.malwarelib.utils.user_list import panel_users
from imav.malwarelib.plugins.detached_scan import DetachedScanPlugin
from imav.malwarelib.scan.scanner import MalwareScanner
from defence360agent.utils import recurring_check, is_cluster

RESCAN_TYPES = (MalwareScanType.RESCAN, MalwareScanType.RESCAN_OUTDATED)

logger = getLogger(__name__)


class Scanner(MessageSink, MessageSource):
    _loop, _sink = None, None
    _targets = None
    _scan_task = None

    async def create_source(self, loop, sink):
        self._loop = loop
        self._sink = sink

        self._scan_task = self._loop.create_task(self._recurring_scan())

    async def create_sink(self, loop):
        self._targets = defaultdict(set)

    async def shutdown(self):
        self._scan_task.cancel()
        await self._scan_task

    def _process_scan_task(self, message):
        scan_type = message.get("scan_type", MalwareScanType.REALTIME)
        bucket = self._targets[scan_type]
        # Snapshot FromConfig descriptors once per batch: each read goes
        # through config_to_dict() which deepcopies the merged config, and
        # with cap=100_000 we would otherwise pay ~2 deepcopies per path.
        max_path_len = Config.MAX_PATH_LEN
        max_targets = Config.MAX_TARGETS_PER_SCAN_TYPE
        dropped_long = 0
        for path in message["filelist"]:
            if not isinstance(path, str):
                t = type(path)
                path = os.fsdecode(path)
                logger.error(
                    "Received path %s as %s instead of %s. Message: %s",
                    path,
                    t,
                    type(str),
                    message,
                )
            if len(path) > max_path_len:
                dropped_long += 1
                continue
            if len(bucket) >= max_targets:
                logger.warning(
                    "MAX_TARGETS_PER_SCAN_TYPE cap (%d) reached for "
                    "scan_type=%s; dropping remaining paths in this batch",
                    max_targets,
                    scan_type,
                )
                break
            bucket.add(path)
        if dropped_long:
            logger.warning(
                "Dropped %d path(s) exceeding MAX_PATH_LEN=%d for "
                "scan_type=%s",
                dropped_long,
                max_path_len,
                scan_type,
            )

    @expect(MessageType.MalwareScanTask)
    async def process_scan_task(self, message):
        self._process_scan_task(message)

    @expect(MessageType.MalwareRescanFiles)
    async def rescan_files(self, message):
        filelist = message["files"]
        msg = MessageType.MalwareScanTask(
            filelist=filelist, scan_type=message.get("type", "rescan")
        )
        self._process_scan_task(msg)

    @staticmethod
    async def _filter_out(targets, require_local_exists=True, registered=None):
        """Filter targets: exclude ignored paths; optionally require file to exist locally.

        In split-container (K8s) rescan, paths live on the scanner's filesystem, so we
        must not require local existence or the wrapper would never be invoked and no
        rescan dir would be created under /var/imunify360/aibolit/rescan/.

        ``registered``: registered k8s app ids — a filename whose first
        component is one of them is treated as tenant-prefixed for the
        ignore check (legacy bare paths stay bare).
        """
        result = list()
        for filename in targets:
            if require_local_exists and not os.path.exists(filename):
                continue
            check_path = filename
            if registered:
                prefix_user, bare = split_prefixed(filename)
                if prefix_user in registered:
                    check_path = TenantPath(bare, user=prefix_user)
            if await MalwareIgnorePath.is_path_ignored(check_path):
                continue
            result.append(filename)
        return result

    @staticmethod
    def _group_files_by_user(files, registered=None):
        """Look up the owner of each file from its MalwareHit.

        Returns a dict ``{user: [bare files]}``. Files whose owner cannot be
        resolved (no hit, or an empty ``hit.user``) are grouped under
        ``None`` — the caller decides the policy for unresolved files (in
        K8s rescan they are dropped with an error, since a file that names
        no tenant cannot be routed to an application).

        This is only for a bare path, which names no tenant. ``orig_file``
        is stored prefixed, so when ``registered`` app ids are given the
        lookup also tries each ``/<app_id>``-prefixed candidate.
        """

        def candidates(fnames):
            # exact form plus, on k8s, each registered app's prefixed form
            out = list(fnames)
            for user in registered or ():
                out.extend(to_prefixed(f, user) for f in fnames)
            return out

        files = list(files)
        user_files = defaultdict(set)
        matched_files = set()
        for hit in MalwareHit.get_hits(candidates(files)):
            if not hit.user:
                continue
            bare = strip_prefix(hit.orig_file, hit.user)
            user_files[hit.user].add(bare)
            # inputs arrive bare or already prefixed (e.g. a rescan for a
            # deregistered app); record both forms so a prefixed input that
            # matched is not also reported as unresolved
            matched_files.add(bare)
            matched_files.add(hit.orig_file)

        # Remaining files go into the None group (unresolved)
        for f in files:
            if f not in matched_files:
                user_files[None].add(f)

        return {user: list(paths) for user, paths in user_files.items()}

    async def _run_scan(self, file_list, scan_type, **kwargs):
        """Start a single MalwareScanner run and publish its results."""
        logger.debug("Scanning files: %s (kwargs=%s)", file_list, kwargs)
        scanner = MalwareScanner(sink=self._sink, hooks=True)
        scanner.start(file_list, scan_type=scan_type, **kwargs)
        result = await scanner.async_wait()
        if scanner is not None:
            message = await DetachedScanPlugin.aggregate_result(result)
            await self._sink.process_message(
                MessageType.MalwareScan(**message)
            )

    async def _scan_targets(self, targets, scan_type):
        if targets:
            logger.info(
                "Checking files to scan with type={}".format(scan_type)
            )

        # In K8s rescan, files are on the scanner container; do not require local existence
        is_k8s_rescan = is_cluster() and scan_type in RESCAN_TYPES
        registered = (
            {u["user"] for u in await panel_users()} if is_k8s_rescan else None
        )
        file_list = await self._filter_out(
            targets,
            require_local_exists=not is_k8s_rescan,
            registered=registered,
        )

        if not file_list:
            return

        # In K8s, run a separate scan per tenant with tenant-prefixed
        # paths in the listing — the prefix is the only identity the
        # shim/syncer can route by. Prefixed input (the canonical
        # server-side form) passes through verbatim; legacy bare paths
        # resolve their tenant via the hit lookup and get re-prefixed.
        if is_k8s_rescan:
            user_groups = defaultdict(set)
            legacy_files = []
            for filename in file_list:
                prefix_user = split_prefixed(filename)[0]
                if prefix_user in registered:
                    user_groups[prefix_user].add(filename)
                else:
                    legacy_files.append(filename)
            grouped = await run_in_executor(
                asyncio.get_event_loop(),
                self._group_files_by_user,
                legacy_files,
                registered,
            )
            for user, files in grouped.items():
                if user is not None:
                    files = (to_prefixed(f, user) for f in files)
                user_groups[user].update(files)
            unresolved = user_groups.pop(None, set())
            if unresolved:
                # A path that names no tenant cannot be prefixed, so
                # there is no safe route — drop these files and report.
                logger.error(
                    "K8s rescan: dropping %d file(s) that name no tenant and"
                    " match no hit: %s",
                    len(unresolved),
                    sorted(unresolved),
                )
            for user, files in user_groups.items():
                logger.info("K8s rescan: user=%s, files=%s", user, files)
                await self._run_scan(sorted(files), scan_type, user=user)
        else:
            await self._run_scan(file_list, scan_type)

    async def _scan(self):
        # copy set to list to prevent race conditions
        targets, self._targets = self._targets, defaultdict(set)
        for scan_type, files in targets.items():
            await self._scan_targets(files, scan_type)

    @recurring_check(Config.INOTIFY_SCAN_PERIOD)
    async def _recurring_scan(self):
        await self._scan()
