323 lines
11 KiB
Python
323 lines
11 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 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 1024–65535",
|
||
}
|
||
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())
|