Files
bookmarks-api/host/firefox_agent_bridge_host.py

323 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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 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"
DEFAULT_PORT = int(os.environ.get("FAB_PORT", "17634"))
PORT = DEFAULT_PORT
MAX_BODY = 256 * 1024
ALLOWED_HOSTS = {f"127.0.0.1:{PORT}", f"localhost:{PORT}"}
HTTP_SERVER: ThreadingHTTPServer | None = None
SERVER_LOCK = threading.Lock()
ALLOWED_METHODS = frozenset(
{
"meta.ping",
"meta.methods",
"bookmarks.getTree",
"bookmarks.getSubTree",
"bookmarks.getChildren",
"bookmarks.get",
"bookmarks.search",
"bookmarks.create",
"bookmarks.update",
"bookmarks.move",
"bookmarks.remove",
"bookmarks.removeTree",
"bookmarks.getRecent",
}
)
STDIN_LOCK = threading.Lock()
PENDING: dict[str, tuple[threading.Event, dict[str, Any]]] = {}
PENDING_LOCK = threading.Lock()
def set_port_globals(port: int) -> None:
global PORT, ALLOWED_HOSTS
PORT = port
ALLOWED_HOSTS = {f"127.0.0.1:{port}", f"localhost:{port}"}
def apply_listen_port(requested: Any) -> dict[str, Any]:
global HTTP_SERVER
try:
port = int(requested)
except (TypeError, ValueError):
return {"type": "config-result", "ok": False, "error": "invalid port"}
if port < 1024 or port > 65535:
return {
"type": "config-result",
"ok": False,
"port": port,
"error": "port must be 102465535",
}
with SERVER_LOCK:
if HTTP_SERVER is not None and PORT == port:
return {"type": "config-result", "ok": True, "port": port}
old = HTTP_SERVER
HTTP_SERVER = None
if old is not None:
old.shutdown()
old.server_close()
try:
set_port_globals(port)
HTTP_SERVER = ThreadingHTTPServer((HOST, port), Handler)
except OSError as exc:
return {"type": "config-result", "ok": False, "port": port, "error": str(exc)}
threading.Thread(
target=HTTP_SERVER.serve_forever, name="fab-http", daemon=True
).start()
log(f"listening on http://{HOST}:{port}")
return {"type": "config-result", "ok": True, "port": port}
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]:
if method not in ALLOWED_METHODS:
raise RuntimeError(f"method not allowed: {method}")
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
if message.get("type") == "config":
send_to_extension(apply_listen_port(message.get("port")))
continue
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:
threading.Thread(target=stdin_loop, name="fab-stdin", daemon=True).start()
apply_listen_port(DEFAULT_PORT)
try:
threading.Event().wait()
except KeyboardInterrupt:
return 0
return 0
if __name__ == "__main__":
raise SystemExit(main())