[EDIT] FILE: messages.py
import asyncio import json import os from enum import Enum from typing import List from defence360agent.contracts.config import Core as CoreConfig class MessageNotFoundError(Exception): pass class UnknownMessage: """ Used as stub for MessageType """ def __init__(self): raise MessageNotFoundError("Message class is not found.") def __getattr__(self, name): return "Unknown" # pragma: no cover class MessageT: _subclasses = [] def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) cls._subclasses.append(cls) @classmethod def get_subclasses(cls): return tuple(cls._subclasses) class _MessageType: """ Used to get specific message class. For example, >>> _MessageType().ConfigUpdate <class 'defence360agent.contracts.messages.ConfigUpdate'> >>> _MessageType().NotExistMessage <class 'defence360agent.contracts.messages.UnknownMessage'> >>> """ def __getattr__(self, name): for subcls in Message.get_subclasses(): # is is supposed that all subclasses have different names if subcls.__name__ == name: return subcls return UnknownMessage MessageType = _MessageType() class ReportTarget(Enum): API = "api" PERSISTENT_CONNECTION = "conn" class Reportable(MessageT): """ Mixin class for messages that should be sent to the server """ TARGET = ReportTarget.PERSISTENT_CONNECTION @classmethod def get_subclass_with_method(cls, method: str): """ Return a subclass with the same DEFAULT_METHOD as *method*. It can be used to detect report target from message method. NOTE: it is not guaranteed that the class with the *method* is unique, in this case the first subclass found is returned, but it is tested that all such subclasses have the same TARGET. """ for subclass in cls.__subclasses__(): if method == getattr(subclass, "DEFAULT_METHOD"): return subclass return None # pragma: no cover class Received(MessageT): """ Mixin class for messages received from the server. These messages are created in the client360 plugin when receiving a request from imunify360.cloudlinux.com. """ @classmethod def get_subclass_with_action(cls, action: str): for subclass in cls.__subclasses__(): received_actions = getattr(subclass, "RECEIVED_ACTIONS", []) or [ getattr(subclass, "DEFAULT_METHOD") ] if action in received_actions: return subclass raise MessageNotFoundError( 'Message class is not found for "{}" action'.format(action) ) class Lockable(MessageT): _lock = None @classmethod async def acquire(cls) -> None: if cls._lock is None: cls._lock = asyncio.Lock() await cls._lock.acquire() @classmethod def locked(cls) -> bool: return cls._lock is not None and cls._lock.locked() @classmethod def release(cls) -> None: if cls._lock is not None: cls._lock.release() class Message(dict, MessageT): """ Base class for messages to be passed as a parameter to plugins.MessageSink.process_message() """ # Default method='...' to send to the Server DEFAULT_METHOD = "" PRIORITY = 10 PROCESSING_TIME_THRESHOLD = 60 # 1 min #: fold collections' repr with more than the threshold number of items _FOLD_LIST_THRESHOLD = 100 #: shorten strings longer than the threshold characters _SHORTEN_STR_THRESHOLD = 320 def __init__(self, *args, **kwargs) -> None: if self.DEFAULT_METHOD: self["method"] = self.DEFAULT_METHOD super(Message, self).__init__(*args, **kwargs) @property def payload(self): return {k: v for k, v in self.items() if k != "method"} def __getattr__(self, name): """ Called when an attribute lookup has not found the attribute in the usual places A shortcut to access an item from dict """ try: return self[name] except KeyError as exc: raise AttributeError(name) from exc def __repr__(self): """Render for logs: collections with more than _FOLD_LIST_THRESHOLD items are collapsed to a count and strings longer than _SHORTEN_STR_THRESHOLD are shortened, recursively through nested payloads, so a single message cannot flood the log.""" folded_msg = { k: _fold_repr_value( v, fold_limit=self._FOLD_LIST_THRESHOLD, str_limit=self._SHORTEN_STR_THRESHOLD, ) for k, v in self.items() } return "{}({})".format(self.__class__.__qualname__, folded_msg) def __str__(self): return self.__repr__() class MessageList(Message): def __init__(self, msg_list): super().__init__(list=msg_list) @property def payload(self): return self.list class ShortenReprListMixin: """ Do not flood console.log with large sequences The method collapses messages that are a list. Instead of showing all the elements of the message, their number will be displayed. """ def __repr__(self: dict): # type: ignore return "{}({})".format( self.__class__.__qualname__, "<{} item(s)>".format(len(self.get("items", []))), ) class Accumulatable(Message): """Messages of this class will be grouped into a list of LIST_CLASS message instance by Accumulate plugin. Messages whose do_accumulate() call returns False will not be added to list.""" LIST_CLASS = MessageList def do_accumulate(self) -> bool: """Return True if this message is worth collecting, False otherwise.""" return True class ServerConnected(Message): pass # alias (for better client code readability) class ServerReconnected(ServerConnected): pass class Ping(Message, Reportable): """ Will send this message on connected, reconnected events to provide central server with agent version """ DEFAULT_METHOD = "PING" PRIORITY = 0 def __init__(self): super().__init__() self["version"] = CoreConfig.VERSION class Ack(Message, Reportable): """ Notify Server that a persistent message with *seq_number* has been received by Agent. """ DEFAULT_METHOD = "ACK" def __init__(self, seq_number, **kwargs): super().__init__(**kwargs) self["_meta"] = dict(per_seq=seq_number) class Noop(Message): """ Sending NOOP to the agent to track the message in agent logs. """ DEFAULT_METHOD = "NOOP" class ServerConfig(Message, Reportable): """ Information about server environment """ DEFAULT_METHOD = "SERVER_CONFIG" TARGET = ReportTarget.API def __repr__(self): return "{}()".format(self.__class__.__qualname__) class WpSecurityPluginStats(Message, Reportable): DEFAULT_METHOD = "WP_SECURITY_PLUGIN_STATS" TARGET = ReportTarget.API class DomainList(Message, Reportable): """ Information about server domains """ DEFAULT_METHOD = "DOMAIN_LIST" TARGET = ReportTarget.API def __repr__(self): return "{}()".format(self.__class__.__qualname__) class FilesUpdated(Message): """ To consume products of files.update() """ def __init__(self, files_type, files_index): """ :param files_type: files.Type :param files_index: files.LocalIndex """ # explicit is better than implicit self["files_type"] = files_type self["files_index"] = files_index def __repr__(self): """ Do not flood console.log with large sequences """ return "{}({{'files_type':'{}', 'files_index':{}}})".format( self.__class__.__qualname__, self["files_type"], self["files_index"], ) class UpdateFiles(Message, Received): """ Update files by getting message from the server """ DEFAULT_METHOD = "UPDATE" class ConfigUpdate(Message): DEFAULT_METHOD = "CONFIG_UPDATE" class Reject(Exception): """ Kinda message filtering facility. Raised in order to stop message processing through plugins. Takes reason of reject as argument. """ pass class Health(Message): DEFAULT_METHOD = "HEALTH" class CommandInvoke(Message, Reportable): DEFAULT_METHOD = "COMMAND_INVOKE" class ScanFailed(Message, Reportable): DEFAULT_METHOD = "SCAN_FAILED" class CleanupFailed(Message, Reportable): DEFAULT_METHOD = "CLEANUP_FAILED" class RestoreFromBackupTask(Message): """ Creates a task to restore files from backup """ DEFAULT_METHOD = "MALWARE_RESTORE_FROM_BACKUP" class cPanelEvent(Message): DEFAULT_METHOD = "PANEL_EVENT" ALLOWED_FIELDS = { "new_pkg", "plan", "exclude", "imunify360_proactive", "imunify360_av", } @classmethod def from_hook_event( cls, username: str, hook: str, ts: float, fields: dict ): data = { k.lower(): v for k, v in fields.items() if k.lower() in cls.ALLOWED_FIELDS } # Check for user rename if ( hook == "Modify" and "user" in fields and "newuser" in fields and fields["user"] != fields["newuser"] ): data["old_username"] = fields["user"] return cls( { "username": username, "hook": hook, "data": data, "timestamp": ts, } ) class IContactSent(Message, Reportable): DEFAULT_METHOD = "ICONTACT_SENT" def _shorten_str(s: str, limit: int) -> str: """Shorten *s* string if its length exceeds *limit*.""" assert limit > 4 return ( f"{s[: limit // 2 - 1]}...{s[-limit // 2 + 2 :]}" if len(s) > limit else s ) def _fold_repr_value(value, *, fold_limit: int, str_limit: int): if isinstance(value, str): return _shorten_str(value, str_limit) if isinstance(value, dict): if len(value) > fold_limit: return "<{} item(s)>".format(len(value)) return { k: _fold_repr_value(v, fold_limit=fold_limit, str_limit=str_limit) for k, v in value.items() } if isinstance(value, (list, tuple, set, frozenset)): if len(value) > fold_limit: return "<{} item(s)>".format(len(value)) return type(value)( _fold_repr_value(v, fold_limit=fold_limit, str_limit=str_limit) for v in value ) return value class BackupInfo(Message, Reportable): """Information about enabled backup backend""" DEFAULT_METHOD = "BACKUP_INFO" class MDSReportList(ShortenReprListMixin, Message, Reportable): DEFAULT_METHOD = "MDS_SCAN_LIST" class MDSReport(Accumulatable): LIST_CLASS = MDSReportList # Target serialized size per outgoing message chunk. Kept far below the # 10 MB NATS max_payload so envelope overhead and size-estimate drift cannot # push a chunk over the transport limit; the transport keeps a split-on- # overflow safety net for the rare cases this estimate misses. MAX_MESSAGE_SIZE = int( os.environ.get("IMUNIFY360_MAX_MESSAGE_SIZE", 1024 * 1024) ) def serialized_size(obj) -> int: from defence360agent.utils.json import ServerJSONEncoder try: return len(json.dumps(obj, cls=ServerJSONEncoder).encode()) except (TypeError, ValueError): return len(repr(obj).encode()) def estimate_size(obj) -> int: """Upper bound on obj's JSON byte size as sent on the wire (ensure_ascii), biased to never undercount. Far cheaper than a full ``serialized_size`` per call on big scans: JSON-native values are measured structurally without building the encoded string, and printable-ASCII strings (the common path for file paths/snippets) are counted with C-level ``str`` ops. Non-native values (peewee Models, IPs, ...) fall back to the exact ``serialized_size`` — their ``repr`` would wildly undercount the ServerJSONEncoder output. The transport keeps a split-on-overflow net for the rare drift this leaves.""" if obj is None: return 4 if isinstance(obj, bool): return 5 if isinstance(obj, int): return max(20, len(str(obj)) + 1) if isinstance(obj, float): return 24 if isinstance(obj, str): if obj.isascii() and obj.isprintable(): return len(obj) + 2 + obj.count('"') + obj.count("\\") return len(json.dumps(obj)) if isinstance(obj, (list, tuple)): return 2 + sum(estimate_size(v) + 1 for v in obj) if isinstance(obj, dict): return 2 + sum( estimate_size(k if isinstance(k, str) else str(k)) + 1 + estimate_size(v) + 1 for k, v in obj.items() ) return serialized_size(obj) class Splittable: """ A message list could be split into multiple batches. The split is possible for a list itself along with internal resources. """ LIST_SIZE = None BATCH_SIZE = None BATCH_FIELD = None @classmethod def _max_message_size(cls) -> int: return MAX_MESSAGE_SIZE @classmethod def _split_items(cls, messages: List[Accumulatable]): """ Split messages' internal lists of things into batches. A field that is meant to split is defined by `BATCH_FIELD`. """ if cls.BATCH_FIELD and cls.BATCH_SIZE: for message in messages: if (items := message.get(cls.BATCH_FIELD)) is None: yield message else: message_class = type(message) for batch in cls._size_bounded_batches(items, message): data = message.copy() data[cls.BATCH_FIELD] = batch new_message = message_class(data) yield new_message else: yield from iter(messages) @classmethod def _unit_size(cls, unit, is_dict: bool, message) -> int: """Serialized byte cost of one BATCH_FIELD unit. Subclasses override to also count data paired with the unit in sibling fields of the message (e.g. a per-hit cleanup result), so those bytes are not excluded from the byte budget.""" return estimate_size({unit[0]: unit[1]} if is_dict else unit) @classmethod def _size_bounded_batches(cls, items, message): """Pack `items` into batches bounded by both the byte budget and the `BATCH_SIZE` count. A single element larger than the budget is emitted alone rather than dropped.""" budget = cls._max_message_size() is_dict = isinstance(items, dict) units = list(items.items()) if is_dict else items def build(buffer): return dict(buffer) if is_dict else list(buffer) buffer = [] size = 0 for unit in units: unit_size = cls._unit_size(unit, is_dict, message) if buffer and ( size + unit_size > budget or len(buffer) >= cls.BATCH_SIZE ): yield build(buffer) buffer, size = [], 0 buffer.append(unit) size += unit_size if buffer: yield build(buffer) @classmethod def batched(cls, messages: List[Accumulatable]): list_size = cls.LIST_SIZE or len(messages) budget = cls._max_message_size() buffer = [] size = 0 for message in cls._split_items(messages): message_size = estimate_size(message) if buffer and ( size + message_size > budget or len(buffer) >= list_size ): yield buffer buffer, size = [], 0 buffer.append(message) size += message_size if buffer: yield buffer class EnsureServiceState(Message): """Ensure the service has the appropriate status""" DEFAULT_METHOD = "ENSURE_SERVICE_STATE" class SensorWordpressIncidentList(MessageList, Reportable): """Aggregated incident list""" DEFAULT_METHOD = "INCIDENT_LIST" class WordpressPluginAction(Message): DEFAULT_METHOD = "WP_SECURITY_PLUGIN_ACTION" class WordpressPluginTelemetry(Message, Reportable): """ Information about telemetry event related to Imunify Security WordPress plugin """ DEFAULT_METHOD = "WP_SECURITY_PLUGIN_EVENT" TARGET = ReportTarget.API def __repr__(self): return "{}()".format(self.__class__.__qualname__) class WPRuleDisabled(Message, Reportable): """WordPress protection rule disabled.""" DEFAULT_METHOD = "RULE_DISABLED" class WPRuleEnabled(Message, Reportable): """WordPress protection rule re-enabled.""" DEFAULT_METHOD = "RULE_ENABLED" class GeneralMetrics(MessageList, Reportable): DEFAULT_METHOD = "GENERAL_METRICS"
SAVE
CANCEL