"""
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 re
import shutil
import time
import uuid
from functools import partial
from logging import getLogger
from pathlib import Path
from typing import Dict, Iterable, List, Optional, Tuple

from defence360agent import utils
from defence360agent.api import inactivity
from defence360agent.contracts.config import (
    Malware as Config,
)
from defence360agent.contracts.config import (
    MalwareTune,
    MyImunifyConfig,
)
from defence360agent.contracts.hook_events import HookEvent
from defence360agent.contracts.license import LicenseCLN
from defence360agent.contracts.messages import MessageType
from defence360agent.contracts.permissions import myimunify_protection_enabled
from defence360agent.contracts.plugins import (
    MessageSink,
    MessageSource,
    expect,
)
from defence360agent.internals.global_scope import g
from defence360agent.subsys.persistent_state import register_lock_file
from defence360agent.utils import (
    Scope,
    is_system_user,
    nice_iterator,
    recurring_check,
    safe_cancel_task,
    split_for_chunk,
)
from defence360agent.utils.check_lock import check_lock
from defence360agent.utils.common import DAY, MINUTE, rate_limit
from imav.contracts.messages import MalwareDatabaseRestoreTask
from imav.malwarelib.cleanup.cleaner import (
    CleanupResult,
    MalwareCleaner,
    MalwareCleanupProxy,
)
from imav.malwarelib.cleanup.storage import CleanupStorage
from imav.malwarelib.config import (
    MalwareHitStatus,
    MalwareScanResourceType,
    MalwareScanType,
)
from imav.malwarelib.model import MalwareHistory, MalwareHit
from imav.malwarelib.scan import ScanAlreadyCompleteError
from imav.malwarelib.scan.crontab import is_crontab
from imav.malwarelib.scan.mds.cleaner import MalwareDatabaseCleaner
from imav.malwarelib.scan.mds.detached import (
    MDSDetachedCleanup,
    MDSDetachedRestore,
)
from imav.malwarelib.scan.mds.restore import MalwareDatabaseRestore
from imav.malwarelib.subsys.malware import HackerTrapHitsSaver, MalwareAction
from imav.malwarelib.tenant_path import gather_exists, safe_exists
from imav.malwarelib.utils import malware_response
from imav.malwarelib.utils.user_list import (
    get_username_by_uid,
    is_uid,
)
from imav.malwarelib.model import MalwareIgnorePath

logger = getLogger(__name__)

COUNT_OF_ATTEMPTS_TO_CLEANUP_PER_DAY = 4
# A path we keep cleaning successfully still burns a full backup copy on
# every failed round, so the reset above needs a ceiling of its own.
COUNT_OF_FAILURES_TO_CLEANUP_PER_DAY = 16
RESCAN_ATTEMPTS = 2
LOCK_FILE = register_lock_file("cleanup", Scope.AV_IM360)

# accounts web servers execute client code as: www-data, apache, EL7 nobody.
# They stay cleanable on our side; the released cleaner still refuses 48/99
# itself until its guard adopts the same set, so until then they behave as
# before this guard existed.
WEB_EXEC_UIDS = frozenset((33, 48, 99))


def _system_owned(uid: int) -> bool:
    return uid not in WEB_EXEC_UIDS and is_system_user(uid)


def _refused_by_owner_guard(hit) -> bool:
    return (
        is_uid(hit.owner)
        and _system_owned(int(hit.owner))
        and not is_crontab(hit.orig_file_path)
    )


_group_by_status = partial(MalwareHit.group_by_attribute, attribute="status")
_group_by_user = partial(MalwareHit.group_by_attribute, attribute="owner")

throttled_log_error = rate_limit(period=DAY, on_drop=logger.warning)(
    logger.error
)


def filter_cleanable(hits: Iterable[MalwareHit]) -> Iterable:
    return (hit for hit in hits if hit.status == MalwareHitStatus.FOUND)


def drop_stored_originals(hits: list) -> list:
    """Forget hits that point inside the cleanup storage.

    They are copies of files detected elsewhere, and they can never be
    cleaned: only root may enter the storage, while the cleanup drops
    privileges to the owner of the file the copy was made from.
    """
    stored = [hit for hit in hits if CleanupStorage.contains(hit.orig_file)]
    if not stored:
        return hits
    logger.warning(
        "Removing %s hit(s) inside %s: a stored original is not a detection"
        " of its own",
        len(stored),
        CleanupStorage.path,
    )
    MalwareHit.delete_instances(stored)
    stored_ids = {hit.id for hit in stored}
    return [hit for hit in hits if hit.id not in stored_ids]


class Cleanup(MessageSink, MessageSource):
    def __init__(self):
        self._cleanup_task = None
        self._store_original_task = None
        self._running = False
        self._loop = None
        self._sink = None
        self._proxy = None
        self._cleaner = None

    async def create_source(self, loop, sink):
        self._loop = loop
        self._sink = sink
        self._proxy = MalwareCleanupProxy()
        self._cleaner = MalwareCleaner(loop=loop, sink=sink)
        self._cleanup_task = loop.create_task(self.cleanup())

    async def create_sink(self, loop):
        pass

    async def shutdown(self):
        if self._cleanup_task:
            await safe_cancel_task(self._cleanup_task)

    @expect(MessageType.MalwareCleanupTask)
    async def process_cleanup_task(self, message: Dict):
        cause = message.get("cause")
        initiator = message.get("initiator")
        post_action = message.get("post_action")
        scan_id = message.get("scan_id")
        standard_only = message.get("standard_only")

        manual_cleanup = cause is None
        # In case another scan already found some of the hits
        # and the cleanup for them has started.
        origin_hits_num = len(message["hits"])
        hits = MalwareHit.refresh_hits(
            message["hits"], include_scan_info=not manual_cleanup
        )
        hits = filter_cleanable(hits)
        hits = [
            hit
            for hit in hits
            if hit.resource_type == MalwareScanResourceType.FILE.value
        ]
        hits = drop_stored_originals(hits)
        if (
            not manual_cleanup
        ):  # don't use any limits when run cleanup manually
            rescan_hits, hits = self._split_hits_by_scan_type(
                hits, [MalwareScanType.RESCAN, MalwareScanType.RESCAN_OUTDATED]
            )
            rescan_hits = await self._filter_rescan_hits(rescan_hits)
            hits = rescan_hits + await self._filter_failed_to_cleanup_hits(
                hits
            )
        if filtered := origin_hits_num - len(hits):
            logger.info(
                "%s/%s hits filtered before cleanup",
                filtered,
                origin_hits_num,
            )
        self._store_original_task = self._loop.create_task(
            self._store_original(
                hits, cause, initiator, post_action, scan_id, standard_only
            )
        )

    @staticmethod
    def _split_hits_by_scan_type(
        hits: list, scan_types: List[MalwareScanType]
    ) -> Tuple[list, list]:
        target_hits, other_hits = [], []
        for hit in hits:
            if hit.scanid.type in scan_types:
                target_hits.append(hit)
            else:
                other_hits.append(hit)
        return target_hits, other_hits

    async def _filter_failed_to_cleanup(
        self,
        hits: list,
        *,
        time_range: float,
        allowed_total: int,
        allowed_after_success: Optional[int] = None,
        report_give_up: bool = False,
    ) -> list:
        hits_to_clean = []
        suppressed = []
        if hits:
            since = time.time() - time_range
            attempts = {}
            for hits_chunk in split_for_chunk(hits, chunk_size=200):
                counted = MalwareHistory.get_failed_cleanup_attempts(
                    [hit.orig_file for hit in hits_chunk], since=since
                )
                for path, total, after_success in counted:
                    attempts[path] = (total, after_success)
                await asyncio.sleep(0)
            for hit in hits:
                total, after_success = attempts.get(hit.orig_file, (0, 0))
                exhausted = total >= allowed_total or (
                    allowed_after_success is not None
                    and after_success >= allowed_after_success
                )
                if exhausted:
                    throttled_log_error(
                        "Skip cleanup file '%s', since there are too many "
                        "attempts to cleanup it in %s sec [%s, %s since the "
                        "last successful cleanup]",
                        hit.orig_file,
                        time_range,
                        total,
                        after_success,
                    )
                    if hit.orig_file not in suppressed:
                        suppressed.append(hit.orig_file)
                    continue
                hits_to_clean.append(hit)
        if suppressed and report_give_up:
            await self._report_suppressed_cleanup(suppressed, time_range)
        return hits_to_clean

    async def _report_suppressed_cleanup(self, paths: list, time_range: float):
        # The only signal that the agent decided not to try: without it a
        # path that is being skipped looks exactly like one with nothing
        # left to do.
        await self._sink.process_message(
            MessageType.CleanupFailed(
                message=(
                    "Automatic cleanup suppressed for {} path(s) that"
                    " exhausted the attempt budget in {} sec: {}".format(
                        len(paths), time_range, ", ".join(paths[:3])
                    )
                ),
                exception="CleanupAttemptsExhausted",
                err="",
                timestamp=int(time.time()),
            )
        )

    async def _filter_rescan_hits(self, hits: list) -> list:
        # A burst guard, not a give-up decision: a success minutes ago is no
        # reason to keep hammering the same path, so no reset applies here.
        return await self._filter_failed_to_cleanup(
            hits, time_range=5 * MINUTE, allowed_total=RESCAN_ATTEMPTS
        )

    async def _filter_failed_to_cleanup_hits(self, hits: list) -> list:
        return await self._filter_failed_to_cleanup(
            hits,
            time_range=DAY,
            report_give_up=True,
            allowed_total=(
                MalwareTune.CLEANUP_FAILURES_PER_DAY
                or COUNT_OF_FAILURES_TO_CLEANUP_PER_DAY
            ),
            allowed_after_success=(
                MalwareTune.CLEANUP_ATTEMPTS_AFTER_SUCCESS
                or COUNT_OF_ATTEMPTS_TO_CLEANUP_PER_DAY
            ),
        )

    async def _store_original(
        self, hits, cause, initiator, post_action, scan_id, standard_only
    ):
        MalwareHit.set_status(hits, MalwareHitStatus.CLEANUP_STARTED)
        original_status = _group_by_status(hits)
        # Re-check ignore list after marking as started, before any destructive action
        hits_to_keep = []
        ignored_hits = []
        for hit in hits:
            try:
                if await MalwareIgnorePath.is_path_ignored(hit.orig_file):
                    ignored_hits.append(hit)
                else:
                    hits_to_keep.append(hit)
            except Exception as exc:
                # Be conservative: if ignore check fails, keep the hit
                logger.exception(
                    "Ignore re-check failed for file %s: %s; keeping hit",
                    hit.orig_file,
                    exc,
                )
                hits_to_keep.append(hit)
        if ignored_hits:
            # Transition ignored hits to a terminal non-in-progress state
            MalwareHit.set_status(ignored_hits, MalwareHitStatus.FOUND)
            # Do not proceed with storage/cleanup for ignored hits
            hits = hits_to_keep
        # _clean_files refuses these before the cleaner ever runs, so a
        # backup copy of them is work nobody will read.
        refused = [hit for hit in hits if _refused_by_owner_guard(hit)]
        if refused:
            refused_set = set(refused)
            hits = [hit for hit in hits if hit not in refused_set]
        with inactivity.track.task("cleanup_storage"):
            succeeded, failed, not_exist = await CleanupStorage.store_all(hits)
        for hit, exc in failed.items():
            await self._sink.process_message(
                MessageType.CleanupFailed(
                    message=(
                        "Failed to store the original from {} to {}".format(
                            hit.orig_file, CleanupStorage.path
                        )
                    ),
                    exception=type(exc).__name__,
                    err=str(exc),
                    timestamp=int(time.time()),
                )
            )

        self._add_to_proxy(
            # succeeded may be a query, and + would union it, not extend it
            [*succeeded, *refused] if refused else succeeded,
            cause,
            initiator,
            post_action,
            scan_id,
            standard_only,
        )

        failed_ids = {hit.id for hit in failed}
        for status, hit_list in original_status.items():
            MalwareHit.set_status(
                [hit for hit in hit_list if hit.id in failed_ids], status
            )
        MalwareHit.delete_instances(not_exist)
        await MalwareAction.not_exist(
            not_exist, cause=cause, initiator=initiator
        )

    def _add_to_proxy(
        self, hits, cause, initiator, post_action, scan_id, standard_only
    ):
        standard_only_hits = []
        advanced_hits = []
        for hit in hits:
            standard_only_user = decide_if_standard_signatures_only(
                initiator, standard_only
            )
            if standard_only_user:
                standard_only_hits.append(hit)
            else:
                advanced_hits.append(hit)

        self._proxy.add(
            cause,
            initiator,
            post_action,
            scan_id,
            True,
            standard_only_hits,
        )
        self._proxy.add(
            cause,
            initiator,
            post_action,
            scan_id,
            standard_only,  # None if default action otherwise False
            advanced_hits,
        )

    @staticmethod
    def _user_hits(hits):
        user_hits = _group_by_user(hits)
        return user_hits

    def _cloud_assisted_hits(self):
        action_hits = self._proxy.flush()

        for (
            cause,
            initiator,
            post_action,
            scan_id,
            standard_only,
            all_hits,
        ) in action_hits:
            blacklist = [
                hit
                for hit in all_hits
                if re.match(r"\w+-BLKH-|cloudhash\.|cld-", hit.type)
            ]
            regular_hits = [hit for hit in all_hits if hit not in blacklist]
            yield (
                regular_hits,
                blacklist,
                cause,
                initiator,
                post_action,
                scan_id,
                standard_only,
            )

    async def _start_hook(self, cleanup_id, started, hits):
        dump = [hit.as_dict() for hit in hits]
        cleanup_started = HookEvent.MalwareCleanupStarted(
            cleanup_id=cleanup_id,
            started=started,
            total_files=len(hits),
            DUMP=dump,
        )
        await self._sink.process_message(cleanup_started)

    async def _clean_files(
        self,
        hits,
        blacklist=None,
        cause=None,
        initiator=None,
        post_action=None,
        scan_id=None,
        standard_only=None,
    ):
        user_hits = self._user_hits(hits)
        user_hits_black = self._user_hits(blacklist or [])

        for user in sorted({*user_hits, *user_hits_black}):
            hits_regular = user_hits.get(user, [])
            hits_black = user_hits_black.get(user, [])

            # Filter out ignored paths before invoking cleaner (DEF-36692)
            filtered_hits_regular = []
            for hit in hits_regular:
                try:
                    if await MalwareIgnorePath.is_path_ignored(hit.orig_file):
                        # Transition to a terminal state instead of silently skipping
                        MalwareHit.set_status([hit], MalwareHitStatus.FOUND)
                    else:
                        filtered_hits_regular.append(hit)
                except Exception as exc:
                    # be conservative: if ignore check fails, do not drop the hit
                    logger.exception(
                        "Ignore check failed for file %s: %s; keeping hit",
                        hit.orig_file,
                        exc,
                    )
                    filtered_hits_regular.append(hit)

            filtered_hits_black = []
            for hit in hits_black:
                try:
                    if await MalwareIgnorePath.is_path_ignored(hit.orig_file):
                        MalwareHit.set_status([hit], MalwareHitStatus.FOUND)
                    else:
                        filtered_hits_black.append(hit)
                except Exception as exc:
                    logger.exception(
                        "Ignore check failed for file %s: %s; keeping hit",
                        hit.orig_file,
                        exc,
                    )
                    filtered_hits_black.append(hit)

            hits_regular = filtered_hits_regular
            hits_black = filtered_hits_black
            user_hits_all = hits_regular + hits_black

            # DEF-40165: skip files that no longer exist on disk so they
            # are not dispatched to procu2 (guaranteed to fail and loop).
            # exists_async() checks concurrently in k8s without blocking
            # the event loop.
            all_hits = hits_regular + hits_black
            existence = await gather_exists(
                [h.orig_file_path for h in all_hits]
            )
            preflight_not_exist = [
                h for h, exists in zip(all_hits, existence) if not exists
            ]
            if preflight_not_exist:
                preflight_set = set(preflight_not_exist)
                hits_regular = [
                    h for h in hits_regular if h not in preflight_set
                ]
                hits_black = [h for h in hits_black if h not in preflight_set]
                MalwareHit.delete_instances(preflight_not_exist)
                await MalwareAction.not_exist(
                    preflight_not_exist,
                    cause=cause,
                    initiator=initiator,
                )
                user_hits_all = [
                    h for h in user_hits_all if h not in preflight_set
                ]

            # Skip if nothing to clean after filtering
            if not hits_regular and not hits_black:
                continue

            # the cleaner wrapper needs bare in-container paths; the
            # tenant rides in --username (k8s prefix stripped)
            files = [str(hit.orig_file_path) for hit in hits_regular]
            black = [str(hit.orig_file_path) for hit in hits_black]

            logger.debug("Cleaning files: %s", files + black)
            cleanup_id = uuid.uuid4().hex
            started = time.time()
            if is_uid(user):  # non panel user
                uid = int(user)

                async def _refuse(error, refused_hits):
                    await self._sink.process_message(
                        MessageType.MalwareCleanup(
                            hits=refused_hits,
                            result={},
                            cleanup_id=cleanup_id,
                            started=started,
                            error=error,
                            cause=cause,
                            initiator=initiator,
                            post_action=post_action,
                            scan_id=scan_id,
                            args=[],
                        )
                    )

                if _system_owned(uid):
                    # crontabs stay cleanable: the cleaner exempts
                    # /var/spool/cron from its own guard by path
                    crontab_ids = {
                        hit.id
                        for hit in user_hits_all
                        if is_crontab(hit.orig_file_path)
                    }
                    refused = [
                        hit
                        for hit in user_hits_all
                        if hit.id not in crontab_ids
                    ]
                    if refused:
                        logger.warning(
                            "Refusing to clean files owned by system"
                            " account %s, manual cleanup required",
                            uid,
                        )
                        await _refuse(
                            "Cleanup skipped. File is owned by a system"
                            " account, manual cleanup required",
                            refused,
                        )
                    if not crontab_ids:
                        continue
                    hits_regular = [
                        h for h in hits_regular if h.id in crontab_ids
                    ]
                    hits_black = [h for h in hits_black if h.id in crontab_ids]
                    user_hits_all = hits_regular + hits_black
                    files = [hit.orig_file for hit in hits_regular]
                    black = [hit.orig_file for hit in hits_black]
                    if not files and not black:
                        continue
                    logger.warning(
                        "Cleaning system-owned crontabs as a deliberate"
                        " guard exception: %s",
                        files + black,
                    )
                if not LicenseCLN.is_unlimited():
                    logger.error(
                        f"Can't clean files for non panel user {uid=}, "
                        "since license is limited"
                    )
                    await _refuse(
                        "Cleanup failed. Automatic cleanup of files not"
                        " owned by a panel user requires an unlimited"
                        " license",
                        user_hits_all,
                    )
                    continue
                if not (username := await get_username_by_uid(uid)):
                    logger.warning(f"Can't find username for {uid=}")
                    await _refuse(
                        "Cleanup failed. File owner cannot be resolved",
                        user_hits_all,
                    )
                    continue
                user = username
            await self._start_hook(cleanup_id, started, user_hits_all)
            result, error, cmd = await self._cleaner.start(
                user,
                files,
                soft=Config.CLEANUP_TRIM,
                blacklist=black,
                standard_only=standard_only,
            )
            await self._sink.process_message(
                MessageType.MalwareCleanup(
                    hits=user_hits_all,
                    result=result,
                    cleanup_id=cleanup_id,
                    started=started,
                    error=error,
                    cause=cause,
                    initiator=initiator,
                    post_action=post_action,
                    scan_id=scan_id,
                    args=cmd,
                )
            )

    async def _cleanup(self):
        if self._running:
            return
        if not self._proxy.hits:
            self._proxy.reset()
            return

        self._running = True

        with inactivity.track.task("cleanup"):
            try:
                data = self._cloud_assisted_hits()
                for (
                    all_hits,
                    blacklist,
                    cause,
                    initiator,
                    post_action,
                    scan_id,
                    standard_only,
                ) in data:
                    await self._clean_files(
                        all_hits,
                        blacklist=blacklist,
                        cause=cause,
                        initiator=initiator,
                        post_action=post_action,
                        scan_id=scan_id,
                        standard_only=standard_only,
                    )
            finally:
                self._running = False

    @recurring_check(1)
    async def cleanup(self):
        await self._cleanup()


class ResultProcessor(MessageSink, MessageSource):
    SCOPE = Scope.AV

    async def create_sink(self, loop):
        pass

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

    @staticmethod
    def _set_hit_status(hits: List[MalwareHit], status: str, cleaned_at=None):
        MalwareHit.set_status(hits, status, cleaned_at)
        for hit in hits:
            hit.status = status
            hit.cleaned_at = cleaned_at

    @expect(MessageType.MalwareCleanup)
    async def store_result(self, message):
        hits: List[MalwareHit] = message["hits"]
        result: CleanupResult = message["result"]
        cause = message.get("cause")
        initiator = message.get("initiator")
        now = time.time()

        processed = [hit for hit in hits if hit in result]
        unprocessed = [hit for hit in hits if hit not in result]

        error = message.get("error")
        if unprocessed:
            sample = [h.orig_file for h in unprocessed[:3]]
            logger.warning(
                "Cleanup: %d/%d hits unmatched in result"
                " (error=%r, sample=%s)",
                len(unprocessed),
                len(hits),
                error,
                sample,
            )
        if error and result:
            logger.warning(
                "Cleanup: partial failure — error=%r"
                " but %d result entries present",
                error,
                len(result),
            )

        not_exist = []
        async for hit in nice_iterator(processed, chunk_size=100):
            # in case if procu2.php tries to clean user file in root dirs,
            # it will be marked as non-existent due to 'Permission denied'
            # error which confuses users, so consider it as unable to cleanup.
            if result[hit].not_exist():  # pragma: no cover
                if await safe_exists(hit.orig_file_path):
                    unprocessed.append(hit)
                else:
                    not_exist.append(hit)

        # DEF-40165: files missing from procu2 result that no longer exist
        # on disk must be removed instead of being restored to FOUND,
        # otherwise they get re-dispatched for cleanup indefinitely.
        still_unprocessed = []
        unprocessed_existence = await gather_exists(
            [h.orig_file_path for h in unprocessed]
        )
        for hit, exists in zip(unprocessed, unprocessed_existence):
            if not exists:
                not_exist.append(hit)
            else:
                still_unprocessed.append(hit)
        unprocessed = still_unprocessed

        await MalwareAction.cleanup_unable(
            unprocessed, cause=cause, initiator=initiator
        )

        requires_myimunify_protection = [
            hit
            for hit in processed
            if result[hit].requires_myimunify_protection()
        ]
        await MalwareAction.cleanup_requires_myimunify_protection(
            requires_myimunify_protection, cause=cause, initiator=initiator
        )
        self._set_hit_status(
            requires_myimunify_protection,
            MalwareHitStatus.CLEANUP_REQUIRES_MYIMUNIFY_PROTECTION,
            now,
        )

        failed = [hit for hit in processed if result[hit].is_failed()]
        await MalwareAction.cleanup_failed(
            failed, cause=cause, initiator=initiator
        )

        cleaned = [hit for hit in processed if result[hit].is_cleaned()]
        await MalwareAction.cleanup_done(
            cleaned, cause=cause, initiator=initiator
        )
        self._set_hit_status(cleaned, MalwareHitStatus.CLEANUP_DONE, now)

        removed = [hit for hit in processed if result[hit].is_removed()]
        await MalwareAction.cleanup_removed(
            removed, cause=cause, initiator=initiator
        )
        self._set_hit_status(removed, MalwareHitStatus.CLEANUP_REMOVED, now)

        # a result row matching no outcome predicate (e.g. an error code
        # this agent does not know yet) must still reach a terminal state,
        # otherwise the hit stays in cleanup_started until agent restart
        unclassified = [
            hit
            for hit in processed
            if not (
                result[hit].not_exist()
                or result[hit].requires_myimunify_protection()
                or result[hit].is_failed()
                or result[hit].is_cleaned()
                or result[hit].is_removed()
            )
        ]
        if unclassified:
            logger.warning(
                "Cleanup: %d/%d hits with unclassified result entries"
                " (sample=%s)",
                len(unclassified),
                len(hits),
                [(h.orig_file, result[h]["e"]) for h in unclassified[:3]],
            )
            await MalwareAction.cleanup_unable(
                unclassified, cause=cause, initiator=initiator
            )

        MalwareHit.delete_instances(not_exist)
        await MalwareAction.not_exist(
            not_exist, cause=cause, initiator=initiator
        )

        for status, hit_list in _group_by_status(
            unprocessed, failed, unclassified
        ).items():
            self._set_hit_status(hit_list, status)

        await self.send_failed_to_cleanup_hits_to_mrs(failed)

        return message

    async def send_failed_to_cleanup_hits_to_mrs(self, failed_to_cleanup_hits):
        if failed_to_cleanup_hits:
            await self._sink.process_message(
                MessageType.MalwareMRSUpload(
                    hits=[
                        malware_response.HitInfo(hit.orig_file, hit.hash)
                        for hit in failed_to_cleanup_hits
                    ],
                    upload_reason="cleanup_failure_current",
                )
            )
            await self._sink.process_message(
                MessageType.MalwareMRSUpload(
                    hits=[
                        malware_response.HitInfo(
                            str(CleanupStorage.get_hit_store_path(hit)),
                            hit.hash,
                        )
                        for hit in failed_to_cleanup_hits
                    ],
                    upload_reason="cleanup_failure_original",
                )
            )


class StorageController(MessageSink):
    """Remove old backed up files from storage"""

    def __init__(self):
        self._clear_task = None
        self._keep = Config.CLEANUP_KEEP

    async def create_sink(self, loop):
        self._clear_task = loop.create_task(self.daily_clear())

    async def shutdown(self):
        if self._clear_task:
            await safe_cancel_task(self._clear_task)

    async def _clear(self):
        now = time.time()
        keep_hits = now - self._keep * DAY
        keep_orig = now - (self._keep + 1) * DAY  # keep files one more day
        MalwareHit.delete().where(MalwareHit.cleaned_at < keep_hits).execute()
        cleared = await CleanupStorage.clear(keep_orig)
        if cleared:
            logger.info(
                "Cleanup storage have cleaned. Files removed: %s", cleared
            )

    @expect(MessageType.ConfigUpdate)
    @utils.log_error_and_ignore()
    async def config_updated(self, _):
        if self._keep != Config.CLEANUP_KEEP:
            self._keep = Config.CLEANUP_KEEP
            await self._clear()

    @recurring_check(
        check_lock,
        check_period_first=True,
        check_lock_period=DAY,
        lock_file=LOCK_FILE,
    )
    async def daily_clear(self):
        await self._clear()


def decide_if_standard_signatures_only(user, standard_only):
    """Root user or user with MyImunify can use advanced signatures"""

    if not MyImunifyConfig.ENABLED:
        return False

    if user is None or user == "root" or myimunify_protection_enabled(user):
        return standard_only

    return True


class ResultProcessorIm360(ResultProcessor):
    """Imunify360 specialization of ResultProcessor, which removes all
    cleaned and removed files from HackerTrap
    """

    SCOPE = Scope.IM360

    @expect(MessageType.MalwareCleanup)
    async def store_result(self, message):
        message = await super().store_result(message)
        to_remove = [
            Path(hit)
            for hit, state in message["result"].items()
            if (state.is_cleaned() or state.is_removed())
        ]
        await HackerTrapHitsSaver.update_sa_hits([], to_remove)


class CleanupDb(MessageSink):
    SCOPE = Scope.IM360

    def __init__(self):
        self._loop = None

    @staticmethod
    async def _start_cleaner(path, app_name):
        cleanup_id = uuid.uuid4().hex
        await MalwareDatabaseCleaner(cleanup_id, path, app_name).start()

    async def _cleanup_next(self):
        if (
            MalwareHit.db_hits_under_cleanup().exists()
            or (
                next_hit := MalwareHit.db_hits_pending_cleanup()
                .order_by(MalwareHit.timestamp.asc())
                .first()
            )
            is None
        ):
            return
        # One --clean pass cleans the whole DB, so batch all its pending hits.
        same_db_hits = list(
            MalwareHit.db_hits_pending_cleanup().where(
                (MalwareHit.orig_file == next_hit.orig_file)
                & (MalwareHit.app_name == next_hit.app_name)
            )
        )
        logger.info(
            "Cleaning hit: (%s::%s::%s)",
            next_hit.orig_file,
            next_hit.app_name,
            next_hit.type,
        )
        MalwareHit.set_status(same_db_hits, MalwareHitStatus.CLEANUP_STARTED)
        await self._start_cleaner(next_hit.orig_file, next_hit.app_name)

    async def create_sink(self, loop):
        self._loop = loop
        await self._cleanup_next()

    @expect(MessageType.MalwareCleanupTask)
    async def process_cleanup_task(self, message):
        hits = MalwareHit.refresh_hits(message["hits"])
        hits_to_clean = filter_cleanable(hits)
        db_hits = [
            hit
            for hit in hits_to_clean
            if hit.resource_type == MalwareScanResourceType.DB.value
        ]
        if not db_hits:
            return

        MalwareHit.set_status(db_hits, MalwareHitStatus.CLEANUP_PENDING)
        await self._cleanup_next()

    @expect(MessageType.MalwareCleanComplete)
    async def parse_cleanup_results(self, message):
        clean_id = message["scan_id"]
        detached_cleanup = MDSDetachedCleanup(clean_id)
        try:
            cleanup_outcome = await detached_cleanup.complete()
        except ScanAlreadyCompleteError:
            # This happens when AV is woken up by AiBolit. See DEF-11078.
            logger.warning(
                "Cannot complete cleanup %s, assuming it is already complete",
                clean_id,
            )
            return
        finally:
            shutil.rmtree(
                str(detached_cleanup.detached_dir), ignore_errors=True
            )
        await g.sink.process_message(cleanup_outcome)

    @expect(MessageType.MalwareDatabaseCleanup)
    async def update_cleaned_hits_status(
        self, message: MessageType.MalwareDatabaseCleanup
    ):
        cleaned_hits = MalwareHit.db_hits_under_cleanup_in(message.succeeded)
        failed_hits = MalwareHit.db_hits_under_cleanup_in(message.failed)
        MalwareHit.set_status(
            cleaned_hits, MalwareHitStatus.CLEANUP_DONE, time.time()
        )
        MalwareHit.set_status(failed_hits, MalwareHitStatus.FOUND)
        # The queue is serial, so any hit still under cleanup belongs to
        # this operation; without a terminal status it would block the
        # queue forever. The sweep is deliberately not scoped to this
        # operation's (path, app): a run that fails before reaching any
        # database reports no hits at all, and only an unscoped sweep
        # finalizes its hits then. The trade-off is that hits stranded in
        # cleanup_started by an unrelated earlier crash are finalized (and
        # get a failed_to_cleanup event) on this run's behalf.
        if leftover_hits := list(MalwareHit.db_hits_under_cleanup()):
            MalwareHit.set_status(leftover_hits, MalwareHitStatus.FOUND)
            await MalwareAction.cleanup_failed(
                leftover_hits, cause=None, initiator=None
            )
        await self._cleanup_next()

    @expect(MessageType.MalwareDatabaseCleanupFailed)
    async def update_failed_hits_status(self, message):
        """
        Clear the queue when the cleanup fails,
        set hits' status back to infected
        """
        # We assume here that all CLEANUP_STARTED hits are part of the
        # same cleanup operation
        hits = list(MalwareHit.db_hits_under_cleanup())
        MalwareHit.set_status(hits, MalwareHitStatus.FOUND)
        await MalwareAction.cleanup_failed(hits, cause=None, initiator=None)
        await self._cleanup_next()

    @expect(MessageType.MalwareDatabaseCleanup)
    async def save_cleanup_events_in_history(
        self, message: MessageType.MalwareDatabaseCleanup
    ):
        cause = None
        initiator = None
        cleaned_hits = MalwareHit.get_db_hits(message.succeeded)
        await MalwareAction.cleanup_done(
            cleaned_hits, cause=cause, initiator=initiator
        )
        failed_hits = MalwareHit.get_db_hits(message.failed)
        await MalwareAction.cleanup_failed(
            failed_hits, cause=cause, initiator=initiator
        )


class RestoreOriginalDb(MessageSink):
    SCOPE = Scope.IM360

    def __init__(self):
        self.loop = None

    @staticmethod
    async def _restore_next():
        if (
            MalwareHit.db_hits_under_restoration().exists()
            or (
                hit_to_restore := (
                    MalwareHit.db_hits_pending_cleanup_restore()
                    .order_by(MalwareHit.timestamp.asc())
                    .first()
                )
            )
            is None
        ):
            return

        signature_id = (
            hit_to_restore.signature_id
            if hit_to_restore.status
            == MalwareHitStatus.CLEANUP_REMOTE_RESTORE_PENDING
            else None
        )

        logger.info(
            "Restoring from cleanup hit: (%s::%s::%s)",
            hit_to_restore.orig_file,
            hit_to_restore.app_name,
            signature_id,
        )

        await MalwareDatabaseRestore(
            path=hit_to_restore.orig_file,
            app_name=hit_to_restore.app_name,
            signature_id=signature_id,
        ).restore()
        MalwareHit.set_status(
            [hit_to_restore], MalwareHitStatus.CLEANUP_RESTORE_STARTED
        )

    async def create_sink(self, loop):
        self.loop = loop
        await self._restore_next()

    @staticmethod
    def _filter_under_restore(
        hits: Iterable[MalwareHit],
    ) -> Iterable[MalwareHit]:
        return (
            hit
            for hit in hits
            if hit.status == MalwareHitStatus.CLEANUP_RESTORE_STARTED
        )

    @expect(MessageType.MalwareDatabaseRestoreTask)
    async def queue_db_restore(self, message: MalwareDatabaseRestoreTask):
        query = (
            MalwareHit.db_hits()
            .where(MalwareHit.orig_file == message.path)
            .where(MalwareHit.app_name == message.app_name)
        )
        restore_status = MalwareHitStatus.CLEANUP_RESTORE_PENDING
        if message.signature_id:
            query = query.where(MalwareHit.type == message.signature_id)
            restore_status = MalwareHitStatus.CLEANUP_REMOTE_RESTORE_PENDING

        MalwareHit.set_status(
            query,
            restore_status,
        )
        await self._restore_next()

    @expect(MessageType.MalwareRestoreComplete)
    async def parse_restore_results(self, message):
        restore_id = message["scan_id"]
        detached_restore = MDSDetachedRestore(restore_id)

        try:
            restore_message = await detached_restore.complete()
        except ScanAlreadyCompleteError:
            # This happens when AV is woken up by AiBolit. See DEF-11078.
            logger.warning(
                "Cannot complete restore %s, assuming it is already complete",
                restore_id,
            )
            return
        finally:
            shutil.rmtree(
                str(detached_restore.detached_dir), ignore_errors=True
            )

        await g.sink.process_message(restore_message)

    @expect(MessageType.MalwareDatabaseRestore)
    async def update_restored_hits_status(self, message):
        restored_hits = MalwareHit.get_db_hits(message.succeeded)
        MalwareHit.set_status(
            self._filter_under_restore(restored_hits), MalwareHitStatus.FOUND
        )
        # Failed hits leave the under-restoration status here so the
        # leftover sweep does not report them a second time; their
        # history event comes from save_restore_events_in_history.
        failed_hits = MalwareHit.get_db_hits(message.failed)
        MalwareHit.set_status(
            self._filter_under_restore(failed_hits),
            MalwareHitStatus.CLEANUP_DONE,
        )
        # The queue is serial, so any hit still under restoration belongs
        # to this operation; without a terminal status it would block the
        # queue forever. Unscoped for the same reason as the cleanup
        # sweep: a run that fails before reaching any database reports no
        # hits at all.
        if leftover_hits := list(MalwareHit.db_hits_under_restoration()):
            MalwareHit.set_status(leftover_hits, MalwareHitStatus.CLEANUP_DONE)
            for hit in leftover_hits:
                await MalwareAction.cleanup_failed_restore(
                    path=hit.orig_file,
                    app_name=hit.app_name,
                    resource_type=MalwareScanResourceType.DB.value,
                    file_owner=hit.owner,
                    file_user=hit.user,
                    initiator=message.get("initiator"),
                    cause=message.get("cause"),
                    db_host=hit.db_host,
                    db_port=hit.db_port,
                    db_name=hit.db_name,
                    signature_id=hit.signature_id,
                )
        await self._restore_next()

    @expect(MessageType.MalwareDatabaseRestore)
    async def save_restore_events_in_history(self, message):
        cause = message.get("cause")
        initiator = message.get("initiator")
        restored_hits = MalwareHit.get_db_hits(message.succeeded)
        for hit in restored_hits:
            # FIXME: change cleanup_restored_original to accept multiple
            # values
            await MalwareAction.cleanup_restored_original(
                path=hit.orig_file,
                app_name=hit.app_name,
                resource_type=MalwareScanResourceType.DB.value,
                file_owner=hit.owner,
                file_user=hit.user,
                initiator=initiator,
                cause=cause,
                db_host=hit.db_host,
                db_port=hit.db_port,
                db_name=hit.db_name,
                signature_id=hit.signature_id,
            )
        failed_hits = MalwareHit.get_db_hits(message.failed)
        for hit in failed_hits:
            await MalwareAction.cleanup_failed_restore(
                path=hit.orig_file,
                app_name=hit.app_name,
                resource_type=MalwareScanResourceType.DB.value,
                file_owner=hit.owner,
                file_user=hit.user,
                initiator=initiator,
                cause=cause,
                db_host=hit.db_host,
                db_port=hit.db_port,
                db_name=hit.db_name,
                signature_id=hit.signature_id,
            )

    @expect(MessageType.MalwareDatabaseRestoreFailed)
    async def update_failed_hits_status(self, message):
        """
        Clear the queue when the restore fails,
        set hits' status back to cleanup_done
        """
        hits = list(MalwareHit.db_hits_under_restoration())
        MalwareHit.set_status(hits, MalwareHitStatus.CLEANUP_DONE)
        for hit in hits:
            await MalwareAction.cleanup_failed_restore(
                path=hit.orig_file,
                app_name=hit.app_name,
                resource_type=MalwareScanResourceType.DB.value,
                file_owner=hit.owner,
                file_user=hit.user,
                initiator=message.get("initiator"),
                cause=message.get("cause"),
                db_host=hit.db_host,
                db_port=hit.db_port,
                db_name=hit.db_name,
                signature_id=hit.signature_id,
            )
        await self._restore_next()
