#!/usr/bin/env python3
"""Dovecot checkpassword for primary + app passwords (argon2id).

Protocol: Dovecot writes "user\\0password\\0" to fd 3; we exit 0 on success,
1 on auth failure, 111 on temporary failure.

Environment (from limristem-mail install / systemd):
  LIMRISTEM_MAIL_DB_HOST, LIMRISTEM_MAIL_DB_PORT, LIMRISTEM_MAIL_DB_NAME
  LIMRISTEM_MAIL_DB_RO_USER, LIMRISTEM_MAIL_DB_RO_PASS
"""

from __future__ import annotations

import os
import sys

try:
    import pymysql
    from passlib.context import CryptContext
except ImportError:
    sys.exit(111)

pwd_context = CryptContext(
    schemes=["argon2", "bcrypt", "pbkdf2_sha512"],
    deprecated="auto",
    default="argon2",
    argon2__type="ID",
)


def read_credentials() -> tuple[str, str] | None:
    try:
        raw = os.read(3, 8192)
    except OSError:
        return None
    if not raw:
        return None
    parts = raw.split(b"\0")
    if len(parts) < 2:
        return None
    try:
        username = parts[0].decode("utf-8")
        password = parts[1].decode("utf-8")
    except UnicodeDecodeError:
        return None
    return username, password


def connect():
    return pymysql.connect(
        host=os.environ.get("LIMRISTEM_MAIL_DB_HOST", "127.0.0.1"),
        port=int(os.environ.get("LIMRISTEM_MAIL_DB_PORT", "3306")),
        user=os.environ.get("LIMRISTEM_MAIL_DB_RO_USER") or os.environ.get("LIMRISTEM_MAIL_DB_USER", "limristem-mail"),
        password=os.environ.get("LIMRISTEM_MAIL_DB_RO_PASS") or os.environ.get("LIMRISTEM_MAIL_DB_PASS", ""),
        database=os.environ.get("LIMRISTEM_MAIL_DB_NAME", "limristem-mail"),
        charset="utf8mb4",
        connect_timeout=5,
        read_timeout=5,
        write_timeout=5,
        cursorclass=pymysql.cursors.DictCursor,
    )


def main() -> int:
    creds = read_credentials()
    if not creds:
        return 1
    username, password = creds
    if not username or password is None:
        return 1

    try:
        conn = connect()
    except Exception:
        return 111

    try:
        with conn.cursor() as cur:
            cur.execute(
                """
                SELECT a.id AS account_id, a.password_hash, COALESCE(a.require_app_password, 0) AS require_app_password
                FROM accounts a
                JOIN domains d ON a.domain_id = d.id
                WHERE CONCAT(a.local_part, '@', d.name) = %s
                  AND a.is_active = 1 AND d.is_active = 1
                LIMIT 1
                """,
                (username,),
            )
            account = cur.fetchone()
            if not account:
                return 1

            candidates: list[str] = []
            if not int(account["require_app_password"] or 0):
                if account.get("password_hash"):
                    candidates.append(account["password_hash"])

            cur.execute(
                """
                SELECT password_hash
                FROM mailbox_app_passwords
                WHERE account_id = %s AND revoked_at IS NULL
                ORDER BY id ASC
                """,
                (account["account_id"],),
            )
            for row in cur.fetchall() or []:
                if row.get("password_hash"):
                    candidates.append(row["password_hash"])

            for password_hash in candidates:
                try:
                    if pwd_context.verify(password, password_hash):
                        # Optional: update last_used_at requires write grants; skip for RO user.
                        return 0
                except (ValueError, TypeError):
                    continue
            return 1
    except Exception:
        return 111
    finally:
        try:
            conn.close()
        except Exception:
            pass


if __name__ == "__main__":
    # After successful auth Dovecot execs the next process from argv; for checkpassword
    # the common pattern is: verify then os.execvp(sys.argv[1], sys.argv[1:])
    status = main()
    if status != 0:
        sys.exit(status)
    if len(sys.argv) > 1:
        try:
            os.execvp(sys.argv[1], sys.argv[1:])
        except OSError:
            sys.exit(111)
    sys.exit(0)
