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()