152 lines
5.4 KiB
Python
152 lines
5.4 KiB
Python
"""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
|