This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
# app.py
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from fastapi import FastAPI, APIRouter, HTTPException, status
|
||||
from fastapi.responses import JSONResponse
|
||||
import logging
|
||||
@@ -16,39 +16,78 @@ logger = logging.getLogger("nft_api")
|
||||
class NftError(RuntimeError):
|
||||
pass
|
||||
|
||||
# ---------- NftManager (small, explicit) ----------
|
||||
# ---------- Helpers ----------
|
||||
_handle_re = re.compile(r"\bhandle\s+(\d+)\b")
|
||||
|
||||
|
||||
def _safe_load_json_unwrap(s: str) -> Any:
|
||||
"""
|
||||
Try to json.loads(s). If result is a string that itself contains JSON,
|
||||
keep unwrapping until we get a non-string or we fail.
|
||||
|
||||
Raises json.JSONDecodeError if initial parse fails.
|
||||
"""
|
||||
# If empty or None, raise
|
||||
if s is None:
|
||||
raise json.JSONDecodeError("empty", "None", 0)
|
||||
cur = s
|
||||
parsed = None
|
||||
# first parse: may throw
|
||||
parsed = json.loads(cur)
|
||||
# unwrap if the parsed result is itself a JSON string
|
||||
unwrap_count = 0
|
||||
while isinstance(parsed, str) and unwrap_count < 5:
|
||||
try:
|
||||
parsed = json.loads(parsed)
|
||||
unwrap_count += 1
|
||||
except json.JSONDecodeError:
|
||||
# can't unwrap further
|
||||
break
|
||||
return parsed
|
||||
|
||||
|
||||
# ---------- NftManager ----------
|
||||
class NftManager:
|
||||
"""
|
||||
Small wrapper around python-nftables:
|
||||
- cmd(text) -> (rc, stdout, stderr) textual
|
||||
- json_cmd(text) -> (rc, stdout, stderr) prefer json wrapper; fallback to -j
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.nft = Nftables()
|
||||
# best-effort: don't rely solely on set_json_output globally
|
||||
# best-effort: prefer JSON output where appropriate, but we'll explicitly choose cmd/json_cmd
|
||||
try:
|
||||
if hasattr(self.nft, "set_json_output"):
|
||||
# do not rely on global toggles for each call
|
||||
self.nft.set_json_output(True)
|
||||
except Exception:
|
||||
logger.debug("set_json_output not available")
|
||||
|
||||
def cmd(self, text_cmd: str) -> Tuple[int, str, str]:
|
||||
"""
|
||||
Run textual nft command via Nftables.cmd() and return (rc, stdout, stderr).
|
||||
Execute a textual nft command via Nftables.cmd()
|
||||
Returns (rc, stdout, stderr) where stdout/stderr are strings (possibly empty).
|
||||
"""
|
||||
rc, out, err = self.nft.cmd(text_cmd)
|
||||
return int(rc), out or "", err or ""
|
||||
|
||||
def json_cmd(self, text_cmd: str) -> Tuple[int, str, str]:
|
||||
"""
|
||||
Run nft expecting JSON output.
|
||||
Prefer Nftables.json_cmd() when present; otherwise call cmd() with -j appended.
|
||||
Returns (rc, stdout, stderr).
|
||||
Execute a command expecting JSON output.
|
||||
Prefer the wrapper json_cmd() if present, otherwise append -j to the command.
|
||||
Returns (rc, stdout, stderr). stdout is the raw text (may be already JSON or double-encoded JSON).
|
||||
"""
|
||||
if hasattr(self.nft, "json_cmd"):
|
||||
try:
|
||||
res = self.nft.json_cmd(text_cmd)
|
||||
# python-nftables json_cmd usually returns tuple (rc, stdout, stderr)
|
||||
if isinstance(res, (list, tuple)) and len(res) >= 3:
|
||||
return int(res[0]), res[1] or "", res[2] or ""
|
||||
except Exception as e:
|
||||
logger.debug("nft.json_cmd failed: %s (falling back to -j)", e)
|
||||
logger.debug("nft.json_cmd failed, falling back to -j: %s", e)
|
||||
|
||||
# fallback: append -j and use cmd()
|
||||
# fallback: append -j if not present
|
||||
cmd_with_j = text_cmd if "-j" in text_cmd else f"{text_cmd} -j"
|
||||
return self.cmd(cmd_with_j)
|
||||
|
||||
@@ -59,42 +98,55 @@ class NftManager:
|
||||
return out
|
||||
|
||||
def list_rules_json(self) -> Dict[str, Any]:
|
||||
"""
|
||||
Return parsed JSON (a Python dict/list) for `nft -j list ruleset`.
|
||||
Handle possible double-encoded outputs by unwrapping.
|
||||
"""
|
||||
rc, out, err = self.json_cmd("list ruleset")
|
||||
if rc != 0:
|
||||
raise NftError(f"nft list ruleset failed (json): {err}")
|
||||
try:
|
||||
return json.loads(out)
|
||||
parsed = _safe_load_json_unwrap(out)
|
||||
except json.JSONDecodeError as e:
|
||||
raise NftError(f"invalid JSON from nft: {e}")
|
||||
if not isinstance(parsed, (dict, list)):
|
||||
raise NftError("nft -j returned non-dict/list JSON")
|
||||
return parsed
|
||||
|
||||
def list_chain_text(self, family: str, table: str, chain: str) -> str:
|
||||
"""
|
||||
Prefer textual listing via cmd(). If that fails, use JSON fallback to reconstruct reasonable lines.
|
||||
Prefer textual listing via cmd(). If textual listing fails or returns JSON,
|
||||
attempt to use json_cmd() and reconstruct readable lines.
|
||||
"""
|
||||
cmd = f"list chain {family} {table} {chain}"
|
||||
rc, out, err = self.cmd(cmd)
|
||||
if rc == 0:
|
||||
if rc == 0 and out and not (out.strip().startswith("{") or out.strip().startswith("[")):
|
||||
# Looks like proper textual output
|
||||
return out
|
||||
|
||||
# fallback to JSON path and try to reconstruct
|
||||
# Either cmd returned error, or returned JSON-like text; try JSON path
|
||||
rcj, outj, errj = self.json_cmd(cmd)
|
||||
if rcj != 0:
|
||||
# prefer original textual error message
|
||||
# prefer the textual error if present
|
||||
raise NftError(f"nft {cmd} failed: {err or errj}")
|
||||
try:
|
||||
parsed = json.loads(outj)
|
||||
except Exception:
|
||||
return out # return whatever textual output we had (maybe empty)
|
||||
|
||||
lines: List[str] = []
|
||||
recs = parsed.get("nftables") if isinstance(parsed, dict) else (parsed or [])
|
||||
try:
|
||||
parsed = _safe_load_json_unwrap(outj)
|
||||
except json.JSONDecodeError:
|
||||
# if parsing fails, return whatever textual output we have (maybe empty)
|
||||
return out
|
||||
|
||||
# parsed expected to be dict { "nftables": [...] } or list
|
||||
recs = parsed.get("nftables") if isinstance(parsed, dict) else parsed
|
||||
if not isinstance(recs, list):
|
||||
return out
|
||||
|
||||
lines: List[str] = []
|
||||
for rec in recs:
|
||||
if "rule" in rec:
|
||||
r = rec["rule"]
|
||||
handle = r.get("handle")
|
||||
expr = r.get("expr")
|
||||
handle = r.get("handle")
|
||||
parts: List[str] = []
|
||||
if isinstance(expr, list):
|
||||
for el in expr:
|
||||
@@ -130,7 +182,7 @@ class NftManager:
|
||||
q = el["queue"]
|
||||
tok = "queue"
|
||||
if isinstance(q, dict):
|
||||
num = q.get("num") or q.get("number")
|
||||
num = q.get("num") or q.get("number") or q.get("queue_number") or q.get("from") or q.get("range")
|
||||
if num is not None:
|
||||
tok += f" num {num}"
|
||||
if q.get("bypass") or q.get("flags") == "bypass":
|
||||
@@ -150,16 +202,14 @@ class NftManager:
|
||||
return "\n".join(lines) if lines else out
|
||||
|
||||
|
||||
# ---------- utilities for predictable representation ----------
|
||||
_handle_re = re.compile(r"\bhandle\s+(\d+)\b")
|
||||
|
||||
def parse_priority(v: Any) -> Optional[int]:
|
||||
if v is None:
|
||||
# ---------- Utilities to convert nft JSON -> predictable model ----------
|
||||
def parse_priority(val: Any) -> Optional[int]:
|
||||
if val is None:
|
||||
return None
|
||||
if isinstance(v, int):
|
||||
return v
|
||||
if isinstance(v, str):
|
||||
s = v.strip()
|
||||
if isinstance(val, int):
|
||||
return val
|
||||
if isinstance(val, str):
|
||||
s = val.strip()
|
||||
try:
|
||||
return int(s)
|
||||
except Exception:
|
||||
@@ -167,47 +217,49 @@ def parse_priority(v: Any) -> Optional[int]:
|
||||
return int(float(s))
|
||||
except Exception:
|
||||
return None
|
||||
if isinstance(v, dict):
|
||||
for k in ("priority", "prio"):
|
||||
if k in v:
|
||||
return parse_priority(v.get(k))
|
||||
for val in v.values():
|
||||
p = parse_priority(val)
|
||||
if isinstance(val, dict):
|
||||
for key in ("priority", "prio"):
|
||||
if key in val:
|
||||
return parse_priority(val.get(key))
|
||||
for v in val.values():
|
||||
p = parse_priority(v)
|
||||
if p is not None:
|
||||
return p
|
||||
return None
|
||||
|
||||
|
||||
def rule_text_from_expr(expr: Any) -> str:
|
||||
"""
|
||||
Conservative, compact textualifier from nft JSON expr.
|
||||
If unknown tokens appear this returns a reasonable fallback (json dump).
|
||||
Conservative compact textualifier for expr.
|
||||
"""
|
||||
if expr is None:
|
||||
return ""
|
||||
if isinstance(expr, list):
|
||||
out: List[str] = []
|
||||
for el in expr:
|
||||
if isinstance(el, dict):
|
||||
if "match" in el:
|
||||
m = el["match"]; left = m.get("left"); right = m.get("right")
|
||||
tokens: List[str] = []
|
||||
for part in expr:
|
||||
if isinstance(part, dict):
|
||||
if "match" in part:
|
||||
m = part["match"]; left = m.get("left"); right = m.get("right")
|
||||
if isinstance(left, dict) and "payload" in left and isinstance(right, (str, int)):
|
||||
p = left["payload"]; prot = p.get("protocol"); field = p.get("field")
|
||||
if prot and field:
|
||||
out.append(f"{prot} {field} {right}"); continue
|
||||
out.append("match"); continue
|
||||
if "payload" in el:
|
||||
p = el["payload"]; prot = p.get("protocol"); field = p.get("field")
|
||||
tokens.append(f"{prot} {field} {right}")
|
||||
continue
|
||||
tokens.append("match")
|
||||
elif "payload" in part:
|
||||
p = part["payload"]; prot = p.get("protocol"); field = p.get("field")
|
||||
if prot and field:
|
||||
out.append(f"payload({prot}.{field})"); continue
|
||||
out.append("payload"); continue
|
||||
if "drop" in el:
|
||||
out.append("drop"); continue
|
||||
if "accept" in el:
|
||||
out.append("accept"); continue
|
||||
if "counter" in el:
|
||||
out.append("counter"); continue
|
||||
if "queue" in el:
|
||||
q = el["queue"]; tok = "queue"
|
||||
tokens.append(f"payload({prot}.{field})")
|
||||
continue
|
||||
tokens.append("payload")
|
||||
elif "drop" in part:
|
||||
tokens.append("drop")
|
||||
elif "accept" in part:
|
||||
tokens.append("accept")
|
||||
elif "counter" in part:
|
||||
tokens.append("counter")
|
||||
elif "queue" in part:
|
||||
q = part["queue"]; tok = "queue"
|
||||
if isinstance(q, dict):
|
||||
num = q.get("num") or q.get("number")
|
||||
if num is not None:
|
||||
@@ -216,22 +268,21 @@ def rule_text_from_expr(expr: Any) -> str:
|
||||
tok += " bypass"
|
||||
elif isinstance(q, (int, float)):
|
||||
tok += f" num {int(q)}"
|
||||
out.append(tok); continue
|
||||
# unknown dict -> show keys
|
||||
out.append("+".join(sorted(el.keys())))
|
||||
tokens.append(tok)
|
||||
else:
|
||||
tokens.append("+".join(sorted(part.keys())))
|
||||
else:
|
||||
out.append(str(el))
|
||||
return " ".join(out)
|
||||
# fallback
|
||||
tokens.append(str(part))
|
||||
return " ".join(tokens)
|
||||
try:
|
||||
return str(expr)
|
||||
except Exception:
|
||||
return json.dumps(expr)
|
||||
|
||||
|
||||
def build_predictable_ruleset(nft_json: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert nft -j ruleset into a deterministic dict:
|
||||
{ "tables": [ {family,name,chains:[{name,type,hook,priority,policy,rules:[{handle,expr,text,...}]}]} ] }
|
||||
Convert nft -j structure into deterministic dict with tables/chains/rules.
|
||||
"""
|
||||
result: Dict[str, Any] = {"tables": []}
|
||||
items = nft_json.get("nftables", []) if isinstance(nft_json, dict) else (nft_json or [])
|
||||
@@ -273,10 +324,9 @@ def build_predictable_ruleset(nft_json: Dict[str, Any]) -> Dict[str, Any]:
|
||||
tables.setdefault((fam, tname), {"family": fam, "name": tname, "chains": {}})
|
||||
chains = tables[(fam, tname)]["chains"]
|
||||
chains.setdefault(cname, {"name": cname, "type": None, "hook": None, "priority": None, "policy": None, "rules": []})
|
||||
rule_obj = {
|
||||
rule_obj: Dict[str, Any] = {
|
||||
"handle": r.get("handle"),
|
||||
"expr": r.get("expr"),
|
||||
# initial text from expr serializer (may be replaced later by exact textual line)
|
||||
"text": rule_text_from_expr(r.get("expr")),
|
||||
}
|
||||
if "position" in r:
|
||||
@@ -285,7 +335,7 @@ def build_predictable_ruleset(nft_json: Dict[str, Any]) -> Dict[str, Any]:
|
||||
rule_obj["comment"] = r["comment"]
|
||||
chains[cname]["rules"].append(rule_obj)
|
||||
|
||||
# salvage chain metadata if nested under rule (some nft outputs)
|
||||
# salvage nested chain metadata if present
|
||||
if isinstance(r.get("chain"), dict):
|
||||
csub = r.get("chain")
|
||||
if chains[cname].get("priority") is None:
|
||||
@@ -299,7 +349,7 @@ def build_predictable_ruleset(nft_json: Dict[str, Any]) -> Dict[str, Any]:
|
||||
if chains[cname].get("policy") is None and csub.get("policy") is not None:
|
||||
chains[cname]["policy"] = csub.get("policy")
|
||||
|
||||
# materialize into lists (sorted for deterministic order)
|
||||
# materialize lists deterministically
|
||||
for (fam, tname) in sorted(tables.keys(), key=lambda k: (k[0], k[1])):
|
||||
tdata = tables[(fam, tname)]
|
||||
chains_list: List[Dict[str, Any]] = []
|
||||
@@ -316,10 +366,10 @@ def build_predictable_ruleset(nft_json: Dict[str, Any]) -> Dict[str, Any]:
|
||||
result["tables"].append({"family": fam, "name": tname, "chains": chains_list})
|
||||
return result
|
||||
|
||||
|
||||
def enrich_text_from_chain(custom: Dict[str, Any], mgr: NftManager) -> None:
|
||||
"""
|
||||
Replace rule['text'] with exact textual lines from `nft list chain` when possible.
|
||||
Looks up rule handles; if not found it tries substring matching on the compact probe.
|
||||
"""
|
||||
for table in custom.get("tables", []):
|
||||
fam = table.get("family"); tname = table.get("name")
|
||||
@@ -331,13 +381,12 @@ def enrich_text_from_chain(custom: Dict[str, Any], mgr: NftManager) -> None:
|
||||
continue
|
||||
try:
|
||||
txt = mgr.list_chain_text(fam, tname, cname) or ""
|
||||
lines = [ln.strip() for ln in txt.splitlines() if ln.strip()]
|
||||
# map handle -> line
|
||||
handle_map = {}
|
||||
lines = [ln.rstrip() for ln in txt.splitlines() if ln.strip()]
|
||||
handle_map: Dict[str, str] = {}
|
||||
for ln in lines:
|
||||
m = _handle_re.search(ln)
|
||||
if m:
|
||||
handle_map[m.group(1)] = ln
|
||||
handle_map[m.group(1)] = ln.strip()
|
||||
for rule in ch.get("rules", []):
|
||||
replaced = False
|
||||
h = rule.get("handle")
|
||||
@@ -351,42 +400,42 @@ def enrich_text_from_chain(custom: Dict[str, Any], mgr: NftManager) -> None:
|
||||
if probe:
|
||||
for ln in lines:
|
||||
if probe in ln:
|
||||
rule["text"] = ln
|
||||
rule["text"] = ln.strip()
|
||||
replaced = True
|
||||
break
|
||||
except Exception as e:
|
||||
logger.debug("failed to enrich chain text for %s %s %s: %s", fam, tname, cname, e)
|
||||
logger.debug("enrich_text_from_chain failed for %s %s %s: %s", fam, tname, cname, e)
|
||||
continue
|
||||
|
||||
|
||||
# ---------- FastAPI app ----------
|
||||
app = FastAPI(title="nft API (clean)")
|
||||
app = FastAPI(title="nft API (robust unwrap)")
|
||||
router = APIRouter(prefix="/firewall", tags=["firewall"])
|
||||
mgr = NftManager()
|
||||
|
||||
@router.get("/rules")
|
||||
|
||||
@router.get("/rules", response_model=None)
|
||||
def get_rules():
|
||||
"""
|
||||
Return: JSONResponse({"ruleset": <dict-or-string>})
|
||||
- If nft -j is available -> returns native dict under "ruleset"
|
||||
- If not, returns the textual ruleset string under "ruleset"
|
||||
If nft -j is available -> returns native dict under "ruleset".
|
||||
Else -> returns textual ruleset string.
|
||||
"""
|
||||
try:
|
||||
try:
|
||||
nft_json = mgr.list_rules_json()
|
||||
except NftError as e:
|
||||
logger.debug("json listing unavailable: %s", e)
|
||||
# fallback to textual listing (return string; not JSON-encoded)
|
||||
text = mgr.list_rules_text()
|
||||
return JSONResponse(content={"ruleset": text.strip() if text is not None else None}, status_code=200)
|
||||
|
||||
custom = build_predictable_ruleset(nft_json)
|
||||
# best-effort: replace compact text with exact lines
|
||||
try:
|
||||
enrich_text_from_chain(custom, mgr)
|
||||
except Exception as e:
|
||||
logger.debug("enrich_text_from_chain failed: %s", e)
|
||||
|
||||
# Return native dict (no json.dumps)
|
||||
# Return native dict (no json.dumps, no double-encoding)
|
||||
return JSONResponse(content={"ruleset": custom}, status_code=200)
|
||||
except NftError as e:
|
||||
logger.exception("get_rules: nft error")
|
||||
@@ -395,27 +444,28 @@ def get_rules():
|
||||
logger.exception("get_rules internal")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@router.post("/raw")
|
||||
def exec_raw(cmd_body: Dict[str, Any]):
|
||||
"""
|
||||
Run raw textual nft command. Body: {"cmd": "add rule ..."}
|
||||
"""
|
||||
cmd = cmd_body.get("cmd")
|
||||
|
||||
@router.post("/raw", response_model=None)
|
||||
def exec_raw(body: Dict[str, Any]):
|
||||
cmd = body.get("cmd")
|
||||
if not isinstance(cmd, str) or not cmd.strip():
|
||||
raise HTTPException(status_code=400, detail="field 'cmd' required")
|
||||
rc, out, err = mgr.cmd(cmd)
|
||||
return JSONResponse(content={"rc": rc, "stdout": out or None, "stderr": err or None}, status_code=200 if rc == 0 else 400)
|
||||
status_code = 200 if rc == 0 else 400
|
||||
return JSONResponse(content={"rc": rc, "stdout": out or None, "stderr": err or None}, status_code=status_code)
|
||||
|
||||
@router.post("/rules")
|
||||
|
||||
@router.post("/rules", response_model=None, status_code=201)
|
||||
def create_rule(req: Dict[str, Any]):
|
||||
"""
|
||||
Minimal JSON-create endpoint: expects {family, table, chain, expr, [position]}
|
||||
Tries to render to textual fragment via rule_text_from_expr and runs add/insert.
|
||||
Minimal JSON->text create:
|
||||
{family, table, chain, expr, [position], [comment]}
|
||||
Renders expr via rule_text_from_expr; if it cannot render, instruct user to use /raw.
|
||||
"""
|
||||
family = req.get("family"); table = req.get("table"); chain = req.get("chain"); expr = req.get("expr")
|
||||
pos = req.get("position")
|
||||
if not (family and table and chain and expr is not None):
|
||||
raise HTTPException(status_code=400, detail="family,table,chain,expr required")
|
||||
raise HTTPException(status_code=400, detail="family, table, chain, expr required")
|
||||
rendered = rule_text_from_expr(expr)
|
||||
if rendered is None:
|
||||
raise HTTPException(status_code=400, detail="cannot render expr to textual rule; use /raw")
|
||||
@@ -431,19 +481,12 @@ def create_rule(req: Dict[str, Any]):
|
||||
else:
|
||||
cmd = f"add rule {family} {table} {chain} {expr_text}"
|
||||
rc, out, err = mgr.cmd(cmd)
|
||||
if rc == 0:
|
||||
return JSONResponse(content={"rc": rc, "stdout": out or None, "stderr": err or None}, status_code=201)
|
||||
# sometimes nft returns non-zero rc but no stderr while rule exists -> verify
|
||||
if (rc != 0) and not (err or "").strip():
|
||||
try:
|
||||
chain_text = mgr.list_chain_text(family, table, chain) or ""
|
||||
if expr_text and expr_text in chain_text:
|
||||
return JSONResponse(content={"rc": 0, "stdout": out or None, "stderr": err or None}, status_code=201)
|
||||
except Exception:
|
||||
pass
|
||||
# else error
|
||||
# if nft returned non-zero but produced no stderr and the rule exists -> treat as success
|
||||
if rc == 0 or (rc != 0 and not (err or "").strip() and expr_text in (mgr.list_chain_text(family, table, chain) or "")):
|
||||
return JSONResponse(content={"rc": 0, "stdout": out or None, "stderr": err or None}, status_code=201)
|
||||
raise HTTPException(status_code=400, detail=f"nft failed rc={rc} stderr={err!r} cmd={cmd}")
|
||||
|
||||
|
||||
@router.delete("/rules/{handle}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
def delete_rule(handle: int, family: str = "inet", table: str = "filter", chain: str = "input"):
|
||||
if not isinstance(handle, int) or handle <= 0:
|
||||
@@ -454,4 +497,5 @@ def delete_rule(handle: int, family: str = "inet", table: str = "filter", chain:
|
||||
raise HTTPException(status_code=500, detail=f"delete failed rc={rc} stderr={err}")
|
||||
return JSONResponse(status_code=204, content={})
|
||||
|
||||
|
||||
app.include_router(router)
|
||||
Reference in New Issue
Block a user