Files
bookmarks-api/host/firefox_agent_bridge_host.py
T

263 lines
9.4 KiB
Python

#!/usr/bin/env python3
"""Native messaging host: 127.0.0.1 HTTP <-> Firefox extension stdio."""
from __future__ import annotations
import hashlib
import json
import os
import struct
import sys
import threading
import uuid
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Any
from urllib.parse import parse_qs, urlparse
if sys.platform == "win32":
import msvcrt
msvcrt.setmode(sys.stdin.fileno(), os.O_BINARY)
msvcrt.setmode(sys.stdout.fileno(), os.O_BINARY)
HOST = "127.0.0.1"
PORT = int(os.environ.get("FAB_PORT", "17634"))
MAX_BODY = 256 * 1024
ALLOWED_HOSTS = {f"127.0.0.1:{PORT}", f"localhost:{PORT}"}
STATE_DIR = Path(os.environ.get("LOCALAPPDATA", str(Path.home()))) / "firefox-agent-bridge"
STDIN_LOCK = threading.Lock()
PENDING: dict[str, tuple[threading.Event, dict[str, Any]]] = {}
PENDING_LOCK = threading.Lock()
def log(msg: str) -> None:
sys.stderr.write(msg + "\n")
sys.stderr.flush()
def send_to_extension(payload: dict[str, Any]) -> None:
raw = json.dumps(payload, separators=(",", ":")).encode("utf-8")
with STDIN_LOCK:
sys.stdout.buffer.write(struct.pack("@I", len(raw)))
sys.stdout.buffer.write(raw)
sys.stdout.buffer.flush()
def read_from_extension() -> dict[str, Any] | None:
header = sys.stdin.buffer.read(4)
if len(header) < 4:
return None
(length,) = struct.unpack("@I", header)
body = sys.stdin.buffer.read(length)
if len(body) < length:
return None
return json.loads(body.decode("utf-8"))
def _header(headers: dict[str, str], name: str) -> str:
want = name.lower()
for key, value in headers.items():
if key.lower() == want:
return (value or "").strip()
return ""
def verify_token_with_extension(token: str, timeout: float = 8.0) -> dict[str, Any]:
if not token.startswith("fab_"):
raise PermissionError("unrecognized token")
digest = hashlib.sha256(token.encode("ascii")).hexdigest()
req_id = str(uuid.uuid4())
event = threading.Event()
slot: dict[str, Any] = {}
with PENDING_LOCK:
PENDING[req_id] = (event, slot)
send_to_extension({"type": "auth", "id": req_id, "token_hash": digest})
if not event.wait(timeout):
with PENDING_LOCK:
PENDING.pop(req_id, None)
raise PermissionError("extension did not verify token")
if not slot.get("ok"):
raise PermissionError(slot.get("error") or "unknown token")
return slot.get("client") or {}
def call_extension(method: str, args: list[Any] | None = None, timeout: float = 15.0) -> dict[str, Any]:
req_id = str(uuid.uuid4())
event = threading.Event()
slot: dict[str, Any] = {}
with PENDING_LOCK:
PENDING[req_id] = (event, slot)
send_to_extension({"id": req_id, "method": method, "args": args or []})
if not event.wait(timeout):
with PENDING_LOCK:
PENDING.pop(req_id, None)
raise TimeoutError(f"extension did not answer {method}")
if not slot.get("ok"):
raise RuntimeError(slot.get("error") or "extension error")
return slot.get("result")
class Handler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def log_message(self, fmt: str, *args: Any) -> None:
log("%s - %s" % (self.address_string(), fmt % args))
def _send(self, code: int, payload: Any) -> None:
body = json.dumps(payload).encode("utf-8")
self.send_response(code)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def _gate_ok(self) -> bool:
"""Block browser-origin calls and non-loopback Host headers."""
origin = self.headers.get("Origin")
if origin:
self._send(403, {"ok": False, "error": "browser origin not allowed"})
return False
host = (self.headers.get("Host") or "").split("%")[0].lower()
if host not in ALLOWED_HOSTS:
self._send(403, {"ok": False, "error": "host not allowed"})
return False
return True
def _read_raw(self) -> bytes:
length = int(self.headers.get("Content-Length") or "0")
if length > MAX_BODY:
raise ValueError("request body too large")
if length <= 0:
return b""
return self.rfile.read(length)
def _authorize(self, http_method: str, path: str, raw_body: bytes) -> bool:
if not self._gate_ok():
return False
try:
auth = _header({k: v for k, v in self.headers.items()}, "Authorization")
if not auth.lower().startswith("bearer "):
raise PermissionError("Bearer token required")
verify_token_with_extension(auth[7:].strip())
except PermissionError as exc:
self._send(401, {"ok": False, "error": str(exc)})
return False
return True
def _parse_json(self, raw: bytes) -> dict[str, Any]:
if not raw:
return {}
data = json.loads(raw.decode("utf-8"))
if not isinstance(data, dict):
raise ValueError("JSON object required")
return data
def do_GET(self) -> None:
parsed = urlparse(self.path)
if parsed.path == "/health":
if not self._gate_ok():
return
self._send(200, {"ok": True, "service": "bookmarks-api", "port": PORT})
return
if not self._authorize("GET", parsed.path, b""):
return
qs = parse_qs(parsed.query)
try:
if parsed.path in ("/v1/ready", "/v1/methods"):
method = "meta.ping" if parsed.path == "/v1/ready" else "meta.methods"
result = call_extension(method)
self._send(200, {"ok": True, "result": result})
return
if parsed.path == "/v1/bookmarks/tree":
result = call_extension("bookmarks.getTree")
self._send(200, {"ok": True, "result": result})
return
if parsed.path.startswith("/v1/bookmarks/") and parsed.path != "/v1/bookmarks/":
node_id = parsed.path[len("/v1/bookmarks/") :].strip("/")
if "children" in qs:
result = call_extension("bookmarks.getChildren", [node_id])
else:
result = call_extension("bookmarks.get", [node_id])
self._send(200, {"ok": True, "result": result})
return
self._send(404, {"ok": False, "error": "not found"})
except Exception as exc:
self._send(502, {"ok": False, "error": str(exc)})
def do_POST(self) -> None:
parsed = urlparse(self.path)
try:
raw = self._read_raw()
except ValueError as exc:
self._send(400, {"ok": False, "error": str(exc)})
return
if not self._authorize("POST", parsed.path, raw):
return
try:
body = self._parse_json(raw)
except ValueError as exc:
self._send(400, {"ok": False, "error": str(exc)})
return
try:
if parsed.path == "/v1/call":
method = body.get("method")
args = body.get("args") or []
if not method:
self._send(400, {"ok": False, "error": "method required"})
return
result = call_extension(str(method), list(args) if not isinstance(args, list) else args)
self._send(200, {"ok": True, "result": result})
return
rest = {
"/v1/bookmarks/search": ("bookmarks.search", [body.get("query", body)]),
"/v1/bookmarks/create": ("bookmarks.create", [body]),
"/v1/bookmarks/update": ("bookmarks.update", [body.get("id"), body.get("changes") or {}]),
"/v1/bookmarks/move": ("bookmarks.move", [body.get("id"), body.get("destination") or {}]),
"/v1/bookmarks/remove": ("bookmarks.remove", [body.get("id")]),
"/v1/bookmarks/remove-tree": ("bookmarks.removeTree", [body.get("id")]),
}
if parsed.path in rest:
method, args = rest[parsed.path]
result = call_extension(method, args)
self._send(200, {"ok": True, "result": result})
return
self._send(404, {"ok": False, "error": "not found"})
except Exception as exc:
self._send(502, {"ok": False, "error": str(exc)})
def stdin_loop() -> None:
while True:
try:
message = read_from_extension()
except Exception as exc:
log(f"stdin read failed: {exc}")
break
if message is None:
break
req_id = str(message.get("id") or "")
with PENDING_LOCK:
pending = PENDING.pop(req_id, None)
if not pending:
continue
event, slot = pending
slot.update(message)
event.set()
os._exit(0)
def main() -> int:
STATE_DIR.mkdir(parents=True, exist_ok=True)
threading.Thread(target=stdin_loop, name="fab-stdin", daemon=True).start()
server = ThreadingHTTPServer((HOST, PORT), Handler)
log(f"listening on http://{HOST}:{PORT}")
try:
server.serve_forever()
except KeyboardInterrupt:
return 0
return 0
if __name__ == "__main__":
raise SystemExit(main())