"""Input validation for client keys, users and the patch graph.""" import base64 import binascii import re import struct KEY_TYPES = { "ssh-ed25519", "ssh-rsa", "ecdsa-sha2-nistp256", "ecdsa-sha2-nistp384", "ecdsa-sha2-nistp521", "sk-ssh-ed25519@openssh.com", "sk-ecdsa-sha2-nistp256@openssh.com", } NODE_TYPES = {"client_source", "client_sink", "public_sink", "splitter", "tunnel_source", "tunnel_sink"} SOURCES = {"client_source", "tunnel_source"} HAS_OUTPUT = SOURCES | {"splitter"} HAS_INPUT = {"client_sink", "public_sink", "splitter", "tunnel_sink"} SINKS = HAS_INPUT - {"splitter"} # How a sink passes the peer address to the service. Whether "transparent" # has a source on a client is checked by the daemon (route error), so editing # the wiring never makes a save fail. ORIGINS = {"", "proxy_v2", "transparent"} IFACE_RE = re.compile(r"^[A-Za-z0-9_.\-]{1,15}$") HOST_RE = re.compile(r"^[A-Za-z0-9.:_\-\[\]%]{0,255}$") NAME_RE = re.compile(r"^[A-Za-z0-9._\-]{1,64}$") EMAIL_RE = re.compile(r"^[^@\s]+@[^@\s]+\.[^@\s]+$") class Invalid(ValueError): pass def pubkey(text): """Returns the normalised "type base64" form of an OpenSSH public key line.""" parts = text.strip().split() if len(parts) < 2: raise Invalid("expected an OpenSSH public key line (type base64 [comment])") ktype, b64 = parts[0], parts[1] if ktype not in KEY_TYPES: raise Invalid(f"unsupported key type {ktype}") try: blob = base64.b64decode(b64, validate=True) (n,) = struct.unpack(">I", blob[:4]) inner = blob[4:4 + n].decode() except (binascii.Error, struct.error, UnicodeDecodeError): raise Invalid("key data is not valid base64") from None if inner != ktype: raise Invalid("key type does not match key data") return f"{ktype} {b64}" def name(text, what="name"): if not NAME_RE.match(text or ""): raise Invalid(f"{what} may only contain letters, digits, '.', '_' and '-' (max 64)") return text def email(text): if not EMAIL_RE.match(text or "") or len(text) > 254: raise Invalid("invalid email address") return text def password(text): if len(text or "") < 10: raise Invalid("password must be at least 10 characters") return text def _port(v): try: p = int(v or 0) except (TypeError, ValueError): raise Invalid("port must be a number") from None if not 0 <= p <= 65535: raise Invalid("port out of range") return p def graph(data, client_ids): """Validates the editor payload. Returns (nodes, links) with clean values. Node ids are positive for existing nodes and negative for new ones. """ if not isinstance(data, dict): raise Invalid("bad payload") nodes, ids = [], set() for n in data.get("nodes", []): t = n.get("type") if t not in NODE_TYPES: raise Invalid(f"unknown node type {t}") nid = int(n.get("id", 0)) if nid == 0 or nid in ids: raise Invalid("duplicate or missing node id") ids.add(nid) cid = n.get("client_id") cid = int(cid) if cid not in (None, "", 0, "0") else None if cid is not None and cid not in client_ids: raise Invalid("unknown client") if t in ("public_sink", "splitter"): cid = None iface = str(n.get("iface") or "").strip() if t in ("tunnel_source", "tunnel_sink"): # client_id None = the target's own interface. if iface and not IFACE_RE.match(iface): raise Invalid(f"invalid interface name {iface!r}") else: iface = "" host = str(n.get("host") or "").strip() if not HOST_RE.match(host): raise Invalid(f"invalid address {host!r}") proto = n.get("proto", "tcp") if proto not in ("tcp", "udp"): raise Invalid("protocol must be tcp or udp") label = str(n.get("label") or "")[:64] origin = n.get("origin") or "" if origin not in ORIGINS: raise Invalid("origin must be proxy_v2 or transparent") if t not in SINKS: origin = "" nodes.append({ "id": nid, "type": t, "client_id": cid, "host": host, "port": _port(n.get("port")), "proto": proto, "label": label, "iface": iface, "origin": origin, "x": float(n.get("x", 0)), "y": float(n.get("y", 0)), }) types = {n["id"]: n["type"] for n in nodes} links, inputs, source_out = [], set(), set() for ln in data.get("links", []): a, b = int(ln.get("from", 0)), int(ln.get("to", 0)) if a not in types or b not in types or a == b: raise Invalid("link refers to unknown node") if types[a] not in HAS_OUTPUT or types[b] not in HAS_INPUT: raise Invalid("link direction invalid") if b in inputs: raise Invalid("an input can only have one link") if types[a] in SOURCES: if a in source_out: raise Invalid("a source output can only have one link; use a splitter") source_out.add(a) inputs.add(b) links.append((a, b)) # Reject cycles through splitters. parent = dict((b, a) for a, b in links) for start in types: seen, cur = set(), start while cur in parent: if cur in seen: raise Invalid("patch contains a loop") seen.add(cur) cur = parent[cur] return nodes, links