#!/usr/bin/env python3
"""Sign in with GitHub or Google and save a sandbox OAuth token privately."""

import argparse
import base64
import hashlib
import http.server
import json
import os
from pathlib import Path
import secrets
import tempfile
import time
import urllib.error
import urllib.parse
import urllib.request
import webbrowser


def challenge(verifier):
    return base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()


def canonical_resource(issuer, resource):
    parsed = urllib.parse.urlsplit(issuer)
    if not parsed.hostname or parsed.username or parsed.password or parsed.query or parsed.fragment or parsed.path not in ("", "/"):
        raise ValueError("Issuer must be an origin without credentials, path, query, or fragment.")
    if parsed.scheme != "https" and not (parsed.scheme == "http" and parsed.hostname in ("localhost", "127.0.0.1")):
        raise ValueError("Issuer requires HTTPS (HTTP is allowed only on localhost or 127.0.0.1).")
    origin = issuer.rstrip("/")
    if resource in ("/mcp", "/vfs"):
        return origin, origin + resource
    if resource not in (origin + "/mcp", origin + "/vfs"):
        raise ValueError("Resource must be /mcp or /vfs on the selected issuer.")
    return origin, resource


def request_json(url, payload=None, form=False):
    headers = {"Accept": "application/json"}
    data = None
    if payload is not None:
        data = (urllib.parse.urlencode(payload) if form else json.dumps(payload)).encode()
        headers["Content-Type"] = "application/x-www-form-urlencoded" if form else "application/json"
    # Do not follow redirects: a token exchange must never forward credentials.
    class NoRedirect(urllib.request.HTTPRedirectHandler):
        def redirect_request(self, *args, **kwargs):
            return None
    try:
        with urllib.request.build_opener(NoRedirect).open(urllib.request.Request(url, data=data, headers=headers), timeout=20) as response:
            result = json.load(response)
    except urllib.error.HTTPError as error:
        raise ValueError(f"OAuth request failed (HTTP {error.code}); check the issuer and retry authorization.") from None
    except (OSError, ValueError):
        raise ValueError("Cannot read OAuth server response; check connectivity and issuer configuration.") from None
    if not isinstance(result, dict):
        raise ValueError("OAuth server returned an invalid JSON object.")
    return result


def save_tokens(path, tokens):
    destination = Path(path).expanduser()
    destination.parent.mkdir(parents=True, exist_ok=True, mode=0o700)
    fd, temporary = tempfile.mkstemp(prefix=".oauth-", dir=destination.parent)
    try:
        os.fchmod(fd, 0o600)
        with os.fdopen(fd, "w") as output:
            json.dump(tokens, output, indent=2)
            output.write("\n")
        os.replace(temporary, destination)
    finally:
        if os.path.exists(temporary):
            os.unlink(temporary)


def callback_handler(expected_state, result, expected_issuer=None):
    class Callback(http.server.BaseHTTPRequestHandler):
        def do_GET(self):
            parsed = urllib.parse.urlsplit(self.path)
            values = urllib.parse.parse_qs(parsed.query)
            valid_state = len(values.get("state", [])) == 1 and secrets.compare_digest(values["state"][0], expected_state)
            valid_issuer = expected_issuer is None or values.get("iss") == [expected_issuer]
            if parsed.path != "/callback" or not valid_state or not valid_issuer:
                self.send_response(400)
                message = "Invalid OAuth callback. Return to the authorization browser window."
            elif "error" in values:
                result["error"] = "Authorization was declined. Run the helper again to retry."
                self.send_response(400)
                message = result["error"]
            elif len(values.get("code", [])) != 1:
                self.send_response(400)
                message = "Authorization callback did not include a code."
            else:
                result["code"] = values["code"][0]
                self.send_response(200)
                message = "Authorization received. You may close this window."
            self.send_header("Content-Type", "text/plain; charset=utf-8")
            self.send_header("Cache-Control", "no-store")
            self.end_headers()
            self.wfile.write(message.encode())

        def log_message(self, *args):
            pass  # Callback URLs contain authorization codes.
    return Callback


def obtain_tokens(args):
    issuer, resource = canonical_resource(args.issuer, "/mcp")
    metadata = request_json(issuer + "/.well-known/oauth-authorization-server")
    if metadata.get("issuer") != issuer:
        raise ValueError("Authorization metadata issuer does not match the requested issuer.")
    endpoints = {}
    for name in ("registration_endpoint", "authorization_endpoint", "token_endpoint"):
        endpoint = metadata.get(name, "")
        parsed = urllib.parse.urlsplit(endpoint) if isinstance(endpoint, str) else None
        origin = urllib.parse.urlsplit(issuer)
        if not parsed or (parsed.scheme, parsed.netloc) != (origin.scheme, origin.netloc) or parsed.query or parsed.fragment:
            raise ValueError(f"Invalid or cross-origin {name} in authorization metadata.")
        endpoints[name] = endpoint
    if args.refresh:
        refresh_tokens(args, issuer, resource, endpoints["token_endpoint"])
        return
    verifier, state = secrets.token_urlsafe(48), secrets.token_urlsafe(32)
    result = {}
    with http.server.HTTPServer(("127.0.0.1", 0), callback_handler(state, result, issuer)) as listener:
        listener.timeout = 1
        redirect_uri = f"http://127.0.0.1:{listener.server_port}/callback"
        client = request_json(endpoints["registration_endpoint"], {
            "client_name": "elixirmcp.dev local demo setup",
            "redirect_uris": [redirect_uri],
            "grant_types": ["authorization_code", "refresh_token"],
            "response_types": ["code"], "token_endpoint_auth_method": "none",
        })
        client_id = client.get("client_id")
        if not isinstance(client_id, str) or not client_id:
            raise ValueError("Registration did not return a client ID.")
        query = urllib.parse.urlencode({
            "client_id": client_id, "redirect_uri": redirect_uri,
            "response_type": "code", "scope": "demo:read demo:write", "resource": resource,
            "state": state, "code_challenge": challenge(verifier), "code_challenge_method": "S256",
        })
        authorization_url = endpoints["authorization_endpoint"] + "?" + query
        print("Approve access to a temporary read/write demo sandbox in your browser. Sign in with GitHub or Google; this app never collects your password.")
        print("If the browser does not open, visit:\n" + authorization_url)
        webbrowser.open(authorization_url)
        deadline = time.monotonic() + args.timeout
        while not result and time.monotonic() < deadline:
            listener.handle_request()
        if not result:
            raise ValueError("Authorization timed out. Run the helper again and approve in the browser.")
        if "error" in result:
            raise ValueError(result["error"])
    tokens = request_json(endpoints["token_endpoint"], {
        "grant_type": "authorization_code", "client_id": client_id,
        "redirect_uri": redirect_uri, "code": result["code"],
        "code_verifier": verifier, "resource": resource,
    }, form=True)
    if not isinstance(tokens.get("access_token"), str) or not tokens["access_token"]:
        raise ValueError("Token exchange did not return an access token.")
    if str(tokens.get("token_type", "")).lower() != "bearer":
        raise ValueError("Token exchange returned an unsupported token type.")
    tokens.update(issuer=issuer, resource=resource, client_id=client_id,
                  redirect_uri=redirect_uri, obtained_at=int(time.time()))
    save_tokens(args.output, tokens)
    print(f"Token saved privately to {Path(args.output).expanduser()}. Treat this file as a credential.")


def refresh_tokens(args, issuer, resource, endpoint):
    try:
        previous = json.loads(Path(args.output).expanduser().read_text())
    except (OSError, ValueError):
        raise ValueError("Cannot read token file; run without --refresh to authorize again.") from None
    if not isinstance(previous, dict):
        raise ValueError("Invalid token file; authorize again without --refresh.")
    if previous.get("issuer") != issuer or previous.get("resource") != resource:
        raise ValueError("Token file issuer/resource mismatch; select the same resource used to authorize.")
    if not previous.get("refresh_token") or not previous.get("client_id"):
        raise ValueError("Token file has no refresh credential; authorize again without --refresh.")
    tokens = request_json(endpoint, {
        "grant_type": "refresh_token", "client_id": previous["client_id"],
        "refresh_token": previous["refresh_token"], "resource": resource,
    }, form=True)
    if not isinstance(tokens.get("access_token"), str) or not tokens["access_token"]:
        raise ValueError("Refresh did not return an access token; authorize again.")
    if str(tokens.get("token_type", "")).lower() != "bearer":
        raise ValueError("Refresh returned an unsupported token type.")
    previous.update(tokens, obtained_at=int(time.time()))
    save_tokens(args.output, previous)
    print("Token refreshed privately. Restart the VFS daemon or reconnect Postgres with the new token.")


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--issuer", default="https://mcp.elixirmcp.dev")
    parser.add_argument("--resource", choices=("/mcp", "/vfs"), default="/mcp")
    parser.add_argument("--output", required=True, help="Private JSON credential file (mode 0600).")
    parser.add_argument("--refresh", action="store_true", help="Rotate the refresh token in --output without opening a browser.")
    parser.add_argument("--timeout", type=int, default=180, help="Browser approval timeout in seconds.")
    args = parser.parse_args()
    if args.timeout <= 0:
        parser.error("--timeout must be positive")
    try:
        obtain_tokens(args)
    except (ValueError, OSError) as error:
        parser.exit(1, f"Setup failed: {error}\n")
    except KeyboardInterrupt:
        parser.exit(130, "Authorization cancelled.\n")


if __name__ == "__main__":
    main()
