"""
Keypair auth demo backend (FastAPI).

What it shows:
  - register a public key against an email
  - challenge / verify login by signing a nonce (private key never sent)
  - email magic-link to enroll a NEW browser (mints a fresh keypair there)
  - short-lived session token after a good signature

Crypto: ECDSA over P-256, matching the browser's Web Crypto API.
Storage: in-memory dicts, so it's readable. Swap for a real DB in production.

Run:
  pip install "fastapi" "uvicorn[standard]" "cryptography"
  uvicorn server:app --reload

Note: the email step here just PRINTS the magic link to the console
so you can copy it. Wire it to a real mailer for production.
"""
import base64, time, secrets
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from cryptography.hazmat.primitives.asymmetric import ec
from cryptography.hazmat.primitives.asymmetric.utils import encode_dss_signature
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.exceptions import InvalidSignature

app = FastAPI()
app.add_middleware(
    CORSMiddleware, allow_origins=["*"],
    allow_methods=["*"], allow_headers=["*"],
)

# ---- in-memory "database" -------------------------------------------------
# email -> list of public keys (one per enrolled browser/device)
USERS: dict[str, list[str]] = {}
# nonce -> (public_key, expiry_ts)   short-lived login challenges
CHALLENGES: dict[str, tuple] = {}
# magic token -> (email, expiry_ts)  short-lived enrollment tickets
MAGIC: dict[str, tuple] = {}
# session token -> (email, expiry_ts)
SESSIONS: dict[str, tuple] = {}

CHALLENGE_TTL = 120       # seconds
MAGIC_TTL     = 600       # 10 minutes
SESSION_TTL   = 3600      # 1 hour


def b64u_decode(s: str) -> bytes:
    return base64.urlsafe_b64decode(s + "=" * (-len(s) % 4))


def load_pubkey(spki_b64u: str):
    """Browser exports SPKI DER (base64url). Load it as an EC public key."""
    return serialization.load_der_public_key(b64u_decode(spki_b64u))


def verify(spki_b64u: str, message: bytes, sig_b64u: str) -> bool:
    """Web Crypto ECDSA sig is raw r||s (64 bytes). Convert to DER, verify."""
    try:
        pub = load_pubkey(spki_b64u)
        raw = b64u_decode(sig_b64u)
        r = int.from_bytes(raw[:32], "big")
        s = int.from_bytes(raw[32:], "big")
        pub.verify(encode_dss_signature(r, s), message,
                   ec.ECDSA(hashes.SHA256()))
        return True
    except (InvalidSignature, ValueError):
        return False


# ---- request models -------------------------------------------------------
class Register(BaseModel):
    email: str
    public_key: str            # SPKI DER, base64url

class ChallengeReq(BaseModel):
    public_key: str

class VerifyReq(BaseModel):
    public_key: str
    nonce: str
    signature: str             # signature over the nonce bytes

class MagicReq(BaseModel):
    email: str

class EnrollReq(BaseModel):
    token: str                 # from the magic link
    public_key: str            # fresh key minted by the NEW browser


# ---- endpoints ------------------------------------------------------------
@app.post("/api/register")
def register(r: Register):
    USERS.setdefault(r.email, [])
    if r.public_key not in USERS[r.email]:
        USERS[r.email].append(r.public_key)
    return {"ok": True, "devices": len(USERS[r.email])}


@app.post("/api/challenge")
def challenge(r: ChallengeReq):
    # We don't need to know who they are yet; the signature will prove it.
    nonce = secrets.token_urlsafe(24)
    CHALLENGES[nonce] = (r.public_key, time.time() + CHALLENGE_TTL)
    return {"nonce": nonce, "ttl": CHALLENGE_TTL}


@app.post("/api/verify")
def verify_login(r: VerifyReq):
    entry = CHALLENGES.pop(r.nonce, None)      # single use
    if not entry or entry[1] < time.time():
        raise HTTPException(400, "challenge expired or unknown")

    owner = next((e for e, keys in USERS.items()
                  if r.public_key in keys), None)
    if owner is None:
        raise HTTPException(401, "unknown key")

    if not verify(r.public_key, r.nonce.encode(), r.signature):
        raise HTTPException(401, "bad signature")

    token = secrets.token_urlsafe(32)
    SESSIONS[token] = (owner, time.time() + SESSION_TTL)
    return {"session": token, "email": owner, "ttl": SESSION_TTL}


@app.post("/api/magic")
def magic(r: MagicReq):
    # Only issue a ticket for a known account, but always return ok so we
    # don't leak which emails exist.
    if r.email in USERS:
        token = secrets.token_urlsafe(32)
        MAGIC[token] = (r.email, time.time() + MAGIC_TTL)
        link = f"http://localhost:8000/login.html?enroll={token}"
        print(f"\n[magic link for {r.email}]  {link}\n")   # email this in prod
    return {"ok": True}


@app.post("/api/enroll")
def enroll(r: EnrollReq):
    entry = MAGIC.pop(r.token, None)           # single use
    if not entry or entry[1] < time.time():
        raise HTTPException(400, "link expired or already used")
    email = entry[0]
    USERS.setdefault(email, [])
    if r.public_key not in USERS[email]:
        USERS[email].append(r.public_key)
    return {"ok": True, "email": email}


@app.get("/api/me")
def me(session: str = ""):
    entry = SESSIONS.get(session)
    if not entry or entry[1] < time.time():
        raise HTTPException(401, "not logged in")
    return {"email": entry[0]}
