#!/usr/bin/env -S uv run --script
# /// script
# requires-python = ">=3.11"
# dependencies = [
#     "boto3",
#     "cryptography",
#     "temporalio",
# ]
# ///
"""Temporal Codec Server for decrypting PostHog Fernet-encrypted payloads.

Runs a local HTTP server implementing the Temporal Remote Codec protocol,
allowing the Temporal Web UI to decode encrypted payloads in-browser.

Fetches TEMPORAL_SECRET_KEY from AWS Secrets Manager (via SSO), derives
the Fernet key, and keeps it in memory for the lifetime of the process.

Usage:
    bin/temporal-codec-server us  # Use US region key
    bin/temporal-codec-server eu  # Use EU region key

Then in the Temporal UI, set the codec server endpoint to http://localhost:8089
"""

import sys
import json
import base64
import subprocess
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import Any, Protocol
from urllib.parse import urlparse

from cryptography.fernet import Fernet, MultiFernet
from temporalio.api.common.v1 import Payload

REGIONS = {
    "us": {"aws_profile": "prod-us-secrets", "aws_region": "us-east-1"},
    "eu": {"aws_profile": "prod-eu-secrets", "aws_region": "eu-central-1"},
}

AWS_SECRET_NAME = "temporal-worker-shared-secrets"

LISTEN_PORT = 8089

ENCRYPTED_ENCODING = base64.b64encode(b"binary/encrypted").decode()  # YmluYXJ5L2VuY3J5cHRlZA==
ALLOWED_ORIGINS = {
    "https://cloud.temporal.io",
    "http://localhost:8233",
}


class _HasDecrypt(Protocol):
    def decrypt(self, data: bytes) -> bytes: ...


def get_fernet(region: str) -> MultiFernet:
    import boto3
    from botocore.exceptions import ClientError, NoCredentialsError, TokenRetrievalError

    cfg = REGIONS[region]
    profile, aws_region = cfg["aws_profile"], cfg["aws_region"]

    try:
        session = boto3.Session(profile_name=profile)
        client = session.client("secretsmanager", region_name=aws_region)
        response = client.get_secret_value(SecretId=AWS_SECRET_NAME)
    except (TokenRetrievalError, NoCredentialsError, ClientError) as e:
        if isinstance(e, ClientError) and e.response["Error"]["Code"] not in (
            "ExpiredToken",
            "ExpiredTokenException",
            "AccessDeniedException",
        ):
            raise
        print(f"AWS SSO session not active for profile '{profile}', logging in...")
        subprocess.run(["aws", "sso", "login", "--profile", profile], check=True)
        session = boto3.Session(profile_name=profile)
        client = session.client("secretsmanager", region_name=aws_region)
        response = client.get_secret_value(SecretId=AWS_SECRET_NAME)

    secret_data = json.loads(response["SecretString"])
    if "TEMPORAL_SECRET_KEY" not in secret_data:
        print(
            f"TEMPORAL_SECRET_KEY not found in {AWS_SECRET_NAME}. Available keys: {sorted(secret_data.keys())}",
            file=sys.stderr,
        )
        sys.exit(1)

    secret_key = Fernet(_prepare_key(_load_as_bytes(secret_data["TEMPORAL_SECRET_KEY"])))

    fallback_keys = (
        Fernet(_prepare_key(_load_as_bytes(secret)))
        for secret in _split_fallback_keys(secret_data.get("TEMPORAL_FALLBACK_SECRET_KEYS", []))
        if secret.strip()  # Filter out empty/whitespace strings
    )


    return MultiFernet([secret_key, *fallback_keys])


def _split_fallback_keys(raw: str | list[str]) -> list[str]:
    """Normalize the fallback keys to a list.

    The secret stores TEMPORAL_FALLBACK_SECRET_KEYS as a comma-separated string, which
    Django parses with get_list (settings/temporal.py). Iterating the raw string would
    yield single characters, so mirror get_list here; tolerate an already-parsed list too.
    """
    if isinstance(raw, list):
        return raw
    if not raw:
        return []
    return [item.strip() for item in raw.split(",")]


def decode_payload(payload: dict, fernet: _HasDecrypt) -> dict:
    """Decode a single Temporal payload if it's Fernet-encrypted."""
    metadata = payload.get("metadata", {})
    if metadata.get("encoding") != ENCRYPTED_ENCODING:
        return payload

    try:
        encrypted_data = base64.b64decode(payload["data"])
        decrypted = fernet.decrypt(encrypted_data)
        inner = Payload.FromString(decrypted)

        decoded_metadata = {key: base64.b64encode(value).decode() for key, value in inner.metadata.items()}
        return {
            "metadata": decoded_metadata,
            "data": base64.b64encode(inner.data).decode(),
        }
    except Exception as e:
        print(f"[codec-server] Decryption failed ({type(e).__name__}): {e}", file=sys.stderr)
        return payload


class CodecHandler(BaseHTTPRequestHandler):
    fernet: _HasDecrypt

    def do_OPTIONS(self) -> None:
        if not self._is_origin_allowed():
            self.send_error(403, "Origin not allowed")
            return
        self.send_response(200)
        self._send_cors_headers()
        self.end_headers()

    def do_POST(self) -> None:
        if not self._is_origin_allowed():
            self.send_error(403, "Origin not allowed")
            return
        path = urlparse(self.path).path
        if path == "/decode":
            self._handle_decode()
        elif path == "/encode":
            self._handle_passthrough()
        else:
            self.send_error(404)

    def _handle_decode(self) -> None:
        try:
            body = self._read_json_body()
        except ValueError as err:
            self.send_error(400, str(err))
            return
        decoded_payloads = [decode_payload(p, self.fernet) for p in body.get("payloads", [])]
        self._send_json_response({"payloads": decoded_payloads})

    def _handle_passthrough(self) -> None:
        try:
            body = self._read_json_body()
        except ValueError as err:
            self.send_error(400, str(err))
            return
        self._send_json_response(body)

    def _read_json_body(self) -> dict[str, Any]:
        content_length = int(self.headers.get("Content-Length", 0))
        if content_length == 0:
            return {}
        try:
            return json.loads(self.rfile.read(content_length))
        except json.JSONDecodeError as err:
            raise ValueError("Request body must be valid JSON") from err

    def _send_json_response(self, data: dict[str, Any]) -> None:
        response = json.dumps(data).encode()
        self.send_response(200)
        self._send_cors_headers()
        self.send_header("Content-Type", "application/json")
        self.end_headers()
        self.wfile.write(response)

    def _is_origin_allowed(self) -> bool:
        origin = self.headers.get("Origin")
        return origin is None or origin in ALLOWED_ORIGINS

    def _send_cors_headers(self) -> None:
        origin = self.headers.get("Origin")
        if origin in ALLOWED_ORIGINS:
            self.send_header("Access-Control-Allow-Origin", origin)
        self.send_header("Access-Control-Allow-Methods", "POST, OPTIONS")
        self.send_header("Access-Control-Allow-Headers", "Content-Type, x-namespace")

    def log_message(self, format: str, *args: object) -> None:
        sys.stderr.write(f"[codec-server] {args[0]}\n")


# region Lifted from posthog.temporal.common.codec


def _load_as_bytes(raw: str | bytes, /, check_length: bool = True) -> bytes:
    if isinstance(raw, bytes):
        loaded = raw

    else:
        prefix, sep, secret = raw.partition(":")

        if sep and prefix == "hex":
            loaded = bytes.fromhex(secret)

        elif sep and prefix == "base64-urlsafe":
            loaded = base64.urlsafe_b64decode(secret)

        elif sep and prefix == "base64":
            loaded = base64.b64decode(secret, validate=True)

        else:
            # Legacy format, kept for backwards compatibility
            # TODO: Remove this branch & raise after rotating secrets
            loaded = raw.encode()

    # TODO: Also, make the check exact after removing legacy format
    if check_length and len(loaded) < 32:
        raise ValueError(f"Expected at least 32 bytes, got '{len(loaded)}'")

    return loaded


def _prepare_key(key: bytes) -> bytes:
    """Prepare an encryption key by padding or truncating it, and encoding it.

    We require a URL-safe, base64 encoded, 32 byte key.
    """
    resized_key = _resize_key(key)
    encoded_key = base64.urlsafe_b64encode(resized_key)
    return encoded_key


def _resize_key(key: bytes, size: int = 32) -> bytes:
    """Resize key to size bytes.

    Adds padding if key is too short, otherwise truncates to first 32 bytes.
    """
    padding = b"\0" * max(size - len(key), 0)
    return (padding + key)[:size]


# endregion


def main() -> None:
    if len(sys.argv) < 2 or sys.argv[1] not in REGIONS:
        print(f"Usage: temporal-codec-server <{'|'.join(REGIONS)}>", file=sys.stderr)
        sys.exit(1)

    region = sys.argv[1]
    print(f"Fetching TEMPORAL_SECRET_KEY from AWS Secrets Manager ({region})...")
    fernet = get_fernet(region)
    CodecHandler.fernet = fernet

    server = HTTPServer(("localhost", LISTEN_PORT), CodecHandler)
    print(f"Temporal Codec Server listening on http://localhost:{LISTEN_PORT} (region: {region})")
    print("Set this as the codec endpoint in the Temporal UI.")

    try:
        server.serve_forever()
    except KeyboardInterrupt:
        print("\nShutting down.")
        server.shutdown()


if __name__ == "__main__":
    main()
