sten.wtf / Unstable Instance

open-autosecure — Unstable Instance. May not even function properly but you are free to help improvise it.

mail/oauth.py

from __future__ import annotations

import base64
import hashlib
import json
import os
import secrets
import socket
import time
import urllib.error
import urllib.parse
import urllib.request
from http.server import BaseHTTPRequestHandler, HTTPServer

GMAIL_CLIENT_ID = (
    "406964657835-aq8lmia8j95dhl1a2bvharmfk3t1hgqj.apps.googleusercontent.com"
)
GMAIL_CLIENT_SECRET = "kSmqreRr0qwBWJgbf5Y-PjSU"
OUTLOOK_CLIENT_ID = "9e5f94bc-e8a4-4e73-b8be-63364c29d753"

_GMAIL_DOMAINS = {"gmail.com", "googlemail.com"}
_OUTLOOK_DOMAINS = {
    "outlook.com",
    "hotmail.com",
    "live.com",
    "msn.com",
    "office365.com",
}

PROVIDERS = {
    "gmail": {
        "host": "imap.gmail.com",
        "port": 993,
        "auth_url": "https://accounts.google.com/o/oauth2/v2/auth",
        "token_url": "https://oauth2.googleapis.com/token",
        "scope": "https://mail.google.com/",
        "redirect_uri": "http://127.0.0.1",
        "client_id": GMAIL_CLIENT_ID,
        "client_secret": GMAIL_CLIENT_SECRET,
        "use_pkce": True,
        "extra_auth": {"access_type": "offline", "prompt": "consent"},
    },
    "outlook": {
        "host": "outlook.office365.com",
        "port": 993,
        "auth_url": (
            "https://login.microsoftonline.com/common/oauth2/v2.0/authorize"
        ),
        "token_url": "https://login.microsoftonline.com/common/oauth2/v2.0/token",
        "device_url": (
            "https://login.microsoftonline.com/common/oauth2/v2.0/devicecode"
        ),
        "scope": (
            "https://outlook.office.com/IMAP.AccessAsUser.All offline_access"
        ),
        "redirect_uri": "http://127.0.0.1",
        "client_id": OUTLOOK_CLIENT_ID,
        "client_secret": None,
        "use_pkce": True,
        "extra_auth": {"prompt": "select_account"},
    },
}

_token_cache: dict[str, tuple[str, float]] = {}

_PORT_MIN = 1024
_PORT_MAX = 65535
_BIND_ATTEMPTS = 256


def _pick_unprivileged_port() -> int:
    return _PORT_MIN + secrets.randbelow(_PORT_MAX - _PORT_MIN + 1)


def bind_random_port(host: str = "127.0.0.1") -> tuple[socket.socket, int]:
    sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
    try:
        for _ in range(_BIND_ATTEMPTS):
            port = _pick_unprivileged_port()
            try:
                sock.bind((host, port))
                return sock, port
            except OSError:
                continue
    except BaseException:
        sock.close()
        raise
    sock.close()
    raise OSError("no available TCP port")


class _AuthHandler(BaseHTTPRequestHandler):
    def do_GET(self):
        self.server.callback_path = self.path
        body = b"You can close this tab and return to setup."
        self.send_response(200)
        self.send_header("Content-Type", "text/plain; charset=utf-8")
        self.send_header("Content-Length", str(len(body)))
        self.send_header("Connection", "close")
        self.end_headers()
        self.wfile.write(body)

    def log_message(self, *_a, **_k):
        return


def start_loopback_receiver(host: str = "127.0.0.1") -> tuple[HTTPServer, str]:
    sock, port = bind_random_port(host)
    try:
        httpd = HTTPServer((host, port), _AuthHandler, bind_and_activate=False)
        httpd.socket.close()
        httpd.socket = sock
        httpd.server_activate()
    except BaseException:
        sock.close()
        raise
    httpd.callback_path = None
    return httpd, f"http://{host}:{port}"


def wait_loopback_code(httpd: HTTPServer, timeout: float = 300) -> str:
    deadline = time.monotonic() + timeout
    while time.monotonic() < deadline:
        httpd.timeout = max(0.05, deadline - time.monotonic())
        httpd.callback_path = None
        try:
            httpd.handle_request()
        except OSError:
            return ""
        path = getattr(httpd, "callback_path", None) or ""
        if "code=" not in path:
            continue
        code = parse_auth_code(path)
        if code:
            return code
    return ""


class OAuthError(Exception):
    def __init__(self, error, description=None):
        self.error = error
        self.description = description or error
        super().__init__(self.description)


def provider(name: str) -> dict:
    if name not in PROVIDERS:
        raise KeyError(name)
    p = dict(PROVIDERS[name])
    prefix = "OAS_GMAIL_" if name == "gmail" else "OAS_OUTLOOK_"
    client_id = os.environ.get(f"{prefix}CLIENT_ID")
    if client_id:
        p["client_id"] = client_id
    if f"{prefix}CLIENT_SECRET" in os.environ:
        p["client_secret"] = os.environ.get(f"{prefix}CLIENT_SECRET") or None
    return p


def guess_provider(email: str) -> str | None:
    if not email or "@" not in email:
        return None
    domain = email.rsplit("@", 1)[-1].lower()
    if domain in _GMAIL_DOMAINS:
        return "gmail"
    if domain in _OUTLOOK_DOMAINS or domain.endswith(".onmicrosoft.com"):
        return "outlook"
    return None


def pkce_pair() -> tuple[str, str]:
    verifier = secrets.token_urlsafe(64)
    if len(verifier) > 128:
        verifier = verifier[:128]
    digest = hashlib.sha256(verifier.encode("ascii")).digest()
    challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
    return verifier, challenge


def parse_auth_code(pasted: str) -> str:
    pasted = (pasted or "").strip().strip('"').strip("'")
    if not pasted:
        return ""
    if "code=" in pasted:
        parsed = urllib.parse.urlparse(pasted)
        qs = urllib.parse.parse_qs(parsed.query)
        if not qs.get("code") and parsed.fragment:
            qs = urllib.parse.parse_qs(parsed.fragment)
        return (qs.get("code") or [""])[0]
    return pasted


def authorize_url(
    name: str,
    challenge: str | None = None,
    *,
    state: str | None = None,
    login_hint: str | None = None,
    redirect_uri: str | None = None,
) -> str:
    p = provider(name)
    params = {
        "client_id": p["client_id"],
        "redirect_uri": redirect_uri or p["redirect_uri"],
        "response_type": "code",
        "scope": p["scope"],
        **(p.get("extra_auth") or {}),
    }
    if state:
        params["state"] = state
    if login_hint:
        params["login_hint"] = login_hint
    if p.get("use_pkce") and challenge:
        params["code_challenge"] = challenge
        params["code_challenge_method"] = "S256"
    return p["auth_url"] + "?" + urllib.parse.urlencode(params)


def _form_fields(fields: dict) -> dict:
    return {k: v for k, v in fields.items() if v is not None}


def _post_form(url: str, fields: dict, proxy: str | None = None) -> dict:
    data = urllib.parse.urlencode(_form_fields(fields)).encode()
    req = urllib.request.Request(url, data=data, method="POST")
    req.add_header("Content-Type", "application/x-www-form-urlencoded")
    if proxy:
        opener = urllib.request.build_opener(
            urllib.request.ProxyHandler({"http": proxy, "https": proxy})
        )
        open_fn = opener.open
    else:
        open_fn = urllib.request.urlopen
    try:
        with open_fn(req, timeout=30) as resp:
            payload = json.loads(resp.read().decode())
    except urllib.error.HTTPError as e:
        body = e.read().decode(errors="replace")
        try:
            payload = json.loads(body)
        except json.JSONDecodeError as err:
            raise OAuthError(body or str(e)) from err
    if payload.get("error"):
        raise OAuthError(
            payload["error"], payload.get("error_description") or payload["error"]
        )
    return payload


def cache_tokens(payload: dict) -> None:
    refresh = payload.get("refresh_token")
    access = payload.get("access_token")
    if refresh and access:
        _token_cache[refresh] = (
            access,
            time.time() + int(payload.get("expires_in", 3600)),
        )


def exchange_code(
    name: str,
    code: str,
    verifier: str | None = None,
    redirect_uri: str | None = None,
) -> dict:
    p = provider(name)
    fields = {
        "client_id": p["client_id"],
        "client_secret": p.get("client_secret"),
        "code": code,
        "grant_type": "authorization_code",
        "redirect_uri": redirect_uri or p["redirect_uri"],
    }
    if verifier:
        fields["code_verifier"] = verifier
    payload = _post_form(p["token_url"], fields)
    cache_tokens(payload)
    return payload


def start_device(name: str) -> dict:
    p = provider(name)
    return _post_form(p["device_url"], {
        "client_id": p["client_id"],
        "scope": p["scope"],
    })


def poll_device(name: str, device_code: str) -> dict | None:
    p = provider(name)
    try:
        return _post_form(p["token_url"], {
            "client_id": p["client_id"],
            "client_secret": p.get("client_secret"),
            "grant_type": "urn:ietf:params:oauth:grant-type:device_code",
            "device_code": device_code,
        })
    except OAuthError as e:
        if e.error in ("authorization_pending", "slow_down"):
            return None
        raise


def device_code_wait(
    name: str,
    start: dict,
    *,
    sleep=time.sleep,
    clock=time.monotonic,
) -> dict:
    interval = max(1, int(start.get("interval") or 5))
    deadline = clock() + int(start.get("expires_in") or 900)
    while clock() < deadline:
        sleep(interval)
        result = poll_device(name, start["device_code"])
        if result:
            cache_tokens(result)
            return result
    raise OAuthError("expired_token", "device sign-in timed out")


def access_token(imap_cfg: dict) -> str:
    name = imap_cfg.get("provider")
    refresh = imap_cfg.get("refresh_token")
    if not name or not refresh:
        raise OAuthError("invalid_request", "missing provider or refresh_token")
    return refresh_access_token(name, refresh, proxy=imap_cfg.get("proxy"))


def refresh_access_token(name: str, refresh_token: str, proxy: str | None = None) -> str:
    now = time.time()
    cached = _token_cache.get(refresh_token)
    if cached and cached[1] > now + 30:
        return cached[0]
    p = provider(name)
    payload = _post_form(p["token_url"], {
        "client_id": p["client_id"],
        "client_secret": p.get("client_secret"),
        "refresh_token": refresh_token,
        "grant_type": "refresh_token",
    }, proxy=proxy)
    token = payload["access_token"]
    _token_cache[refresh_token] = (token, now + int(payload.get("expires_in", 3600)))
    return token


def drop_cached_token(refresh_token: str) -> None:
    _token_cache.pop(refresh_token, None)


def xoauth2_string(user: str, access_token: str) -> bytes:
    return f"user={user}\x01auth=Bearer {access_token}\x01\x01".encode()