"""SMTP receiver. Listens for incoming mail and stores it in MySQL.

Locally:  python smtp_server.py         (listens on port 2525)
On a VPS: set TM_SMTP_PORT=25 and point your domain's MX record here.

Only mail addressed to a known inbox (an address we created) is kept;
everything else is accepted-and-dropped so senders don't get errors.
"""
import asyncio
from email import message_from_bytes
from email.policy import default as default_policy

from aiosmtpd.controller import Controller

import config
import db


def _extract_bodies(msg):
    """Return (text, html) from a parsed email.message.EmailMessage."""
    text, html = None, None
    if msg.is_multipart():
        for part in msg.walk():
            ctype = part.get_content_type()
            disp = part.get("Content-Disposition", "")
            if "attachment" in (disp or ""):
                continue
            try:
                payload = part.get_content()
            except Exception:
                continue
            if ctype == "text/plain" and text is None:
                text = payload
            elif ctype == "text/html" and html is None:
                html = payload
    else:
        payload = msg.get_content()
        if msg.get_content_type() == "text/html":
            html = payload
        else:
            text = payload
    return text, html


class Handler:
    async def handle_RCPT(self, server, session, envelope, address, rcpt_options):
        envelope.rcpt_tos.append(address)
        return "250 OK"

    async def handle_DATA(self, server, session, envelope):
        msg = message_from_bytes(envelope.content, policy=default_policy)
        subject = str(msg.get("Subject", ""))
        from_addr = str(msg.get("From", envelope.mail_from or ""))
        text, html = _extract_bodies(msg)

        conn = db.get_conn()
        try:
            with conn.cursor() as cur:
                for rcpt in envelope.rcpt_tos:
                    cur.execute(
                        "SELECT id FROM inboxes WHERE address=%s", (rcpt.lower(),)
                    )
                    row = cur.fetchone()
                    if not row:
                        continue  # unknown address -> silently drop
                    cur.execute(
                        "INSERT INTO emails "
                        "(inbox_id, from_addr, to_addr, subject, body_text, body_html) "
                        "VALUES (%s, %s, %s, %s, %s, %s)",
                        (row["id"], from_addr, rcpt, subject, text, html),
                    )
                    print(f"[smtp] stored mail for {rcpt}: {subject!r}")
        finally:
            conn.close()
        return "250 Message accepted for delivery"


def main():
    db.init_db()
    controller = Controller(Handler(), hostname=config.SMTP_HOST,
                            port=config.SMTP_PORT)
    controller.start()
    print(f"[smtp] listening on {config.SMTP_HOST}:{config.SMTP_PORT} "
          f"for domain '{config.MAIL_DOMAIN}'")
    try:
        asyncio.get_event_loop().run_forever()
    except KeyboardInterrupt:
        controller.stop()


if __name__ == "__main__":
    main()
