Skip to content
← Public packages

@kentcdodds/x

X API v2 helpers for tweets, search, legacy DMs, and encrypted X Chat via a Fly XDK sidecar.

sidecar/server.py

291 lines · 10.5 KB · Python
"""On-demand X Chat XDK sidecar.

Unlocks the existing Juicebox identity with the user's PIN (never generates a
new keypair) and decrypts/encrypts events for the owning Kody X package.

Auth: `Authorization: Bearer $SIDECAR_TOKEN`
PIN:  `X-Chat-Pin` header (never logged, never sent to api.x.com)
"""

from __future__ import annotations

import base64
import hashlib
import hmac
import json
import os
import threading
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Any

from chat_xdk import Chat

try:
    from chat_xdk import guesses_remaining as juicebox_guesses_remaining
except ImportError:
    def juicebox_guesses_remaining(exc: BaseException) -> int | None:
        text = str(exc)
        marker = 'guesses_remaining='
        index = text.find(marker)
        if index < 0:
            return None
        raw = text[index + len(marker) :].split()[0].rstrip(',;')
        try:
            return int(raw)
        except ValueError:
            return None

PORT = int(os.environ.get("PORT", "8080"))
SIDECAR_TOKEN = os.environ.get("SIDECAR_TOKEN", "")

_sessions: dict[str, Chat] = {}
_session_lock = threading.Lock()


def _jsonable(value: Any) -> Any:
    if isinstance(value, dict):
        return {str(key): _jsonable(item) for key, item in value.items()}
    if isinstance(value, (list, tuple)):
        return [_jsonable(item) for item in value]
    if isinstance(value, (bytes, bytearray)):
        return base64.b64encode(bytes(value)).decode("ascii")
    if hasattr(value, "model_dump"):
        return _jsonable(value.model_dump())
    if isinstance(value, (str, int, float, bool)) or value is None:
        return value
    return str(value)


def _digest(*parts: str) -> str:
    hasher = hashlib.sha256()
    for part in parts:
        hasher.update(part.encode("utf-8"))
        hasher.update(b"\0")
    return hasher.hexdigest()


def _config_json(juicebox_config: Any) -> str:
    if isinstance(juicebox_config, str) and juicebox_config.strip():
        return juicebox_config
    if isinstance(juicebox_config, dict):
        return json.dumps(juicebox_config)
    raise ValueError("juicebox_config is required")


def _as_string_list(value: Any) -> list[str]:
    if not isinstance(value, list):
        return []
    out: list[str] = []
    for item in value:
        if isinstance(item, str) and item:
            out.append(item)
        elif isinstance(item, dict):
            encoded = item.get("encoded_event") or item.get("encodedEvent")
            if isinstance(encoded, str) and encoded:
                out.append(encoded)
    return out


def _unlock(user_id: str, pin: str, juicebox_config: Any, signing_key_version: str) -> Chat:
    config_json = _config_json(juicebox_config)
    cache_key = _digest(user_id, signing_key_version, pin, config_json)
    with _session_lock:
        existing = _sessions.get(cache_key)
        if existing is not None and existing.is_unlocked():
            return existing
        chat = Chat(config_json)
        chat.unlock(pin)
        chat.set_identity(user_id, signing_key_version)
        chat.set_cache_keys(True)
        _sessions.clear()
        _sessions[cache_key] = chat
        return chat


def _prepare_chat(body: dict[str, Any], pin: str) -> Chat:
    user_id = str(body.get("user_id") or "").strip()
    if not user_id:
        raise ValueError("user_id is required")
    signing_key_version = str(body.get("signing_key_version") or "").strip()
    if not signing_key_version:
        raise ValueError("signing_key_version is required")
    chat = _unlock(user_id, pin, body.get("juicebox_config"), signing_key_version)
    signing_keys = body.get("signing_keys")
    if isinstance(signing_keys, list) and signing_keys:
        chat.set_signing_keys(signing_keys)
    if hasattr(chat, "set_reject_unverified"):
        chat.set_reject_unverified(body.get("reject_unverified") is not False)
    return chat


def _send_payload(payload: Any) -> dict[str, Any]:
    return {
        "message_id": getattr(payload, "message_id", None),
        "encoded_message_create_event": getattr(payload, "encrypted_content", None),
        "encoded_message_event_signature": getattr(payload, "encoded_event_signature", None),
        "conversation_key_version": getattr(payload, "conversation_key_version", None),
        "should_notify": getattr(payload, "should_notify", None),
    }


def _authorized(handler: BaseHTTPRequestHandler) -> bool:
    if not SIDECAR_TOKEN:
        return False
    header = handler.headers.get("Authorization", "")
    scheme, _, token = header.partition(" ")
    if scheme.lower() != "bearer" or not token:
        return False
    return hmac.compare_digest(token, SIDECAR_TOKEN)


class Handler(BaseHTTPRequestHandler):
    server_version = "kody-x-chat/1"

    def log_message(self, fmt: str, *args: Any) -> None:
        path = self.path.split("?", 1)[0]
        super().log_message("%s %s", path, fmt % args)

    def _write(self, status: int, payload: dict[str, Any]) -> None:
        body = json.dumps(payload).encode("utf-8")
        self.send_response(status)
        self.send_header("Content-Type", "application/json")
        self.send_header("Content-Length", str(len(body)))
        self.end_headers()
        self.wfile.write(body)

    def _read_json(self) -> dict[str, Any]:
        length = int(self.headers.get("Content-Length", "0") or "0")
        raw = self.rfile.read(length) if length else b"{}"
        parsed = json.loads(raw.decode("utf-8") or "{}")
        if not isinstance(parsed, dict):
            raise ValueError("JSON object body required")
        return parsed

    def do_GET(self) -> None:
        path = self.path.split("?", 1)[0]
        if path == "/health":
            self._write(
                200,
                {
                    "ok": True,
                    "service": "kody-x-chat",
                    "unlocked": any(chat.is_unlocked() for chat in _sessions.values()),
                },
            )
            return
        self._write(404, {"error": "not_found"})

    def do_POST(self) -> None:
        path = self.path.split("?", 1)[0]
        if not _authorized(self):
            self._write(401, {"error": "unauthorized"})
            return
        pin = (self.headers.get("X-Chat-Pin") or "").strip()
        if not pin:
            self._write(400, {"error": "missing_pin"})
            return
        try:
            body = self._read_json()
        except Exception:
            self._write(400, {"error": "invalid_json"})
            return
        try:
            if path == "/v1/decrypt-events":
                self._write(200, self._decrypt(body, pin))
                return
            if path == "/v1/encrypt-message":
                self._write(200, self._encrypt(body, pin))
                return
        except ValueError as error:
            self._write(400, {"error": "invalid_request", "message": str(error)})
            return
        except Exception as error:
            remaining = juicebox_guesses_remaining(error)
            payload: dict[str, Any] = {"error": "sidecar_failed", "message": str(error)}
            if remaining is not None:
                payload["guesses_remaining"] = remaining
            self._write(500, payload)
            return
        self._write(404, {"error": "not_found"})

    def _decrypt(self, body: dict[str, Any], pin: str) -> dict[str, Any]:
        chat = _prepare_chat(body, pin)
        encoded_events = _as_string_list(body.get("encoded_events"))
        if not encoded_events:
            raise ValueError("encoded_events is required")
        result = chat.decrypt_events(encoded_events)
        messages = []
        for item in result.get("messages") or []:
            event = item.get("event") if isinstance(item, dict) else None
            messages.append({"event": _jsonable(event or item)})
        errors = result.get("errors") or {}
        if errors and body.get("reject_unverified") is False:
            conversation_keys = {}
            try:
                extracted = chat.extract_conversation_keys(encoded_events)
                conversation_keys = extracted.get("keys") or {}
            except Exception:
                conversation_keys = {}
            if conversation_keys:
                retried_errors: dict[str, str] = {}
                recovered: dict[int, dict[str, Any]] = {}
                for index_str, message in errors.items():
                    try:
                        index = int(index_str)
                        event = chat.decrypt_event(
                            encoded_events[index], conversation_keys=conversation_keys
                        )
                        recovered[index] = {"event": _jsonable(event)}
                    except Exception as retry_error:
                        retried_errors[index_str] = str(retry_error)
                for index in sorted(recovered):
                    messages.append(recovered[index])
                errors = retried_errors
        return {
            "ok": True,
            "messages": messages,
            "error_count": len(errors) if isinstance(errors, dict) else 0,
            "errors": _jsonable(errors) if errors else {},
        }

    def _encrypt(self, body: dict[str, Any], pin: str) -> dict[str, Any]:
        chat = _prepare_chat(body, pin)
        warmup = _as_string_list(body.get("encoded_events"))
        if warmup:
            chat.decrypt_events(warmup)
        conversation_id = str(body.get("conversation_id") or "").strip()
        text = str(body.get("text") or "")
        if not conversation_id:
            raise ValueError("conversation_id is required")
        if not text:
            raise ValueError("text is required")
        try:
            payload = chat.encrypt_message(conversation_id, text)
        except Exception:
            keys = {}
            if warmup:
                try:
                    keys = chat.extract_conversation_keys(warmup).get("keys") or {}
                except Exception:
                    keys = {}
            if not keys:
                raise
            version = max(keys, key=lambda item: int(item) if item.isdigit() else -1)
            payload = chat.encrypt_message(
                conversation_id,
                text,
                conversation_key=keys[version],
                conversation_key_version=version,
            )
        return {"ok": True, "payload": _send_payload(payload)}


def main() -> None:
    if not SIDECAR_TOKEN:
        raise SystemExit("SIDECAR_TOKEN is required")
    server = ThreadingHTTPServer(("0.0.0.0", PORT), Handler)
    server.serve_forever()


if __name__ == "__main__":
    main()