#!/usr/bin/env python3
"""Minimal RFC 5321 SMTP sink used to capture Keycloak password-reset emails.

Listens on 0.0.0.0:PORT and appends every accepted message (headers + body)
to the capture file given as argv[2]. Dependency-free (Python 3.12+ removed
smtpd, so this is a hand-rolled single-threaded server)."""

import socketserver
import sys

LISTEN_PORT = int(sys.argv[1]) if len(sys.argv) > 1 else 1025
CAPTURE_FILE = sys.argv[2] if len(sys.argv) > 2 else "/tmp/smtp_capture.log"


class SMTPHandler(socketserver.StreamRequestHandler):
    def _send(self, line: str) -> None:
        self.wfile.write((line + "\r\n").encode("utf-8", "replace"))
        self.wfile.flush()

    def handle(self) -> None:
        self._send("220 smtp-sink ESMTP ready")
        mail_from = None
        rcpt = []
        while True:
            raw = self.rfile.readline()
            if not raw:
                return
            line = raw.decode("utf-8", "replace").rstrip("\r\n")
            upper = line.upper()
            if upper.startswith(("EHLO", "HELO")):
                self._send("250-smtp-sink greets you")
                self._send("250 8BITMIME")
            elif upper.startswith("MAIL FROM:"):
                mail_from = line
                rcpt = []
                self._send("250 OK")
            elif upper.startswith("RCPT TO:"):
                rcpt.append(line)
                self._send("250 OK")
            elif upper == "DATA":
                if mail_from is None or not rcpt:
                    self._send("503 Bad sequence")
                    continue
                self._send("354 End data with <CR><LF>.<CR><LF>")
                data_lines = []
                while True:
                    d = self.rfile.readline()
                    if not d:
                        break
                    decoded = d.decode("utf-8", "replace")
                    if decoded.rstrip("\r\n") == ".":
                        break
                    data_lines.append(decoded)
                with open(CAPTURE_FILE, "a", encoding="utf-8") as fh:
                    fh.write("=== MESSAGE ===\n")
                    fh.write(mail_from + "\n")
                    fh.write("\n".join(rcpt) + "\n")
                    fh.write("".join(data_lines))
                    fh.write("\n=== END ===\n")
                self._send("250 Message accepted")
                mail_from = None
                rcpt = []
            elif upper == "RSET":
                mail_from = None
                rcpt = []
                self._send("250 OK")
            elif upper == "NOOP":
                self._send("250 OK")
            elif upper == "QUIT":
                self._send("221 Bye")
                return
            elif upper.startswith("STARTTLS"):
                self._send("454 TLS not available")
            elif upper.startswith("AUTH"):
                self._send("535 Authentication not supported")
            else:
                self._send("250 OK")


class SMTPServer(socketserver.ThreadingTCPServer):
    allow_reuse_address = True
    daemon_threads = True


if __name__ == "__main__":
    server = SMTPServer(("0.0.0.0", LISTEN_PORT), SMTPHandler)
    print(f"smtp-sink listening on 0.0.0.0:{LISTEN_PORT}, capture -> {CAPTURE_FILE}", flush=True)
    server.serve_forever()
