#!/usr/bin/python3
# log_digest - mail and XMPP digests of warnings and errors from the journal and rsyslog files
#
# Copyright (C) 2026 Thomas Wagner <wagner-thomas@gmx.at>
# SPDX-License-Identifier: GPL-2.0-or-later
"""Collect the warnings and errors of a time span from the systemd journal
and from the files rsyslog writes, merge messages that appear in both, and
send a digest by mail and XMPP (go-sendxmpp).

The files rsyslog writes, and which severities each of them holds, are read
from rsyslog's configuration.
"""

import argparse
import bz2
import collections
import datetime
import email.message
import email.utils
import glob
import gzip
import json
import lzma
import os
import re
import shlex
import shutil
import smtplib
import socket
import subprocess
import sys
import time

VERSION = "0.1"

LEVELS = ["emerg", "alert", "crit", "err", "warning", "notice", "info", "debug"]
LEVEL_ALIASES = {"emergency": 0, "panic": 0, "critical": 2, "error": 3, "warn": 4, "informational": 6}
ALL_LEVELS = frozenset(range(8))

DEFAULTS = {
    "since": "-1h",
    "until": "now",
    "priority": "warning",
    "max_lines": 100,
    "top": 10,
    "duplicate_window": 5,
    "send_empty": False,
    "subject": "[log_digest] {host}: {count} message(s) {span}",
    "include": [],
    "exclude": [],
    "journal": {"enabled": "auto", "journalctl": "journalctl", "args": ["--merge"]},
    "rsyslog": {"enabled": "auto", "config": "/etc/rsyslog.conf", "files": []},
    "mail": {"to": [], "from": None, "method": "smtp", "host": "localhost", "port": 25,
             "sendmail": "/usr/sbin/sendmail"},
    "xmpp": {"to": [], "high_priority_to": [], "chatroom": False, "config": None,
             "go_sendxmpp": "go-sendxmpp", "args": [], "full": False},
}


class LogDigestError(Exception):
    pass


# ---------------------------------------------------------------------------
# severities and time stamps, in the syntax of journalctl
# ---------------------------------------------------------------------------

def parse_level(text):
    value = str(text).strip().lower()
    if value.isdigit() and int(value) <= 7:
        return int(value)
    if value in LEVELS:
        return LEVELS.index(value)
    if value in LEVEL_ALIASES:
        return LEVEL_ALIASES[value]
    raise ValueError("unknown severity %r" % text)


def parse_priority(text):
    """The levels journalctl -p selects: LEVEL and everything more severe, or FROM..TO."""
    value = str(text)
    if ".." in value:
        first, last = (parse_level(part) for part in value.split("..", 1))
        return frozenset(range(min(first, last), max(first, last) + 1))
    return frozenset(range(0, parse_level(value) + 1))


def level_label(entry):
    if entry.priority is None:
        return "?"
    return LEVELS[entry.priority] + ("+" if entry.approximate else "")


TIME_UNITS = [
    (("usec", "us", "µs"), 1e-6), (("msec", "ms"), 1e-3),
    (("seconds", "second", "sec", "s"), 1), (("minutes", "minute", "min", "m"), 60),
    (("hours", "hour", "hr", "h"), 3600), (("days", "day", "d"), 86400),
    (("weeks", "week", "w"), 604800), (("months", "month", "M"), 2629800),
    (("years", "year", "y"), 31557600),
]
UNIT_SECONDS = {name: seconds for names, seconds in TIME_UNITS for name in names}


def parse_span(text):
    """A time span of systemd.time(7), e.g. "1h 30min" or "2days"."""
    total = 0.0
    found = False
    for number, unit in re.findall(r"(\d+(?:\.\d+)?)\s*([a-zA-Zµ]*)", text):
        unit = unit or "s"
        seconds = UNIT_SECONDS.get(unit, UNIT_SECONDS.get(unit.lower()))
        if seconds is None:
            raise ValueError("unknown time unit %r" % unit)
        total += float(number) * seconds
        found = True
    if not found or re.sub(r"[\d.\sa-zA-Zµ]", "", text):
        raise ValueError("invalid time span %r" % text)
    return total


def parse_time_fallback(text, now):
    """The common forms of systemd.time(7), for systems without systemd-analyze."""
    value = text.strip()
    today = datetime.datetime.fromtimestamp(now).replace(hour=0, minute=0, second=0, microsecond=0)
    words = {"now": now, "today": today.timestamp(),
             "yesterday": (today - datetime.timedelta(days=1)).timestamp(),
             "tomorrow": (today + datetime.timedelta(days=1)).timestamp()}
    if value in words:
        return words[value]
    if value.startswith("@"):
        return float(value[1:])
    if value[:1] in "+-":
        return now + (1 if value[0] == "+" else -1) * parse_span(value[1:])
    if value.endswith(" ago"):
        return now - parse_span(value[:-4])
    if value.endswith(" left"):
        return now + parse_span(value[:-5])
    # an optional weekday in front, as systemd prints it
    value = re.sub(r"^(Mon|Tue|Wed|Thu|Fri|Sat|Sun)[a-z]*\s+", "", value)
    for pattern in ("%Y-%m-%d %H:%M:%S", "%Y-%m-%d %H:%M", "%Y-%m-%d", "%Y-%m-%dT%H:%M:%S"):
        try:
            return datetime.datetime.strptime(value, pattern).timestamp()
        except ValueError:
            pass
    for pattern in ("%H:%M:%S", "%H:%M"):
        try:
            clock = datetime.datetime.strptime(value, pattern)
        except ValueError:
            continue
        return today.replace(hour=clock.hour, minute=clock.minute, second=clock.second).timestamp()
    raise ValueError("cannot parse time %r" % text)


def parse_time(text, now):
    """A time stamp in the syntax of journalctl --since/--until, as seconds since the epoch.

    systemd-analyze timestamp understands exactly that syntax; without it a
    reimplementation of its common forms is used.
    """
    text = str(text)
    if text == "now":
        return now
    exe = shutil.which("systemd-analyze")
    if exe:
        try:
            result = subprocess.run([exe, "timestamp", "--", text], stdout=subprocess.PIPE, stderr=subprocess.PIPE,
                                    universal_newlines=True, timeout=10)
        except (OSError, subprocess.TimeoutExpired):
            result = None
        if result is not None:
            match = re.search(r"UNIX seconds:\s*@(\d+(?:\.\d+)?)", result.stdout)
            if result.returncode == 0 and match:
                return float(match.group(1))
            if result.returncode != 0:
                raise ValueError("cannot parse time %r: %s" % (text, result.stderr.strip() or "invalid"))
    return parse_time_fallback(text, now)


def format_time(seconds):
    return time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(seconds))


# ---------------------------------------------------------------------------
# messages and filter rules
# ---------------------------------------------------------------------------

class Entry:
    __slots__ = ("time", "host", "program", "pid", "message", "priority", "approximate", "origin", "unit")

    def __init__(self, time, host, program, pid, message, priority, origin, approximate=False, unit=""):
        self.time = time
        self.host = host
        self.program = program
        self.pid = pid
        self.message = message
        self.priority = priority
        self.approximate = approximate
        self.origin = origin
        self.unit = unit


RULE_FIELDS = {
    "host": "host", "hostname": "host",
    "program": "program", "source": "program", "identifier": "program", "tag": "program",
    "message": "message", "content": "message", "msg": "message",
    "unit": "unit", "origin": "origin",
}


class Rule:
    """Matches a message if every condition's regular expression is found in its field."""

    def __init__(self, conditions, text):
        self.conditions = conditions
        self.text = text

    def matches(self, entry):
        return all(regex.search(getattr(entry, field) or "") for field, regex in self.conditions)


def parse_rule(value):
    """A rule from YAML (a mapping of fields to regular expressions) or the command line (FIELD=REGEX)."""
    if isinstance(value, str):
        field, sep, regex = value.partition("=")
        if not sep:
            raise ValueError("a rule is FIELD=REGEX, not %r" % value)
        value = {field.strip(): regex}
    if not isinstance(value, dict) or not value:
        raise ValueError("a rule needs at least one of %s" % ", ".join(sorted(set(RULE_FIELDS.values()))))
    conditions = []
    for field, regex in value.items():
        name = RULE_FIELDS.get(str(field).strip().lower())
        if name is None:
            raise ValueError("unknown field %r in rule, use %s" % (field, ", ".join(sorted(set(RULE_FIELDS.values())))))
        try:
            conditions.append((name, re.compile(str(regex))))
        except re.error as err:
            raise ValueError("invalid regular expression %r: %s" % (regex, err))
    return Rule(conditions, ", ".join("%s=%s" % (field, regex) for field, regex in value.items()))


# ---------------------------------------------------------------------------
# the journal
# ---------------------------------------------------------------------------

def field_text(value):
    """A journal field: a string, or a list of byte values for binary content."""
    if isinstance(value, list):
        if all(isinstance(item, int) for item in value):
            return bytes(value).decode("utf-8", "replace")
        return field_text(value[0]) if value else ""
    return "" if value is None else str(value)


def read_journal(config, since, until, levels, all_levels):
    """The journal entries of the time span; all levels if include rules need them."""
    command = [config["journalctl"], "-o", "json", "--no-pager", "-q",
               "--since", "@%.6f" % since, "--until", "@%.6f" % until] + [str(arg) for arg in config["args"]]
    if not all_levels:
        command += ["-p", "%d..%d" % (min(levels), max(levels))]
    try:
        process = subprocess.Popen(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
    except OSError as err:
        raise LogDigestError("cannot run %s: %s" % (config["journalctl"], err.strerror))
    entries = []
    for line in process.stdout:
        try:
            record = json.loads(line)
        except ValueError:
            continue
        try:
            stamp = int(record["__REALTIME_TIMESTAMP"]) / 1e6
        except (KeyError, ValueError):
            continue
        program = field_text(record.get("SYSLOG_IDENTIFIER")) or field_text(record.get("_COMM"))
        if not program and record.get("_TRANSPORT") == "kernel":
            program = "kernel"
        try:
            priority = int(field_text(record.get("PRIORITY")) or 6)
        except ValueError:
            priority = 6
        entries.append(Entry(stamp, field_text(record.get("_HOSTNAME")), program,
                             field_text(record.get("SYSLOG_PID")) or field_text(record.get("_PID")),
                             field_text(record.get("MESSAGE")), priority, "journal",
                             unit=field_text(record.get("_SYSTEMD_UNIT"))))
    stderr = process.stderr.read().decode("utf-8", "replace").strip()
    if process.wait() != 0:
        raise LogDigestError("%s failed: %s" % (config["journalctl"], stderr or "exit code %d" % process.returncode))
    return entries


# ---------------------------------------------------------------------------
# rsyslog's configuration: which files it writes, with which severities
# ---------------------------------------------------------------------------

FACILITIES = ["kern", "user", "mail", "daemon", "auth", "syslog", "lpr", "news", "uucp", "cron", "authpriv",
              "ftp", "ntp", "security", "console", "clock"] + ["local%d" % n for n in range(8)]

SELECTOR = r"[\w*]+(?:,[\w*]+)*\.!?=?[\w*]+"
SELECTOR_RE = re.compile(r"^(%s(?:;\s*%s)*)\s+(.+)$" % (SELECTOR, SELECTOR), re.S)
PARAM_RE = re.compile(r"""([\w.-]+)\s*=\s*(?:"((?:[^"\\]|\\.)*)"|'((?:[^'\\]|\\.)*)')""")
IGNORED_STATEMENTS = re.compile(r"^(module|input|global|main_queue|lookup_table|parser|timezone|license|set|unset|"
                                r"call|reset|foreach|dyn_stats|percentile_stats|ratelimit)\b", re.I)


def selector_levels(selector):
    """The severities a legacy selector such as *.*;mail.none or *.=warning;*.=err lets through."""
    per_facility = collections.defaultdict(set)
    for part in selector.split(";"):
        part = part.strip()
        if not part:
            continue
        facilities, _, spec = part.rpartition(".")
        names = FACILITIES if "*" in facilities.split(",") else [name.strip() for name in facilities.split(",")]
        negate = spec.startswith("!")
        spec = spec.lstrip("!")
        exact = spec.startswith("=")
        spec = spec.lstrip("=")
        if spec == "none":
            for name in names:
                per_facility[name] = set()
            continue
        if spec == "*":
            chosen = set(ALL_LEVELS)
        else:
            level = parse_level(spec)
            chosen = {level} if exact else set(range(0, level + 1))
        for name in names:
            if negate:
                per_facility[name] -= chosen
            else:
                per_facility[name] |= chosen
    result = set()
    for levels in per_facility.values():
        result |= levels
    return frozenset(result)


def condition_levels(condition):
    """The severities an if condition lets through; all of them unless it is a plain conjunction."""
    if re.search(r"\bor\b|\bnot\b", condition):
        return ALL_LEVELS
    levels = set(ALL_LEVELS)
    for match in re.finditer(r"\$syslogseverity(?:-text)?\s*(==|!=|<=|>=|<|>)\s*['\"]?(\w+)['\"]?", condition):
        operator, value = match.groups()
        try:
            level = parse_level(value)
        except ValueError:
            return ALL_LEVELS
        allowed = {"==": {level}, "!=": ALL_LEVELS - {level}, "<=": set(range(0, level + 1)),
                   "<": set(range(0, level)), ">=": set(range(level, 8)), ">": set(range(level + 1, 8))}[operator]
        levels &= allowed
    for match in re.finditer(r"prifilt\(\s*['\"]([^'\"]+)['\"]\s*\)", condition):
        try:
            levels &= selector_levels(match.group(1))
        except ValueError:
            return ALL_LEVELS
    return frozenset(levels)


def strip_comments(text):
    """rsyslog's configuration without # and /* */ comments, outside of quoted strings."""
    out = []
    index = 0
    quote = None
    while index < len(text):
        char = text[index]
        if quote:
            out.append(char)
            if char == "\\" and index + 1 < len(text):
                out.append(text[index + 1])
                index += 2
                continue
            if char == quote:
                quote = None
            index += 1
        elif char == '"':
            quote = char
            out.append(char)
            index += 1
        elif char == "#":
            end = text.find("\n", index)
            index = len(text) if end < 0 else end
        elif text.startswith("/*", index) and not (index and (text[index - 1].isalnum() or text[index - 1] in "._-/")):
            # a comment, unlike the glob in $IncludeConfig /etc/rsyslog.d/*.conf
            end = text.find("*/", index + 2)
            index = len(text) if end < 0 else end + 2
        else:
            out.append(char)
            index += 1
    return "".join(out)


def split_statements(text):
    """Statements of rsyslog's configuration: lines, with parentheses spanning lines, braces separate."""
    text = re.sub(r"\\[ \t]*\n", " ", text)
    statements = []
    buffer = []
    depth = 0
    quote = None
    index = 0
    while index < len(text):
        char = text[index]
        if quote:
            buffer.append(char)
            if char == "\\" and index + 1 < len(text):
                buffer.append(text[index + 1])
                index += 2
                continue
            if char == quote:
                quote = None
        elif char == '"':
            quote = char
            buffer.append(char)
        elif char == "(":
            depth += 1
            buffer.append(char)
        elif char == ")":
            depth = max(0, depth - 1)
            buffer.append(char)
        elif depth == 0 and char in "{}\n":
            statement = "".join(buffer).strip()
            if statement:
                statements.append(statement)
            buffer = []
            if char != "\n":
                statements.append(char)
        else:
            buffer.append(char)
        index += 1
    statement = "".join(buffer).strip()
    if statement:
        statements.append(statement)
    return statements


def parameters(text):
    return {key.lower(): (double if double is not None else single) for key, double, single in PARAM_RE.findall(text)}


class RsyslogConfig:
    """The files rsyslog's configuration writes, each with the severities it can hold."""

    def __init__(self, path):
        self.templates = {}
        self.targets = collections.OrderedDict()   # file or glob -> set of levels
        self.problems = []
        self._last_levels = ALL_LEVELS
        self._depth = 0
        self._read(path, ALL_LEVELS, required=True)

    def _read(self, path, levels, required=False):
        if self._depth > 20:
            self.problems.append("rsyslog includes nest too deeply at %s" % path)
            return
        try:
            with open(path, encoding="utf-8", errors="replace") as handle:
                text = handle.read()
        except OSError as err:
            if required:
                raise LogDigestError("cannot read %s: %s" % (path, err.strerror))
            self.problems.append("cannot read %s: %s" % (path, err.strerror))
            return
        self._depth += 1
        try:
            self._process(split_statements(strip_comments(text)), levels)
        finally:
            self._depth -= 1

    def _include(self, pattern, levels):
        if os.path.isdir(pattern):
            paths = sorted(os.path.join(pattern, name) for name in os.listdir(pattern))
        else:
            paths = sorted(glob.glob(pattern))
        for path in paths:
            if os.path.isfile(path):
                self._read(path, levels)

    def _process(self, statements, levels):
        stack = [levels]
        pending = None          # levels for the block or statement that follows
        pending_condition = None
        for statement in statements:
            context = stack[-1]
            if statement == "{":
                stack.append(pending if pending is not None else context)
                pending = None
                continue
            if statement == "}":
                if len(stack) > 1:
                    stack.pop()
                pending = None
                continue
            if pending_condition is not None:
                statement = "if " + pending_condition + " " + statement
                pending_condition = None
            lower = statement.lower()
            if re.match(r"^if\b", lower):
                match = re.search(r"\bthen\b", statement)
                if not match:
                    pending_condition = statement[2:]
                    continue
                allowed = condition_levels(statement[2:match.start()]) & context
                rest = statement[match.end():].strip()
                if rest:
                    self._statement(rest, allowed)
                else:
                    pending = allowed
                continue
            if re.match(r"^else\b", lower):
                rest = statement[4:].strip()
                if rest:
                    self._statement(rest, context)
                else:
                    pending = context
                continue
            if lower.startswith("ruleset("):
                pending = context
                continue
            self._statement(statement, context)

    def _statement(self, statement, context):
        lower = statement.lower()
        if statement.startswith("$"):
            self._directive(statement, context)
            return
        if lower.startswith("include("):
            path = parameters(statement).get("file")
            if path:
                self._include(path, context)
            return
        if lower.startswith("template("):
            params = parameters(statement)
            if params.get("name") and params.get("string") is not None:
                self.templates[params["name"]] = params["string"]
            return
        if IGNORED_STATEMENTS.match(statement) or statement in ("stop", "~"):
            return
        if statement.startswith("&"):
            self._action(statement[1:].strip(), self._last_levels)
            return
        if statement.startswith(":"):
            # a property based filter: the severities it lets through are not known
            match = re.match(r'^:\s*[^,]+,\s*!?\s*[\w-]+\s*,\s*"(?:[^"\\]|\\.)*"\s*(.*)$', statement, re.S)
            self._last_levels = context
            if match:
                self._action(match.group(1), context)
            return
        match = SELECTOR_RE.match(statement)
        if match:
            try:
                levels = selector_levels(match.group(1)) & context
            except ValueError as err:
                self.problems.append("rsyslog selector %r: %s" % (match.group(1), err))
                return
            self._last_levels = levels
            self._action(match.group(2), levels)
            return
        self._action(statement, context)

    def _directive(self, statement, context):
        match = re.match(r"^\$IncludeConfig\s+(\S+)", statement, re.I)
        if match:
            self._include(match.group(1), context)
            return
        match = re.match(r'^\$template\s+([^,\s]+)\s*,\s*"((?:[^"\\]|\\.)*)"', statement, re.I)
        if match:
            self.templates[match.group(1)] = match.group(2)

    def _action(self, action, levels):
        action = action.strip()
        path = dynamic = None
        match = re.match(r"^action\((.*)\)\s*$", action, re.S | re.I)
        if match:
            params = parameters(match.group(1))
            if params.get("type", "").lower() != "omfile":
                return
            path = params.get("file")
            dynamic = params.get("dynafile")
        else:
            match = re.match(r"^-?(/[^;\s]+)", action)
            if match:
                path = match.group(1)
            match = re.match(r"^-?\?([^;\s]+)", action)
            if match:
                dynamic = match.group(1)
        if dynamic:
            template = self.templates.get(dynamic)
            if template is None:
                self.problems.append("rsyslog uses the unknown template %r as a file name" % dynamic)
                return
            path = re.sub(r"\*+", "*", re.sub(r"%[^%]*%", "*", template))
        if not path or path.startswith("/dev/"):
            return
        if levels:
            self.targets.setdefault(path, set()).update(levels)


# ---------------------------------------------------------------------------
# reading rsyslog's files
# ---------------------------------------------------------------------------

MONTHS = {name: number for number, name in enumerate(
    ("Jan", "Feb", "Mar", "Apr", "May", "Jun", "Jul", "Aug", "Sep", "Oct", "Nov", "Dec"), 1)}
RFC5424_RE = re.compile(r"^<(?P<pri>\d{1,3})>1 (?P<ts>\S+) (?P<host>\S+) (?P<app>\S+) (?P<procid>\S+) \S+ "
                        r"(?:-|(?:\[(?:[^\]\\]|\\.)*\])+) ?(?P<msg>.*)$")
LINE_RE = re.compile(r"^(?:<(?P<pri>\d{1,3})>)?(?P<ts>\d{4}-\d\d-\d\dT\d\d:\d\d:\d\d(?:\.\d+)?(?:Z|[+-]\d\d:?\d\d)|"
                     r"[A-Z][a-z]{2} [ \d]\d \d\d:\d\d:\d\d) (?P<host>\S+) ?(?P<rest>.*)$")
TAG_RE = re.compile(r"^(?P<program>[^\s\[\]:]+)(?:\[(?P<pid>[^\]]*)\])?:\s?(?P<msg>.*)$")
REPEATED_RE = re.compile(r"^(last message repeated \d+ times|message repeated \d+ times: \[)")
ROTATED_RE = re.compile(r"^(?:\.\d+|-\d{8}(?:\d{2})?)(?:\.(?:gz|xz|bz2))?$")


def rfc3339_time(text):
    text = text.replace("Z", "+00:00")
    match = re.match(r"^(.*?T\d\d:\d\d:\d\d)(\.\d+)?([+-]\d\d):?(\d\d)$", text)
    if not match:
        raise ValueError(text)
    head, fraction, hours, minutes = match.groups()
    fraction = (fraction or ".0")[:7]
    return datetime.datetime.fromisoformat("%s%s%s:%s" % (head, fraction, hours, minutes)).timestamp()


def traditional_time(text, until):
    """A time stamp like 'Oct  2 18:28:01', in local time, of the year that puts it before until."""
    month, day, clock = text.split()
    hour, minute, second = (int(part) for part in clock.split(":"))
    year = time.localtime(until).tm_year
    stamp = datetime.datetime(year, MONTHS[month], int(day), hour, minute, second).timestamp()
    if stamp > until + 2 * 86400:
        stamp = datetime.datetime(year - 1, MONTHS[month], int(day), hour, minute, second).timestamp()
    return stamp


def parse_line(line, until):
    """time, host, program, pid, message and severity (or None) of a line rsyslog wrote."""
    match = RFC5424_RE.match(line)
    if match:
        message = match.group("msg")
        if message.startswith("\ufeff"):
            message = message[1:]
        procid = match.group("procid")
        return (rfc3339_time(match.group("ts")), match.group("host"), match.group("app").strip("-"),
                "" if procid == "-" else procid, message, int(match.group("pri")) & 7)
    match = LINE_RE.match(line)
    if not match:
        return None
    stamp = match.group("ts")
    when = rfc3339_time(stamp) if stamp[0].isdigit() else traditional_time(stamp, until)
    rest = match.group("rest")
    tag = TAG_RE.match(rest)
    if tag:
        program, pid, message = tag.group("program"), tag.group("pid") or "", tag.group("msg")
    else:
        program, pid, message = "", "", rest
    severity = int(match.group("pri")) & 7 if match.group("pri") else None
    return when, match.group("host"), program, pid, message, severity


def log_files(pattern, since):
    """The files of a configured file or glob, with their rotated copies, that may hold lines after since."""
    found = []
    for path in sorted(glob.glob(pattern)):
        candidates = [path] + sorted(other for other in glob.glob(glob.escape(path) + "*")
                                     if ROTATED_RE.match(other[len(path):]))
        for candidate in candidates:
            try:
                if os.path.isfile(candidate) and os.path.getmtime(candidate) >= since:
                    found.append(candidate)
            except OSError:
                continue
    return found


def open_log(path):
    if path.endswith(".gz"):
        return gzip.open(path, "rb")
    if path.endswith(".xz"):
        return lzma.open(path, "rb")
    if path.endswith(".bz2"):
        return bz2.open(path, "rb")
    return open(path, "rb")


def seek_near(handle, size, target, until):
    """Move a plain file to shortly before the first line at target, by bisection; lines are mostly in time order."""
    low, high = 0, size
    while high - low > 1 << 16:
        middle = (low + high) // 2
        handle.seek(middle)
        handle.readline()
        stamp = None
        for _ in range(50):
            line = handle.readline()
            if not line:
                break
            parsed = parse_line(line.decode("utf-8", "replace").rstrip("\n"), until)
            if parsed:
                stamp = parsed[0]
                break
        if stamp is None or stamp >= target:
            high = middle
        else:
            low = middle
    handle.seek(low)
    if low:
        handle.readline()


def read_log_file(path, levels_of_file, levels, since, until, margin):
    """The lines of one file between since and until, as entries.

    A line's severity is its own, if the file's format includes it (<PRI>).
    Otherwise it is known only if every severity the file can hold is
    selected: then it is the least severe of them, marked as approximate.
    """
    entries = []
    if levels_of_file <= levels:
        file_level, approximate = max(levels_of_file), len(levels_of_file) > 1
    else:
        file_level, approximate = None, False
    with open_log(path) as handle:
        if path.endswith((".gz", ".xz", ".bz2")):
            lines = handle
        else:
            seek_near(handle, os.fstat(handle.fileno()).st_size, since - margin - 600, until)
            lines = handle
        for raw in lines:
            line = raw.decode("utf-8", "replace").rstrip("\n")
            try:
                parsed = parse_line(line, until)
            except (ValueError, KeyError, OverflowError):
                continue
            if not parsed:
                continue
            stamp, host, program, pid, message, severity = parsed
            if stamp < since - margin or stamp > until + margin or REPEATED_RE.match(message):
                continue
            if severity is not None:
                entries.append(Entry(stamp, host, program, pid, message, severity, path))
            else:
                entries.append(Entry(stamp, host, program, pid, message, file_level, path,
                                     approximate=approximate))
    return entries


# ---------------------------------------------------------------------------
# selecting and merging messages
# ---------------------------------------------------------------------------

def normalize_message(text):
    # rsyslog writes control characters as #011 etc.
    text = re.sub(r"#([0-3][0-7]{2})", lambda match: chr(int(match.group(1), 8)), text)
    return " ".join(text.split())[:1000]


def short_host(host, local):
    return (host or local).lower().split(".")[0]


def merge_duplicates(entries, window, local_host):
    """Merge a message that appears in several origins (the journal, rsyslog files) into one.

    Repetitions within one origin are kept: they are separate messages.
    The journal's copy is kept, as it carries the exact severity.
    """
    entries = sorted(entries, key=lambda entry: (entry.time, entry.origin != "journal"))
    kept = []
    by_key = collections.defaultdict(list)
    merged = 0
    for entry in entries:
        key = (short_host(entry.host, local_host), entry.program[:32].lower(), normalize_message(entry.message))
        candidates = by_key[key]
        hit = None
        for record in reversed(candidates):
            if entry.time - record[0].time > window:
                break
            if entry.origin not in record[1]:
                hit = record
                break
        if hit is not None:
            hit[1].add(entry.origin)
            if entry.origin == "journal":
                hit[0] = entry
            merged += 1
            continue
        record = [entry, {entry.origin}]
        candidates.append(record)
        kept.append(record)
    return sorted((record[0] for record in kept), key=lambda entry: entry.time), merged


def selected(entry, levels, include, exclude):
    if any(rule.matches(entry) for rule in exclude):
        return False
    if any(rule.matches(entry) for rule in include):
        return True
    return entry.priority is not None and entry.priority in levels


# ---------------------------------------------------------------------------
# the digest
# ---------------------------------------------------------------------------

def table(rows, header):
    widths = [max(len(str(row[index])) for row in [header] + rows) for index in range(len(header))]
    lines = []
    for row in [header] + rows:
        cells = [str(cell).rjust(width) if isinstance(cell, int) else str(cell).ljust(width)
                 for cell, width in zip(row, widths)]
        lines.append("  " + "  ".join(cells).rstrip())
    return lines


def describe_levels(levels):
    if levels == frozenset(range(0, max(levels) + 1)):
        return "%s or more severe" % LEVELS[max(levels)]
    return "%s to %s" % (LEVELS[min(levels)], LEVELS[max(levels)])


def build_digest(entries, run):
    """Subject, full digest, short summary and whether the digest is of high priority."""
    count = len(entries)
    high = count > run["max_lines"]
    span = "%s - %s" % (format_time(run["since"]), format_time(run["until"]))
    subject = run["subject"].format(host=run["host"], count=count, span=span)
    if high:
        subject = "HIGH PRIORITY: " + subject

    head = ["Log digest of %s for %s" % (run["host"], span),
            "Severity: %s%s%s" % (
                describe_levels(run["levels"]),
                "; %d include rule(s)" % len(run["include"]) if run["include"] else "",
                "; %d exclude rule(s)" % len(run["exclude"]) if run["exclude"] else ""),
            "Sources: %s" % ", ".join(run["sources"]) if run["sources"] else "Sources: none"]
    by_origin = collections.Counter("journal" if entry.origin == "journal" else "rsyslog" for entry in entries)
    head.append("%d message(s)%s%s" % (
        count, " (%s)" % ", ".join("%s: %d" % item for item in sorted(by_origin.items())) if by_origin else "",
        ", %d duplicate(s) merged" % run["merged"] if run["merged"] else ""))
    if run["problems"]:
        head.append("")
        head.append("Problems:")
        head.extend("  " + problem for problem in run["problems"])

    lines = list(head)
    summary = list(head)
    if entries:
        counts = collections.Counter((entry.host or run["host"], entry.program or "?") for entry in entries)
        errors = collections.Counter((entry.host or run["host"], entry.program or "?") for entry in entries
                                     if entry.priority is not None and entry.priority <= 3)
        ranked = sorted(counts.items(), key=lambda item: (-errors[item[0]], -item[1], item[0]))
        rows = [[total, errors[key], key[0], key[1]] for key, total in ranked[:run["top"]]]
        overview = ["", "Most messages by host and program:"] + table(rows, ["Count", "Errors", "Host", "Program"])
        if len(ranked) > run["top"]:
            overview.append("  ... and %d more" % (len(ranked) - run["top"]))
        lines += overview
        summary += overview
        lines.append("")
        if high:
            lines.append("More than %d messages, so they are not listed. To see them:" % run["max_lines"])
            lines.append("  log_digest --print --since %s --until %s" % (
                shlex.quote(format_time(run["since"])), shlex.quote(format_time(run["until"] + 1))))
        else:
            lines.append("Messages:")
            for entry in entries:
                program = entry.program + ("[%s]" % entry.pid if entry.pid else "")
                lines.append("%s %s %s %s: %s" % (format_time(entry.time), entry.host or run["host"],
                                                  program or "?", level_label(entry), entry.message))
    return subject, "\n".join(lines) + "\n", "\n".join(summary) + "\n", high


# ---------------------------------------------------------------------------
# sending
# ---------------------------------------------------------------------------

def send_mail(config, subject, body, high, host):
    message = email.message.EmailMessage()
    message["From"] = config.get("from") or "log_digest@%s" % host
    message["To"] = ", ".join(config["to"])
    message["Subject"] = subject
    message["Date"] = email.utils.formatdate(localtime=True)
    message["Message-ID"] = email.utils.make_msgid(domain=host)
    message["Auto-Submitted"] = "auto-generated"
    if high:
        message["X-Priority"] = "1 (Highest)"
        message["Importance"] = "high"
        message["Priority"] = "urgent"
    message.set_content(body)
    if config.get("method", "smtp") == "sendmail":
        result = subprocess.run([config["sendmail"], "-t", "-oi"], input=message.as_bytes(),
                                stdout=subprocess.PIPE, stderr=subprocess.PIPE, timeout=120)
        if result.returncode != 0:
            raise LogDigestError("%s failed: %s" % (config["sendmail"], result.stderr.decode("utf-8", "replace").strip()))
        return
    try:
        with smtplib.SMTP(config.get("host", "localhost"), int(config.get("port", 25)), timeout=60) as smtp:
            smtp.send_message(message)
    except (OSError, smtplib.SMTPException) as err:
        raise LogDigestError("cannot send mail via %s:%s: %s" % (config.get("host"), config.get("port"), err))


def send_xmpp(config, recipients, text):
    command = [config["go_sendxmpp"]]
    if config.get("config"):
        command += ["-f", config["config"]]
    if config.get("chatroom"):
        command.append("-c")
    command += [str(arg) for arg in config.get("args", [])] + list(recipients)
    try:
        result = subprocess.run(command, input=text.encode("utf-8"), stdout=subprocess.PIPE, stderr=subprocess.PIPE,
                                timeout=120)
    except OSError as err:
        raise LogDigestError("cannot run %s: %s" % (config["go_sendxmpp"], err.strerror))
    except subprocess.TimeoutExpired:
        raise LogDigestError("%s did not finish within 120s" % config["go_sendxmpp"])
    if result.returncode != 0:
        raise LogDigestError("%s failed: %s" % (config["go_sendxmpp"],
                                                result.stderr.decode("utf-8", "replace").strip() or result.returncode))


# ---------------------------------------------------------------------------
# configuration
# ---------------------------------------------------------------------------

def merge_config(base, override):
    result = dict(base)
    for key, value in override.items():
        if isinstance(value, dict) and isinstance(result.get(key), dict):
            result[key] = merge_config(result[key], value)
        else:
            result[key] = value
    return result


def load_yaml(path):
    try:
        import yaml
    except ImportError:
        raise LogDigestError("reading %s needs PyYAML (python3-PyYAML / python3-yaml)" % path)
    try:
        with open(path, encoding="utf-8") as handle:
            data = yaml.safe_load(handle) or {}
    except OSError as err:
        raise LogDigestError("cannot read %s: %s" % (path, err.strerror))
    except yaml.YAMLError as err:
        raise LogDigestError("invalid YAML in %s: %s" % (path, err))
    if not isinstance(data, dict):
        raise LogDigestError("%s must contain a mapping" % path)
    return data


def build_parser():
    parser = argparse.ArgumentParser(
        prog="log_digest",
        description="Send a digest of the warnings and errors of a time span from the systemd journal and the "
                    "files rsyslog writes, by mail and XMPP. Messages that appear in both are counted once.",
        epilog="--since and --until take the syntax of journalctl, e.g. -1h, \"1 hour ago\", today, "
               "\"2026-10-02 08:00\". Options given here override the configuration file.")
    parser.add_argument("-V", "--version", action="version", version="log_digest %s" % VERSION)
    parser.add_argument("-C", "--config", help="YAML configuration file")
    parser.add_argument("-S", "--since", help="start of the time span (default: -1h)")
    parser.add_argument("-U", "--until", help="end of the time span (default: now)")
    parser.add_argument("-p", "--priority",
                        help="severity, like journalctl -p: a level and all more severe ones, or FROM..TO "
                             "(default: warning)")
    parser.add_argument("-n", "--max-lines", type=int,
                        help="list the messages only up to this many, above it send with high priority (default: 100)")
    parser.add_argument("--top", type=int, help="rows of the host and program table (default: 10)")
    parser.add_argument("--include", action="append", metavar="FIELD=REGEX",
                        help="report messages that match, whatever their severity; FIELD is host, program, "
                             "message, unit or origin; repeatable")
    parser.add_argument("--exclude", action="append", metavar="FIELD=REGEX",
                        help="leave out messages that match; repeatable")
    parser.add_argument("--no-journal", action="store_true", help="do not read the journal")
    parser.add_argument("--no-rsyslog", action="store_true", help="do not read rsyslog's files")
    parser.add_argument("--rsyslog-config", help="rsyslog's configuration (default: /etc/rsyslog.conf)")
    parser.add_argument("--journal-arg", action="append", metavar="ARG",
                        help="argument for journalctl instead of --merge, e.g. --directory=/var/log/journal/remote; "
                             "repeatable")
    parser.add_argument("--mail-to", action="append", metavar="ADDRESS", help="mail recipient; repeatable")
    parser.add_argument("--xmpp-to", action="append", metavar="JID", help="XMPP recipient; repeatable")
    parser.add_argument("--print", action="store_true", help="print the digest instead of sending it")
    parser.add_argument("--show-sources", action="store_true",
                        help="print the files rsyslog writes with the severities they hold, and exit")
    parser.add_argument("-v", "--verbose", action="store_true", help="explain what is read on stderr")
    return parser


def load_config(args):
    config = merge_config(DEFAULTS, load_yaml(args.config)) if args.config else dict(DEFAULTS)
    cli = {"since": args.since, "until": args.until, "priority": args.priority, "max_lines": args.max_lines,
           "top": args.top}
    config.update({key: value for key, value in cli.items() if value is not None})
    if args.include:
        config["include"] = list(config.get("include") or []) + args.include
    if args.exclude:
        config["exclude"] = list(config.get("exclude") or []) + args.exclude
    if args.no_journal:
        config["journal"] = dict(config["journal"], enabled=False)
    if args.no_rsyslog:
        config["rsyslog"] = dict(config["rsyslog"], enabled=False)
    if args.rsyslog_config:
        config["rsyslog"] = dict(config["rsyslog"], config=args.rsyslog_config)
    if args.journal_arg:
        config["journal"] = dict(config["journal"], args=args.journal_arg)
    if args.mail_to:
        config["mail"] = dict(config["mail"], to=args.mail_to)
    if args.xmpp_to:
        config["xmpp"] = dict(config["xmpp"], to=args.xmpp_to)
    for section in ("mail", "xmpp"):
        for key in ("to", "high_priority_to"):
            value = config[section].get(key)
            if isinstance(value, str):
                config[section] = dict(config[section], **{key: [value]})
    return config


def enabled(value):
    """True, False or "auto" from the configuration."""
    if isinstance(value, str):
        value = value.strip().lower()
        if value == "auto":
            return "auto"
        return value in ("1", "yes", "true", "on")
    return bool(value)


# ---------------------------------------------------------------------------
# main
# ---------------------------------------------------------------------------

def collect(config, since, until, levels, include, verbose, problems, sources):
    entries = []
    margin = float(config["duplicate_window"])

    journal = config["journal"]
    use_journal = enabled(journal.get("enabled", "auto"))
    if use_journal == "auto":
        use_journal = shutil.which(journal["journalctl"]) is not None
    if use_journal:
        try:
            found = read_journal(journal, since, until, levels, bool(include))
            entries += found
            sources.append("journal")
            if verbose:
                sys.stderr.write("journal: %d entries\n" % len(found))
        except LogDigestError as err:
            problems.append(str(err))

    rsyslog = config["rsyslog"]
    use_rsyslog = enabled(rsyslog.get("enabled", "auto"))
    if use_rsyslog == "auto":
        use_rsyslog = os.path.exists(rsyslog["config"])
    if not use_rsyslog:
        return entries
    targets = collections.OrderedDict()
    try:
        parsed = RsyslogConfig(rsyslog["config"])
        targets.update(parsed.targets)
        problems += parsed.problems
    except LogDigestError as err:
        problems.append(str(err))
    for extra in rsyslog.get("files") or []:
        try:
            targets.setdefault(extra["path"], set()).update(parse_priority(extra.get("priority", "debug")))
        except (KeyError, TypeError, ValueError) as err:
            problems.append("invalid rsyslog.files entry %r: %s" % (extra, err))
    read_any = False
    for pattern, file_levels in targets.items():
        if not file_levels:
            continue
        if not file_levels & levels and not include:
            if verbose:
                sys.stderr.write("rsyslog: skipping %s, it holds none of the severities\n" % pattern)
            continue
        for path in log_files(pattern, since - margin):
            try:
                found = read_log_file(path, frozenset(file_levels), levels, since, until, margin)
            except (OSError, EOFError, lzma.LZMAError) as err:
                problems.append("cannot read %s: %s" % (path, getattr(err, "strerror", None) or err))
                continue
            read_any = True
            entries += found
            if verbose:
                sys.stderr.write("rsyslog: %s (%s): %d lines\n" % (
                    path, ",".join(LEVELS[level] for level in sorted(file_levels)), len(found)))
    if read_any or targets:
        sources.append("rsyslog")
    return entries


def join_time_arguments(argv):
    """Let --since -1h work as in journalctl: argparse would take -1h for an option."""
    names = {"-S": "--since", "--since": "--since", "-U": "--until", "--until": "--until"}
    result = []
    index = 0
    while index < len(argv):
        arg = argv[index]
        if arg in names and index + 1 < len(argv):
            result.append("%s=%s" % (names[arg], argv[index + 1]))
            index += 2
            continue
        result.append(arg)
        index += 1
    return result


def run(argv=None):
    args = build_parser().parse_args(join_time_arguments(sys.argv[1:] if argv is None else list(argv)))
    config = load_config(args)
    host = socket.getfqdn()
    local_short = host.split(".")[0]

    if args.show_sources:
        parsed = RsyslogConfig(config["rsyslog"]["config"])
        for pattern, levels in parsed.targets.items():
            print("%-40s %s" % (pattern, ",".join(LEVELS[level] for level in sorted(levels))))
        for problem in parsed.problems:
            print("problem: %s" % problem)
        return 0

    try:
        levels = parse_priority(config["priority"])
        include = [parse_rule(rule) for rule in config.get("include") or []]
        exclude = [parse_rule(rule) for rule in config.get("exclude") or []]
        now = time.time()
        since = parse_time(config["since"], now)
        until = parse_time(config["until"], now)
    except ValueError as err:
        raise LogDigestError(str(err))
    if since > until:
        raise LogDigestError("--since is after --until")

    problems = []
    sources = []
    entries = collect(config, since, until, levels, include, args.verbose, problems, sources)
    entries = [entry for entry in entries if since <= entry.time <= until and selected(entry, levels, include, exclude)]
    entries, merged = merge_duplicates(entries, float(config["duplicate_window"]), local_short)

    digest = {"host": host, "since": since, "until": until, "levels": levels, "include": include,
              "exclude": exclude, "sources": sources, "merged": merged, "problems": problems,
              "max_lines": int(config["max_lines"]), "top": int(config["top"]), "subject": config["subject"]}
    subject, body, summary, high = build_digest(entries, digest)

    mail, xmpp = config["mail"], config["xmpp"]
    if args.print or not (mail.get("to") or xmpp.get("to")):
        sys.stdout.write("Subject: %s\n\n%s" % (subject, body))
        return 0
    if not entries and not problems and not config.get("send_empty"):
        if args.verbose:
            sys.stderr.write("nothing to send\n")
        return 0
    failures = []
    if mail.get("to"):
        try:
            send_mail(mail, subject, body, high, host)
        except LogDigestError as err:
            failures.append(str(err))
    if xmpp.get("to"):
        recipients = list(xmpp["to"]) + (list(xmpp.get("high_priority_to") or []) if high else [])
        text = subject + "\n\n" + (body if xmpp.get("full") else summary)
        try:
            send_xmpp(xmpp, recipients, text)
        except LogDigestError as err:
            failures.append(str(err))
    for failure in failures:
        sys.stderr.write("log_digest: %s\n" % failure)
    return 1 if failures else 0


def main():
    try:
        sys.exit(run())
    except LogDigestError as err:
        sys.stderr.write("log_digest: %s\n" % err)
        sys.exit(2)
    except KeyboardInterrupt:
        sys.exit(130)


if __name__ == "__main__":
    main()
