Initial Awawawa
This commit is contained in:
141
web/patchbay_web/validate.py
Normal file
141
web/patchbay_web/validate.py
Normal file
@@ -0,0 +1,141 @@
|
||||
"""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"}
|
||||
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]
|
||||
nodes.append({
|
||||
"id": nid, "type": t, "client_id": cid, "host": host, "port": _port(n.get("port")),
|
||||
"proto": proto, "label": label, "iface": iface,
|
||||
"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
|
||||
Reference in New Issue
Block a user